Skip to main content

competitive/data_structure/
segment_tree.rs

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