Use this to learn the idea, then write your own version.
12 3 4import math5from bisect import bisect_left, bisect_right, insort6from typing import Generic, Iterable, Iterator, List, TypeVar, Union7 8T = TypeVar("T")9 10 11class SortedMultiset(Generic[T]):12 """Sorted multi set (set) in C++.13 14 See:15 https:qiita.com/tatyam/items/492c70ac4c955c05560216 https:github.com/tatyam-prime/SortedSet/blob/main/SortedMultiset.py17 """18 19 BUCKET_RATIO = 5020 REBUILD_RATIO = 17021 22 def _build(self, a=None) -> None:23 "Evenly divide `a` into buckets."24 if a is None:25 a = list(self)26 27 size = self.size = len(a)28 bucket_size = int(math.ceil(math.sqrt(size / self.BUCKET_RATIO)))29 self.a = [30 a[size * i bucket_size : size * (i + 1) bucket_size]31 for i in range(bucket_size)32 ]33 34 def __init__(self, a: Iterable[T] = []) -> None:35 "Make a new SortedMultiset from iterable. / O(N) if sorted / O(N log N)"36 a = list(a)37 38 if not all(a[i] <= a[i + 1] for i in range(len(a) - 1)): 39 a = sorted(a) 40 41 self._build(a)42 43 def __iter__(self) -> Iterator[T]:44 for i in self.a:45 for j in i:46 yield j 47 48 def __reversed__(self) -> Iterator[T]:49 for i in reversed(self.a):50 for j in reversed(i):51 yield j52 53 def __len__(self) -> int:54 return self.size55 56 def __repr__(self) -> str:57 return "SortedMultiset" + str(self.a)58 59 def __str__(self) -> str:60 s = str(list(self))61 return "{" + s[1 : len(s) - 1] + "}"62 63 def _find_bucket(self, x: T) -> List[T]:64 "Find the bucket which should contain x. self must not be empty."65 for a in self.a:66 if x <= a[-1]: 67 return a68 return a 69 70 def __contains__(self, x: T) -> bool:71 if self.size == 0:72 return False73 74 a = self._find_bucket(x)75 i = bisect_left(a, x) 76 return i != len(a) and a[i] == x77 78 def count(self, x: T) -> int:79 "Count the number of x."80 return self.index_right(x) - self.index(x)81 82 def add(self, x: T) -> None:83 "Add an element. / O(√N)"84 if self.size == 0:85 self.a = [[x]]86 self.size = 187 return88 89 a = self._find_bucket(x)90 insort(a, x) 91 self.size += 192 93 if len(a) > len(self.a) * self.REBUILD_RATIO:94 self._build()95 96 def discard(self, x: T) -> bool:97 "Remove an element and return True if removed. / O(√N)"98 if self.size == 0:99 return False100 101 a = self._find_bucket(x)102 i = bisect_left(a, x) 103 104 if i == len(a) or a[i] != x:105 return False106 107 a.pop(i)108 self.size -= 1109 110 if len(a) == 0:111 self._build()112 113 return True114 115 def lt(self, x: T) -> Union[T, None]:116 "Find the largest element < x, or None if it doesn't exist."117 for a in reversed(self.a):118 if a[0] < x: 119 return a[bisect_left(a, x) - 1] 120 return None121 122 def le(self, x: T) -> Union[T, None]:123 "Find the largest element <= x, or None if it doesn't exist."124 for a in reversed(self.a):125 if a[0] <= x: 126 return a[bisect_right(a, x) - 1] 127 return None128 129 def gt(self, x: T) -> Union[T, None]:130 "Find the smallest element > x, or None if it doesn't exist."131 for a in self.a:132 if a[-1] > x: 133 return a[bisect_right(a, x)] 134 return None135 136 def ge(self, x: T) -> Union[T, None]:137 "Find the smallest element >= x, or None if it doesn't exist."138 for a in self.a:139 if a[-1] >= x: 140 return a[bisect_left(a, x)] 141 return None142 143 def __getitem__(self, x: int) -> T:144 "Return the x-th element, or IndexError if it doesn't exist."145 if x < 0:146 x += self.size147 if x < 0:148 raise IndexError149 150 for a in self.a:151 if x < len(a):152 return a[x] 153 154 x -= len(a)155 raise IndexError156 157 def index(self, x: T) -> int:158 "Count the number of elements < x."159 ans = 0160 161 for a in self.a:162 if a[-1] >= x: 163 return ans + bisect_left(a, x) 164 ans += len(a)165 return ans166 167 def index_right(self, x: T) -> int:168 "Count the number of elements <= x."169 ans = 0170 171 for a in self.a:172 if a[-1] > x: 173 return ans + bisect_right(a, x) 174 ans += len(a)175 return ans176 177 178def main():179 import sys180 181 input = sys.stdin.readline182 183 n, m = map(int, input().split())184 p = list(map(int, input().split()))185 inf = 10**12186 s = SortedMultiset(p + [inf])187 l = list(map(int, input().split()))188 d = list(map(int, input().split()))189 dl = sorted([(di, li) for di, li in zip(d, l)], reverse=True)190 191 192 193 ans = sum(p)194 195 for di, li in dl:196 value = s.ge(li)197 198 if value != inf:199 s.discard(value)200 ans -= di201 202 print(ans)203 204 205if __name__ == "__main__":206 main()207