Skip to main content

competitive/data_structure/
lazy_segment_tree_map.rs

1use super::{FibHashMap, LazyMapMonoid, RangeBoundsExt};
2use std::{
3    fmt::{self, Debug, Formatter},
4    mem::replace,
5    ops::RangeBounds,
6};
7
8pub struct LazySegmentTreeMap<M>
9where
10    M: LazyMapMonoid,
11{
12    n: usize,
13    seg: FibHashMap<usize, (M::Agg, M::Act)>,
14}
15
16impl<M> Clone for LazySegmentTreeMap<M>
17where
18    M: LazyMapMonoid,
19{
20    fn clone(&self) -> Self {
21        Self {
22            n: self.n,
23            seg: self.seg.clone(),
24        }
25    }
26}
27
28impl<M> Debug for LazySegmentTreeMap<M>
29where
30    M: LazyMapMonoid<Agg: Debug, Act: Debug>,
31{
32    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
33        f.debug_struct("LazySegmentTreeMap")
34            .field("n", &self.n)
35            .field("seg", &self.seg)
36            .finish()
37    }
38}
39
40impl<M> LazySegmentTreeMap<M>
41where
42    M: LazyMapMonoid,
43{
44    pub fn new(n: usize) -> Self {
45        Self {
46            n,
47            seg: Default::default(),
48        }
49    }
50    #[inline]
51    fn get_mut(&mut self, k: usize) -> &mut (M::Agg, M::Act) {
52        self.seg.entry(k).or_insert((M::agg_unit(), M::act_unit()))
53    }
54    #[inline]
55    fn update_at(&mut self, k: usize, x: &M::Act) {
56        if M::is_act_unit(x) {
57            return;
58        }
59        let n = self.n;
60        let a = self.get_mut(k);
61        let nx = M::act_agg(&a.0, x);
62        if k < n {
63            a.1 = M::act_operate(&a.1, x);
64        }
65        if let Some(nx) = nx {
66            a.0 = nx;
67        } else if k < n {
68            self.propagate_at(k);
69            self.recalc_at(k);
70        } else {
71            panic!("act failed on leaf");
72        }
73    }
74    #[inline]
75    fn recalc_at(&mut self, k: usize) {
76        let x = match (self.seg.get(&(2 * k)), self.seg.get(&(2 * k + 1))) {
77            (None, None) => M::agg_unit(),
78            (None, Some((y, _))) => y.clone(),
79            (Some((x, _)), None) => x.clone(),
80            (Some((x, _)), Some((y, _))) => M::agg_operate(x, y),
81        };
82        self.get_mut(k).0 = x;
83    }
84    #[inline]
85    fn propagate_at(&mut self, k: usize) {
86        debug_assert!(k < self.n);
87        let x = match self.seg.get_mut(&k) {
88            Some((_, x)) => replace(x, M::act_unit()),
89            None => M::act_unit(),
90        };
91        if M::is_act_unit(&x) {
92            return;
93        }
94        self.update_at(2 * k, &x);
95        self.update_at(2 * k + 1, &x);
96    }
97    #[inline]
98    fn propagate(&mut self, k: usize, right: bool, nofilt: bool) {
99        let right = right as usize;
100        for i in (1..(k + 1 - right).next_power_of_two().trailing_zeros()).rev() {
101            if nofilt || (k >> i) << i != k {
102                self.propagate_at((k - right) >> i);
103            }
104        }
105    }
106    #[inline]
107    fn recalc(&mut self, k: usize, right: bool, nofilt: bool) {
108        let right = right as usize;
109        for i in 1..(k + 1 - right).next_power_of_two().trailing_zeros() {
110            if nofilt || (k >> i) << i != k {
111                self.recalc_at((k - right) >> i);
112            }
113        }
114    }
115    pub fn update<R>(&mut self, range: R, x: M::Act)
116    where
117        R: RangeBounds<usize>,
118    {
119        let range = range.to_range_bounded(0, self.n).expect("invalid range");
120        if M::is_act_unit(&x) {
121            return;
122        }
123        let mut a = range.start + self.n;
124        let mut b = range.end + self.n;
125        self.propagate(a, false, false);
126        self.propagate(b, true, false);
127        while a < b {
128            if a & 1 != 0 {
129                self.update_at(a, &x);
130                a += 1;
131            }
132            if b & 1 != 0 {
133                b -= 1;
134                self.update_at(b, &x);
135            }
136            a /= 2;
137            b /= 2;
138        }
139        self.recalc(range.start + self.n, false, false);
140        self.recalc(range.end + self.n, true, false);
141    }
142    pub fn fold<R>(&mut self, range: R) -> M::Agg
143    where
144        R: RangeBounds<usize>,
145    {
146        let range = range.to_range_bounded(0, self.n).expect("invalid range");
147        let mut l = range.start + self.n;
148        let mut r = range.end + self.n;
149        self.propagate(l, false, true);
150        self.propagate(r, true, true);
151        let mut vl = M::agg_unit();
152        let mut vr = M::agg_unit();
153        while l < r {
154            if l & 1 != 0 {
155                if let Some((x, _)) = self.seg.get(&l) {
156                    vl = M::agg_operate(&vl, x);
157                }
158                l += 1;
159            }
160            if r & 1 != 0 {
161                r -= 1;
162                if let Some((x, _)) = self.seg.get(&r) {
163                    vr = M::agg_operate(x, &vr);
164                }
165            }
166            l /= 2;
167            r /= 2;
168        }
169        M::agg_operate(&vl, &vr)
170    }
171    pub fn set(&mut self, k: usize, x: M::Agg) {
172        let k = k + self.n;
173        self.propagate(k, false, true);
174        *self.get_mut(k) = (x, M::act_unit());
175        self.recalc(k, false, true);
176    }
177    pub fn get(&mut self, k: usize) -> M::Agg {
178        assert!(k < self.n);
179        let k = k + self.n;
180        self.propagate(k, false, true);
181        self.seg
182            .get(&k)
183            .map(|(x, _)| x.clone())
184            .unwrap_or_else(M::agg_unit)
185    }
186    pub fn fold_all(&mut self) -> M::Agg {
187        self.fold(0..self.n)
188    }
189    fn partition_point_perfect<P>(
190        &mut self,
191        mut pos: usize,
192        mut acc: M::Agg,
193        mut pred: P,
194    ) -> (usize, M::Agg)
195    where
196        P: FnMut(&M::Agg) -> bool,
197    {
198        while pos < self.n {
199            self.propagate_at(pos);
200            pos <<= 1;
201            let nacc = match self.seg.get(&pos) {
202                Some((x, _)) => M::agg_operate(&acc, x),
203                None => acc.clone(),
204            };
205            if pred(&nacc) {
206                acc = nacc;
207                pos += 1;
208            }
209        }
210        (pos - self.n, acc)
211    }
212    fn rpartition_point_perfect<P>(
213        &mut self,
214        mut pos: usize,
215        mut acc: M::Agg,
216        mut pred: P,
217    ) -> (usize, M::Agg)
218    where
219        P: FnMut(&M::Agg) -> bool,
220    {
221        while pos < self.n {
222            self.propagate_at(pos);
223            pos = pos * 2 + 1;
224            let nacc = match self.seg.get(&pos) {
225                Some((x, _)) => M::agg_operate(x, &acc),
226                None => acc.clone(),
227            };
228            if pred(&nacc) {
229                acc = nacc;
230                pos -= 1;
231            }
232        }
233        (pos - self.n, acc)
234    }
235    pub fn partition_point_acc<P>(&mut self, left: usize, mut pred: P) -> usize
236    where
237        P: FnMut(&M::Agg) -> bool,
238    {
239        let mut acc = M::agg_unit();
240        if left == self.n {
241            return self.n;
242        }
243        let mut l = left + self.n;
244        let r = 2 * self.n;
245        self.propagate(l, false, true);
246        self.propagate(r, true, true);
247        let mut k = 0usize;
248        while l < r >> k {
249            if l & 1 != 0 {
250                let nacc = match self.seg.get(&l) {
251                    Some((x, _)) => M::agg_operate(&acc, x),
252                    None => acc.clone(),
253                };
254                if !pred(&nacc) {
255                    return self.partition_point_perfect(l, acc, pred).0;
256                }
257                acc = nacc;
258                l += 1;
259            }
260            l >>= 1;
261            k += 1;
262        }
263        for k in (0..k).rev() {
264            let r = r >> k;
265            if r & 1 != 0 {
266                let nacc = match self.seg.get(&(r - 1)) {
267                    Some((x, _)) => M::agg_operate(&acc, x),
268                    None => acc.clone(),
269                };
270                if !pred(&nacc) {
271                    return self.partition_point_perfect(r - 1, acc, pred).0;
272                }
273                acc = nacc;
274            }
275        }
276        self.n
277    }
278    pub fn rpartition_point_acc<P>(&mut self, right: usize, mut pred: P) -> usize
279    where
280        P: FnMut(&M::Agg) -> bool,
281    {
282        let mut acc = M::agg_unit();
283        if right == 0 {
284            return 0;
285        }
286        let mut l = self.n;
287        let mut r = right + self.n;
288        self.propagate(l, false, true);
289        self.propagate(r, true, true);
290        let mut c = 0usize;
291        let mut k = 0usize;
292        while l >> k < r {
293            c <<= 1;
294            if l & (1 << k) != 0 {
295                l += 1 << k;
296                c += 1;
297            }
298            if r & 1 != 0 {
299                r -= 1;
300                let nacc = match self.seg.get(&r) {
301                    Some((x, _)) => M::agg_operate(x, &acc),
302                    None => acc.clone(),
303                };
304                if !pred(&nacc) {
305                    return self.rpartition_point_perfect(r, acc, pred).0 + 1;
306                }
307                acc = nacc;
308            }
309            r >>= 1;
310            k += 1;
311        }
312        for k in (0..k).rev() {
313            if c & 1 != 0 {
314                l -= 1 << k;
315                let l = l >> k;
316                let nacc = match self.seg.get(&l) {
317                    Some((x, _)) => M::agg_operate(x, &acc),
318                    None => acc.clone(),
319                };
320                if !pred(&nacc) {
321                    return self.rpartition_point_perfect(l, acc, pred).0 + 1;
322                }
323                acc = nacc;
324            }
325            c >>= 1;
326        }
327        0
328    }
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334    use crate::{
335        algebra::{RangeMaxRangeUpdate, RangeSumRangeAdd},
336        rand,
337        tools::{NotEmptySegment, Xorshift},
338    };
339
340    const N: usize = 1_000;
341    const Q: usize = 20_000;
342    const A: i64 = 1_000_000_000;
343
344    #[test]
345    fn test_lazy_segment_tree_map() {
346        let mut rng = Xorshift::default();
347        // Range Sum Query & Range Add Query
348        let mut arr = vec![0i64; N];
349        let mut seg = LazySegmentTreeMap::<RangeSumRangeAdd<_>>::new(N);
350        for i in 0..N {
351            seg.set(i, (0i64, 1i64));
352        }
353        for _ in 0..Q {
354            rand!(rng, (l, r): NotEmptySegment(N));
355            match rng.rand(3) {
356                0 => {
357                    // Range Add Query
358                    rand!(rng, x: -A..A);
359                    seg.update(l..r, x);
360                    for a in arr[l..r].iter_mut() {
361                        *a += x;
362                    }
363                }
364                1 => {
365                    // Point Set Query
366                    rand!(rng, k: 0..N, x: -A..A);
367                    seg.set(k, (x, 1));
368                    arr[k] = x;
369                }
370                _ => {
371                    // Range Sum Query
372                    let res = arr[l..r].iter().sum();
373                    assert_eq!(seg.fold(l..r).0, res);
374                }
375            }
376            rand!(rng, k: 0..N);
377            assert_eq!(seg.get(k).0, arr[k]);
378            assert_eq!(seg.fold_all().0, arr.iter().sum());
379        }
380
381        // Range Max Query & Range Update Query & Binary Search Query
382        let mut arr = vec![i64::MIN; N];
383        let mut seg = LazySegmentTreeMap::<RangeMaxRangeUpdate<_>>::new(N);
384        for _ in 0..Q {
385            rand!(rng, ty: 0..5, (l, r): NotEmptySegment(N));
386            match ty {
387                0 => {
388                    // Range Update Query
389                    rand!(rng, x: -A..A);
390                    seg.update(l..r, Some(x));
391                    arr[l..r].iter_mut().for_each(|a| *a = x);
392                }
393                1 => {
394                    // Range Max Query
395                    let res = arr[l..r].iter().max().cloned().unwrap_or_default();
396                    assert_eq!(seg.fold(l..r), res);
397                }
398                2 => {
399                    // Binary Search Query
400                    rand!(rng, left: ..=N, x: -A..A);
401                    assert_eq!(
402                        seg.partition_point_acc(left, |&d| d < x),
403                        arr[left..]
404                            .iter()
405                            .scan(i64::MIN, |acc, &a| {
406                                *acc = a.max(*acc);
407                                Some(*acc)
408                            })
409                            .position(|acc| acc >= x)
410                            .map_or(N, |i| i + left),
411                    );
412                }
413                3 => {
414                    // Binary Search Query
415                    rand!(rng, right: ..=N, x: -A..A);
416                    assert_eq!(
417                        seg.rpartition_point_acc(right, |&d| d < x),
418                        arr[..right]
419                            .iter()
420                            .rev()
421                            .scan(i64::MIN, |acc, &a| {
422                                *acc = a.max(*acc);
423                                Some(*acc)
424                            })
425                            .position(|acc| acc >= x)
426                            .map_or(0, |i| right - i),
427                    );
428                }
429                _ => {
430                    // Point Set Query
431                    rand!(rng, k: 0..N, x: -A..A);
432                    seg.set(k, x);
433                    arr[k] = x;
434                }
435            }
436            rand!(rng, k: 0..N);
437            assert_eq!(seg.get(k), arr[k]);
438            assert_eq!(seg.fold_all(), *arr.iter().max().unwrap());
439        }
440    }
441}