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, TypeVar, Union, List7T = TypeVar('T')8 9 10class SortedMultiset(Generic[T]):11 """Sorted multi set (set) in C++.12 See:13 https:qiita.com/tatyam/items/492c70ac4c955c05560214 https:github.com/tatyam-prime/SortedSet/blob/main/SortedMultiset.py15 """16 17 BUCKET_RATIO = 5018 REBUILD_RATIO = 17019 20 def _build(self, a=None) -> None:21 "Evenly divide `a` into buckets."22 if a is None:23 a = list(self)24 25 size = self.size = len(a)26 bucket_size = int(math.ceil(math.sqrt(size / self.BUCKET_RATIO)))27 self.a = [a[size * i bucket_size: size * (i + 1) bucket_size] for i in range(bucket_size)]28 29 def __init__(self, a: Iterable[T] = []) -> None:30 "Make a new SortedMultiset from iterable. / O(N) if sorted / O(N log N)"31 a = list(a)32 33 if not all(a[i] <= a[i + 1] for i in range(len(a) - 1)): 34 a = sorted(a) 35 36 self._build(a)37 38 def __iter__(self) -> Iterator[T]:39 for i in self.a:40 for j in i:41 yield j 42 43 def __reversed__(self) -> Iterator[T]:44 for i in reversed(self.a):45 for j in reversed(i):46 yield j47 48 def __len__(self) -> int:49 return self.size50 51 def __repr__(self) -> str:52 return "SortedMultiset" + str(self.a)53 54 def __str__(self) -> str:55 s = str(list(self))56 return "{" + s[1: len(s) - 1] + "}"57 58 def _find_bucket(self, x: T) -> List[T]:59 "Find the bucket which should contain x. self must not be empty."60 for a in self.a:61 if x <= a[-1]: 62 return a63 return a 64 65 def __contains__(self, x: T) -> bool:66 if self.size == 0:67 return False68 69 a = self._find_bucket(x)70 i = bisect_left(a, x) 71 return i != len(a) and a[i] == x72 73 def count(self, x: T) -> int:74 "Count the number of x."75 return self.index_right(x) - self.index(x)76 77 def add(self, x: T) -> None:78 "Add an element. / O(√N)"79 if self.size == 0:80 self.a = [[x]]81 self.size = 182 return83 84 a = self._find_bucket(x)85 insort(a, x) 86 self.size += 187 88 if len(a) > len(self.a) * self.REBUILD_RATIO:89 self._build()90 91 def discard(self, x: T) -> bool:92 "Remove an element and return True if removed. / O(√N)"93 if self.size == 0:94 return False95 96 a = self._find_bucket(x)97 i = bisect_left(a, x) 98 99 if i == len(a) or a[i] != x:100 return False101 102 a.pop(i)103 self.size -= 1104 105 if len(a) == 0:106 self._build()107 108 return True109 110 def lt(self, x: T) -> Union[T, None]:111 "Find the largest element < x, or None if it doesn't exist."112 for a in reversed(self.a):113 if a[0] < x: 114 return a[bisect_left(a, x) - 1] 115 return None116 117 def le(self, x: T) -> Union[T, None]:118 "Find the largest element <= x, or None if it doesn't exist."119 for a in reversed(self.a):120 if a[0] <= x: 121 return a[bisect_right(a, x) - 1] 122 return None123 124 def gt(self, x: T) -> Union[T, None]:125 "Find the smallest element > x, or None if it doesn't exist."126 for a in self.a:127 if a[-1] > x: 128 return a[bisect_right(a, x)] 129 return None130 131 def ge(self, x: T) -> Union[T, None]:132 "Find the smallest element >= x, or None if it doesn't exist."133 for a in self.a:134 if a[-1] >= x: 135 return a[bisect_left(a, x)] 136 return None137 138 def __getitem__(self, x: int) -> T:139 "Return the x-th element, or IndexError if it doesn't exist."140 if x < 0:141 x += self.size142 if x < 0:143 raise IndexError144 145 for a in self.a:146 if x < len(a):147 return a[x] 148 149 x -= len(a)150 raise IndexError151 152 def index(self, x: T) -> int:153 "Count the number of elements < x."154 ans = 0155 156 for a in self.a:157 if a[-1] >= x: 158 return ans + bisect_left(a, x) 159 ans += len(a)160 return ans161 162 def index_right(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_right(a, x) 169 ans += len(a)170 return ans171 172 173def calc(b, x, y):174 inf = 10 ** 9175 s_min = SortedMultiset([inf])176 s_max = SortedMultiset([inf])177 178 for i, bi in enumerate(b):179 if bi == y:180 s_min.add(i)181 if bi == x:182 s_max.add(i)183 184 size = len(b)185 results = 0186 187 for i in range(size):188 value_min = s_min.ge(i)189 value_max = s_max.ge(i)190 191 if value_min == inf or value_max == inf:192 continue193 194 results += size - max(value_min, value_max)195 return results196 197 198def main():199 import sys200 201 input = sys.stdin.readline202 203 n, x, y = map(int, input().split())204 a = list(map(int, input().split()))205 206 i = 0207 ans = 0208 209 while i < n:210 b = []211 212 while i < n:213 if y <= a[i] <= x:214 b.append(a[i])215 else:216 break217 218 i += 1219 220 ans += calc(b, x, y)221 i += 1222 223 print(ans)224 225 226if __name__ == "__main__":227 main()228