Use this to learn the idea, then write your own version.
12 3 4import math5from bisect import bisect_left, bisect_right6from typing import Generic, Iterable, Iterator, TypeVar, Union, List7 8T = TypeVar("T")9 10 11class SortedSet(Generic[T]):12 """Sorted set (set) in C++.13 14 See:15 https:qiita.com/tatyam/items/492c70ac4c955c05560216 https:github.com/tatyam-prime/SortedSet/blob/main/SortedSet.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 SortedSet from iterable.36 / O(N) if sorted and unique / O(N log N)37 """38 a = list(a)39 40 if not all(a[i] < a[i + 1] for i in range(len(a) - 1)): 41 a = sorted(set(a)) 42 43 self._build(a)44 45 def __iter__(self) -> Iterator[T]:46 for i in self.a:47 for j in i:48 yield j 49 50 def __reversed__(self) -> Iterator[T]:51 for i in reversed(self.a):52 for j in reversed(i):53 yield j54 55 def __len__(self) -> int:56 return self.size57 58 def __repr__(self) -> str:59 return "SortedSet" + str(self.a)60 61 def __str__(self) -> str:62 s = str(list(self))63 return "{" + s[1 : len(s) - 1] + "}"64 65 def _find_bucket(self, x: T) -> List[T]:66 "Find the bucket which should contain x. self must not be empty."67 for a in self.a:68 if x <= a[-1]: 69 return a70 return a71 72 def __contains__(self, x: T) -> bool:73 if self.size == 0:74 return False75 76 a = self._find_bucket(x)77 i = bisect_left(a, x) 78 79 return i != len(a) and a[i] == x80 81 def add(self, x: T) -> bool:82 "Add an element and return True if added. / O(√N)"83 if self.size == 0:84 self.a = [[x]]85 self.size = 186 return True87 88 a = self._find_bucket(x)89 i = bisect_left(a, x) 90 91 if i != len(a) and a[i] == x:92 return False93 94 a.insert(i, x)95 self.size += 196 97 if len(a) > len(self.a) * self.REBUILD_RATIO:98 self._build()99 100 return True101 102 def discard(self, x: T) -> bool:103 "Remove an element and return True if removed. / O(√N)"104 if self.size == 0:105 return False106 107 a = self._find_bucket(x)108 i = bisect_left(a, x) 109 110 if i == len(a) or a[i] != x:111 return False112 113 a.pop(i)114 self.size -= 1115 116 if len(a) == 0:117 self._build()118 return True119 120 def lt(self, x: T) -> Union[T, None]:121 "Find the largest element < x, or None if it doesn't exist."122 for a in reversed(self.a):123 if a[0] < x: 124 return a[bisect_left(a, x) - 1] 125 return None126 127 def le(self, x: T) -> Union[T, None]:128 "Find the largest element <= x, or None if it doesn't exist."129 for a in reversed(self.a):130 if a[0] <= x: 131 return a[bisect_right(a, x) - 1] 132 return None133 134 def gt(self, x: T) -> Union[T, None]:135 "Find the smallest element > x, or None if it doesn't exist."136 for a in self.a:137 if a[-1] > x: 138 return a[bisect_right(a, x)] 139 return None140 141 def ge(self, x: T) -> Union[T, None]:142 "Find the smallest element >= x, or None if it doesn't exist."143 for a in self.a:144 if a[-1] >= x: 145 return a[bisect_left(a, x)] 146 return None147 148 def __getitem__(self, x: int) -> T:149 "Return the x-th element, or IndexError if it doesn't exist."150 if x < 0:151 x += self.size152 if x < 0:153 raise IndexError154 155 for a in self.a:156 if x < len(a):157 return a[x] 158 159 x -= len(a)160 raise IndexError161 162 def index(self, x: T) -> int:163 "Count the number of elements < x."164 ans = 0165 166 for a in self.a:167 if a[-1] >= x: 168 return ans + bisect_left(a, x) 169 ans += len(a)170 return ans171 172 def index_right(self, x: T) -> int:173 "Count the number of elements <= x."174 ans = 0175 176 for a in self.a:177 if a[-1] > x: 178 return ans + bisect_right(a, x) 179 ans += len(a)180 return ans181 182 183def main():184 import sys185 from collections import defaultdict186 187 input = sys.stdin.readline188 189 n = int(input())190 inf = 10**18191 x = [0] + list(map(int, input().split())) + [inf]192 193 d_min = defaultdict(int)194 d_min[x[0]] = x[1]195 d_min[inf] = inf196 197 ids = defaultdict(int)198 199 for i, xi in enumerate(x):200 ids[xi] = i201 202 s = SortedSet([0, inf])203 prev_sum = 0204 205 for i, xi in enumerate(x[1:-1]):206 cur_sum = prev_sum207 208 p1, p2 = s.lt(xi), s.gt(xi)209 id1, id2 = ids[p1], ids[p2]210 211 if i != 0:212 cur_sum -= d_min[p1]213 214 if p2 != inf:215 cur_sum -= d_min[p2]216 217 d_min[p1] = min(d_min[p1], abs(xi - x[id1]))218 d_min[p2] = min(d_min[p2], abs(x[id2] - xi))219 d_min[xi] = min(abs(xi - x[id1]), abs(x[id2] - xi))220 221 cur_sum += d_min[p1] + d_min[xi]222 223 if p2 != inf:224 cur_sum += d_min[p2]225 226 print(cur_sum)227 228 s.add(xi)229 prev_sum = cur_sum230 231 232if __name__ == "__main__":233 main()234