Skip to main content

competitive/data_structure/
lazy_segment_tree.rs

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