Skip to main content

competitive/data_structure/
segment_tree_map.rs

1use 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}