competitive/data_structure/
segment_tree_map.rs1use super::{AbelianMonoid, FibHashMap, Monoid, RangeBoundsExt};
2use std::{
3 fmt::{self, Debug, Formatter},
4 ops::RangeBounds,
5};
6
7pub struct SegmentTreeMap<M>
8where
9 M: Monoid,
10{
11 n: usize,
12 seg: FibHashMap<usize, M::T>,
13 u: M::T,
14}
15
16impl<M> Clone for SegmentTreeMap<M>
17where
18 M: Monoid,
19{
20 fn clone(&self) -> Self {
21 Self {
22 n: self.n,
23 seg: self.seg.clone(),
24 u: self.u.clone(),
25 }
26 }
27}
28
29impl<M> Debug for SegmentTreeMap<M>
30where
31 M: Monoid<T: Debug>,
32{
33 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
34 f.debug_struct("SegmentTreeMap")
35 .field("n", &self.n)
36 .field("seg", &self.seg)
37 .field("u", &self.u)
38 .finish()
39 }
40}
41
42impl<M> SegmentTreeMap<M>
43where
44 M: Monoid,
45{
46 pub fn new(n: usize) -> Self {
47 let u = M::unit();
48 Self {
49 n,
50 seg: Default::default(),
51 u,
52 }
53 }
54 #[inline]
55 fn get_ref(&self, k: usize) -> &M::T {
56 self.seg.get(&k).unwrap_or(&self.u)
57 }
58 pub fn set(&mut self, k: usize, x: M::T) {
59 debug_assert!(k < self.n);
60 let mut k = k + self.n;
61 *self.seg.entry(k).or_insert(M::unit()) = x;
62 k /= 2;
63 while k > 0 {
64 *self.seg.entry(k).or_insert(M::unit()) =
65 M::operate(self.get_ref(2 * k), self.get_ref(2 * k + 1));
66 k /= 2;
67 }
68 }
69 pub fn update(&mut self, k: usize, x: M::T) {
70 debug_assert!(k < self.n);
71 let mut k = k + self.n;
72 let t = self.seg.entry(k).or_insert(M::unit());
73 *t = M::operate(t, &x);
74 k /= 2;
75 while k > 0 {
76 *self.seg.entry(k).or_insert(M::unit()) =
77 M::operate(self.get_ref(2 * k), self.get_ref(2 * k + 1));
78 k /= 2;
79 }
80 }
81 pub fn get(&self, k: usize) -> M::T {
82 debug_assert!(k < self.n);
83 self.seg.get(&(k + self.n)).cloned().unwrap_or_else(M::unit)
84 }
85 pub fn fold<R>(&self, range: R) -> M::T
86 where
87 R: RangeBounds<usize>,
88 {
89 let range = range.to_range();
90 debug_assert!(range.end <= self.n);
91 let mut l = range.start + self.n;
92 let mut r = range.end + self.n;
93 let mut vl = M::unit();
94 let mut vr = M::unit();
95 while l < r {
96 if l & 1 != 0 {
97 vl = M::operate(&vl, self.get_ref(l));
98 l += 1;
99 }
100 if r & 1 != 0 {
101 r -= 1;
102 vr = M::operate(self.get_ref(r), &vr);
103 }
104 l /= 2;
105 r /= 2;
106 }
107 M::operate(&vl, &vr)
108 }
109 fn partition_point_perfect<P>(
110 &self,
111 mut pos: usize,
112 mut acc: M::T,
113 mut pred: P,
114 ) -> (usize, M::T)
115 where
116 P: FnMut(&M::T) -> bool,
117 {
118 while pos < self.n {
119 pos <<= 1;
120 let nacc = M::operate(&acc, self.get_ref(pos));
121 if pred(&nacc) {
122 acc = nacc;
123 pos += 1;
124 }
125 }
126 (pos - self.n, acc)
127 }
128 fn rpartition_point_perfect<P>(
129 &self,
130 mut pos: usize,
131 mut acc: M::T,
132 mut pred: P,
133 ) -> (usize, M::T)
134 where
135 P: FnMut(&M::T) -> bool,
136 {
137 while pos < self.n {
138 pos = pos * 2 + 1;
139 let nacc = M::operate(self.get_ref(pos), &acc);
140 if pred(&nacc) {
141 acc = nacc;
142 pos -= 1;
143 }
144 }
145 (pos - self.n, acc)
146 }
147 pub fn partition_point_acc<P>(&self, left: usize, mut pred: P) -> usize
148 where
149 P: FnMut(&M::T) -> bool,
150 {
151 let mut l = left + self.n;
152 let r = 2 * self.n;
153 let mut k = 0usize;
154 let mut acc = M::unit();
155 while l < r >> k {
156 if l & 1 != 0 {
157 let nacc = M::operate(&acc, self.get_ref(l));
158 if !pred(&nacc) {
159 return self.partition_point_perfect(l, acc, pred).0;
160 }
161 acc = nacc;
162 l += 1;
163 }
164 l >>= 1;
165 k += 1;
166 }
167 for k in (0..k).rev() {
168 let r = r >> k;
169 if r & 1 != 0 {
170 let nacc = M::operate(&acc, self.get_ref(r - 1));
171 if !pred(&nacc) {
172 return self.partition_point_perfect(r - 1, acc, pred).0;
173 }
174 acc = nacc;
175 }
176 }
177 self.n
178 }
179 pub fn rpartition_point_acc<P>(&self, right: usize, mut pred: P) -> usize
180 where
181 P: FnMut(&M::T) -> bool,
182 {
183 let mut l = self.n;
184 let mut r = right + self.n;
185 let mut c = 0usize;
186 let mut k = 0usize;
187 let mut acc = M::unit();
188 while l >> k < r {
189 c <<= 1;
190 if l & (1 << k) != 0 {
191 l += 1 << k;
192 c += 1;
193 }
194 if r & 1 != 0 {
195 r -= 1;
196 let nacc = M::operate(self.get_ref(r), &acc);
197 if !pred(&nacc) {
198 return self.rpartition_point_perfect(r, acc, pred).0 + 1;
199 }
200 acc = nacc;
201 }
202 r >>= 1;
203 k += 1;
204 }
205 for k in (0..k).rev() {
206 if c & 1 != 0 {
207 l -= 1 << k;
208 let l = l >> k;
209 let nacc = M::operate(self.get_ref(l), &acc);
210 if !pred(&nacc) {
211 return self.rpartition_point_perfect(l, acc, pred).0 + 1;
212 }
213 acc = nacc;
214 }
215 c >>= 1;
216 }
217 0
218 }
219}
220
221impl<M> SegmentTreeMap<M>
222where
223 M: AbelianMonoid,
224{
225 pub fn fold_all(&self) -> M::T {
226 self.seg.get(&1).cloned().unwrap_or_else(M::unit)
227 }
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233 use crate::{
234 algebra::{AdditiveOperation, MaxOperation},
235 algorithm::SliceBisectExt as _,
236 rand,
237 tools::{NotEmptySegment as Nes, Xorshift},
238 };
239
240 const N: usize = 1_000;
241 const Q: usize = 10_000;
242 const A: i64 = 1_000_000_000;
243
244 #[test]
245 fn test_segment_tree_map() {
246 let mut rng = Xorshift::default();
247 let mut arr = vec![0; N + 1];
248 let mut seg = SegmentTreeMap::<AdditiveOperation<_>>::new(N);
249 for (k, v) in rng.random_iter((..N, 1..=A)).take(Q) {
250 seg.set(k, v);
251 arr[k + 1] = v;
252 }
253 for i in 0..N {
254 arr[i + 1] += arr[i];
255 }
256 for i in 0..N {
257 for j in i + 1..N + 1 {
258 assert_eq!(seg.fold(i..j), arr[j] - arr[i]);
259 }
260 }
261 for (left, v) in rng.random_iter((..=N, 1..=A * N as i64)).take(Q) {
262 assert_eq!(
263 seg.partition_point_acc(left, |&x| x < v),
264 arr[left + 1..].position_bisect(|&x| x - arr[left] >= v) + left
265 );
266 }
267 for (right, v) in rng.random_iter((..=N, 1..=A)).take(Q) {
268 assert_eq!(
269 seg.rpartition_point_acc(right, |&x| x < v),
270 arr[..right].rposition_bisect(|&x| arr[right] - x >= v)
271 );
272 }
273
274 rand!(rng, mut arr: [-A..=A; N]);
275 let mut seg = SegmentTreeMap::<MaxOperation<_>>::new(N);
276 for (i, a) in arr.iter().cloned().enumerate() {
277 seg.set(i, a);
278 }
279 for (k, v) in rng.random_iter((..N, -A..=A)).take(Q) {
280 seg.set(k, v);
281 arr[k] = v;
282 }
283 for (l, r) in rng.random_iter(Nes(N)).take(Q) {
284 let res = arr[l..r].iter().max().cloned().unwrap_or_default();
285 assert_eq!(seg.fold(l..r), res);
286 }
287 }
288}