Skip to main content

competitive/data_structure/
compressed_binary_indexed_tree.rs

1use super::Monoid;
2use std::{
3    fmt::{self, Debug},
4    marker::PhantomData,
5    ops::{Bound, RangeBounds},
6};
7
8pub struct CompressedBinaryIndexedTree<M, X, Inner>
9where
10    M: Monoid,
11{
12    compress: Vec<X>,
13    bits: Vec<Inner>,
14    _marker: PhantomData<fn() -> M>,
15}
16impl<M, X, Inner> Debug for CompressedBinaryIndexedTree<M, X, Inner>
17where
18    M: Monoid,
19    X: Debug,
20    Inner: Debug,
21{
22    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
23        f.debug_struct("CompressedBinaryIndexedTree")
24            .field("compress", &self.compress)
25            .field("bits", &self.bits)
26            .finish()
27    }
28}
29impl<M, X, Inner> Clone for CompressedBinaryIndexedTree<M, X, Inner>
30where
31    M: Monoid,
32    X: Clone,
33    Inner: Clone,
34{
35    fn clone(&self) -> Self {
36        Self {
37            compress: self.compress.clone(),
38            bits: self.bits.clone(),
39            _marker: self._marker,
40        }
41    }
42}
43impl<M, X, Inner> Default for CompressedBinaryIndexedTree<M, X, Inner>
44where
45    M: Monoid,
46{
47    fn default() -> Self {
48        Self {
49            compress: Default::default(),
50            bits: Default::default(),
51            _marker: Default::default(),
52        }
53    }
54}
55#[repr(transparent)]
56pub struct Tag<M>(M::T)
57where
58    M: Monoid;
59impl<M> Debug for Tag<M>
60where
61    M: Monoid<T: Debug>,
62{
63    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64        self.0.fmt(f)
65    }
66}
67impl<M> Clone for Tag<M>
68where
69    M: Monoid,
70{
71    fn clone(&self) -> Self {
72        Self(self.0.clone())
73    }
74}
75
76macro_rules! impl_compressed_binary_indexed_tree {
77    (@tuple ($($l:tt)*) ($($r:tt)*) $T:ident) => {
78        ($($l)* $T $($r)*,)
79    };
80    (@tuple ($($l:tt)*) ($($r:tt)*) $T:ident $($Rest:ident)+) => {
81        ($($l)* $T $($r)*, impl_compressed_binary_indexed_tree!(@tuple ($($l)*) ($($r)*) $($Rest)+))
82    };
83    (@cst $M:ident) => {
84        Tag<$M>
85    };
86    (@cst $M:ident $T:ident $($Rest:ident)*) => {
87        CompressedBinaryIndexedTree<$M, $T, impl_compressed_binary_indexed_tree!(@cst $M $($Rest)*)>
88    };
89    (@from_iter $M:ident $points:ident $T:ident) => {{
90        let mut compress: Vec<_> = $points.into_iter().map(|t| t.0.clone()).collect();
91        compress.sort_unstable();
92        compress.dedup();
93        let n = compress.len();
94        Self {
95            compress,
96            bits: vec![Tag(M::unit()); n + 1],
97            _marker: PhantomData,
98        }
99    }};
100    (@from_iter $M:ident $points:ident $T:ident $U:ident $($Rest:ident)*) => {{
101        let mut points: Vec<_> = $points.into_iter().collect();
102        points.sort_unstable_by(|left, right| left.0.cmp(&right.0));
103        let mut compress = Vec::new();
104        let mut offsets = vec![0];
105        let mut start = 0;
106        while start < points.len() {
107            let mut end = start + 1;
108            while end < points.len() && points[end].0 == points[start].0 {
109                end += 1;
110            }
111            compress.push(points[start].0.clone());
112            offsets.push(end);
113            start = end;
114        }
115        let n = compress.len();
116        let mut bits: Vec<impl_compressed_binary_indexed_tree!(@cst $M $U $($Rest)*)> =
117            vec![Default::default(); n + 1];
118        for i in 1..=n {
119            let start = i - (i & (!i + 1));
120            bits[i] = <impl_compressed_binary_indexed_tree!(@cst $M $U $($Rest)*)>::from_iter(
121                points[offsets[start]..offsets[i]]
122                    .iter()
123                    .map(|point| &point.1),
124            );
125        }
126        Self {
127            compress,
128            bits,
129            _marker: PhantomData,
130        }
131    }};
132    (@acc $e:expr, $rng:ident $T:ident) => {
133        $e.0
134    };
135    (@acc $e:expr, $rng:ident $T:ident $($Rest:ident)+) => {
136        $e.accumulate(&$rng.1)
137    };
138    (@update $e:expr, $M:ident $key:ident $x:ident $T:ident) => {
139        $M::operate_assign(&mut $e.0, $x);
140    };
141    (@update $e:expr, $M:ident $key:ident $x:ident $T:ident $($Rest:ident)+) => {
142        $e.update(&$key.1, $x);
143    };
144    (@partition_method $T:ident, $Q:ident) => {
145        pub fn partition_point_acc<P>(&self, mut pred: P) -> (Option<&$T>, M::T)
146        where
147            P: FnMut(&M::T) -> bool,
148        {
149            let n = self.compress.len();
150            let mut acc = M::unit();
151            let mut pos = 0;
152            let mut k = n.next_power_of_two();
153            if k > n {
154                k >>= 1;
155            }
156            while k > 0 {
157                if k + pos <= n {
158                    let nacc = M::operate(&acc, &self.bits[k + pos].0);
159                    if pred(&nacc) {
160                        pos += k;
161                        acc = nacc;
162                    }
163                }
164                k >>= 1;
165            }
166            (self.compress.get(pos), acc)
167        }
168    };
169    (@partition_method $T:ident $($RestT:ident)+, $Q:ident $($RestQ:ident)+) => {
170        pub fn partition_point_acc<P, $($RestQ,)*>(
171            &self,
172            inner_ranges: &impl_compressed_binary_indexed_tree!(@tuple () () $($RestQ)*),
173            mut pred: P,
174        ) -> (Option<&$T>, M::T)
175        where
176            P: FnMut(&M::T) -> bool,
177            $($RestQ: RangeBounds<$RestT>,)*
178        {
179            let n = self.compress.len();
180            let mut acc = M::unit();
181            let mut pos = 0;
182            let mut k = n.next_power_of_two();
183            if k > n {
184                k >>= 1;
185            }
186            while k > 0 {
187                if k + pos <= n {
188                    let nacc = M::operate(
189                        &acc,
190                        &self.bits[k + pos].accumulate(inner_ranges),
191                    );
192                    if pred(&nacc) {
193                        pos += k;
194                        acc = nacc;
195                    }
196                }
197                k >>= 1;
198            }
199            (self.compress.get(pos), acc)
200        }
201    };
202    (@impl $C:ident $($T:ident)*, $($Q:ident)*) => {
203        impl<M, $($T,)*> impl_compressed_binary_indexed_tree!(@cst M $($T)*)
204        where
205            M: Monoid,
206            $($T: Clone + Ord,)*
207        {
208            pub fn new(points: &[impl_compressed_binary_indexed_tree!(@tuple () () $($T)*)]) -> Self {
209                Self::from_iter(points)
210            }
211            fn from_iter<'a, Iter>(points: Iter) -> Self
212            where
213                $($T: 'a,)*
214                Iter: IntoIterator<Item = &'a impl_compressed_binary_indexed_tree!(@tuple () () $($T)*)> + Clone,
215            {
216                impl_compressed_binary_indexed_tree!(@from_iter M points $($T)*)
217            }
218            pub fn accumulate<$($Q,)*>(&self, range: &impl_compressed_binary_indexed_tree!(@tuple () () $($Q)*)) -> M::T
219            where
220                $($Q: RangeBounds<$T>,)*
221            {
222                match range.0.start_bound() {
223                    Bound::Unbounded => (),
224                    _ => panic!("expected `Bound::Unbounded`"),
225                };
226                let mut k = match range.0.end_bound() {
227                    Bound::Included(index) => self.compress.partition_point(|x| x <= index),
228                    Bound::Excluded(index) => self.compress.partition_point(|x| x < index),
229                    Bound::Unbounded => self.compress.len(),
230                };
231                let mut x = M::unit();
232                while k > 0 {
233                    x = M::operate(&x, &impl_compressed_binary_indexed_tree!(@acc self.bits[k], range $($T)*));
234                    k -= k & (!k + 1);
235                }
236                x
237            }
238            pub fn update(&mut self, key: &impl_compressed_binary_indexed_tree!(@tuple () () $($T)*), x: &M::T) {
239                let mut k = self.compress.binary_search(&key.0).expect("not exist key") + 1;
240                while k < self.bits.len() {
241                    impl_compressed_binary_indexed_tree!(@update self.bits[k], M key x $($T)*);
242                    k += k & (!k + 1);
243                }
244            }
245            impl_compressed_binary_indexed_tree!(@partition_method $($T)*, $($Q)*);
246        }
247        pub type $C<M, $($T),*> = impl_compressed_binary_indexed_tree!(@cst M $($T)*);
248    };
249    (@inner [$C:ident][$($T:ident)*][$($Q:ident)*][]) => {
250        impl_compressed_binary_indexed_tree!(@impl $C $($T)*, $($Q)*);
251    };
252    (@inner [$C:ident][$($T:ident)*][$($Q:ident)*][$D:ident $U:ident $R:ident $($Rest:ident)*]) => {
253        impl_compressed_binary_indexed_tree!(@impl $C $($T)*, $($Q)*);
254        impl_compressed_binary_indexed_tree!(@inner [$D][$($T)* $U][$($Q)* $R][$($Rest)*]);
255    };
256    ($C:ident $T:ident $Q:ident $($Rest:ident)* $(;$($t:tt)*)?) => {
257        impl_compressed_binary_indexed_tree!(@inner [$C][$T][$Q][$($Rest)*]);
258    };
259    ($($t:tt)*) => {
260        compile_error!($($t:tt)*)
261    }
262}
263
264impl_compressed_binary_indexed_tree!(
265    CompressedBinaryIndexedTree1d A QA
266    CompressedBinaryIndexedTree2d B QB
267    CompressedBinaryIndexedTree3d C QC
268    CompressedBinaryIndexedTree4d D QD;
269    CompressedBinaryIndexedTree5d E QE
270    CompressedBinaryIndexedTree6d F QF
271    CompressedBinaryIndexedTree7d G QG
272    CompressedBinaryIndexedTree8d H QH
273    CompressedBinaryIndexedTree9d I QI
274);
275
276#[cfg(test)]
277mod tests {
278    use super::*;
279    use crate::{algebra::AdditiveOperation, tools::Xorshift};
280    use std::{collections::HashMap, ops::RangeTo};
281
282    #[test]
283    fn test_bit1d() {
284        let mut rng = Xorshift::default();
285        const N: usize = 100;
286        const Q: usize = 5000;
287        const A: RangeTo<u64> = ..1_000;
288        let mut points: Vec<_> = rng.random_iter(A).take(N).map(|x| (x,)).collect();
289        points.sort();
290        points.dedup();
291        let mut values: HashMap<_, _> = points.iter().map(|p| (p.0, 0u64)).collect();
292        let mut bit = CompressedBinaryIndexedTree1d::<AdditiveOperation<u64>, _>::new(&points);
293        for _ in 0..Q {
294            let p = &points[rng.random(0..points.len())];
295            let x = rng.random(A);
296            *values.get_mut(&p.0).unwrap() += x;
297            bit.update(p, &x);
298
299            let range = ((
300                Bound::Unbounded,
301                match rng.rand(3) {
302                    0 => Bound::Excluded(rng.random(A)),
303                    1 => Bound::Included(rng.random(A)),
304                    _ => Bound::Unbounded,
305                },
306            ),);
307            let expected: u64 = values
308                .iter()
309                .filter_map(|(p, x)| RangeBounds::contains(&range.0, p).then_some(*x))
310                .sum();
311            assert_eq!(bit.accumulate(&range), expected);
312
313            let target = rng.random(1..A.end * Q as u64);
314            let mut expected_acc = 0;
315            let mut expected_pos = None;
316            for p in &points {
317                let nacc = expected_acc + values[&p.0];
318                if nacc < target {
319                    expected_acc = nacc;
320                } else {
321                    expected_pos = Some(&p.0);
322                    break;
323                }
324            }
325            let result = bit.partition_point_acc(|&acc| acc < target);
326            assert_eq!(result, (expected_pos, expected_acc));
327        }
328    }
329
330    #[test]
331    fn test_bit2d_and_4d() {
332        let mut rng = Xorshift::default();
333        for _ in 0..12 {
334            let domain = rng.rand(128) + 1;
335            let point_count = rng.rand(96) as usize + 1;
336            let registered: Vec<_> = rng
337                .random_iter((..domain, (..domain,)))
338                .take(point_count)
339                .collect();
340            let mut points = registered.clone();
341            points.sort_unstable();
342            points.dedup();
343            let mut values: HashMap<_, _> =
344                points.iter().copied().map(|point| (point, 0u64)).collect();
345            let mut bit =
346                CompressedBinaryIndexedTree2d::<AdditiveOperation<u64>, _, _>::new(&registered);
347            let query_count = rng.rand(300) as usize + 300;
348            for _ in 0..query_count {
349                let point = &points[rng.rand(points.len() as u64) as usize];
350                let value = rng.rand(domain);
351                *values.get_mut(point).unwrap() += value;
352                bit.update(point, &value);
353
354                let end_x = rng.rand(domain + 1);
355                let end_y = rng.rand(domain + 1);
356                let expected = values
357                    .iter()
358                    .filter_map(|((x, (y,)), value)| (*x < end_x && *y < end_y).then_some(*value))
359                    .sum();
360                assert_eq!(bit.accumulate(&(..end_x, (..end_y,))), expected);
361            }
362        }
363
364        const N: usize = 100;
365        const Q: usize = 5000;
366        const A: RangeTo<u64> = ..1_000;
367        let mut points: Vec<_> = rng.random_iter(((A), (A, (A, (A,))))).take(N).collect();
368        points.sort();
369        points.dedup();
370        let mut map: HashMap<_, _> = points.iter().map(|p| (p, 0u64)).collect();
371        let mut bit =
372            CompressedBinaryIndexedTree4d::<AdditiveOperation<u64>, _, _, _, _>::new(&points);
373        for _ in 0..Q {
374            let p = &points[rng.random(0..points.len())];
375            let x = rng.random(A);
376            *map.get_mut(p).unwrap() += x;
377            bit.update(p, &x);
378
379            let mut f = || {
380                (
381                    Bound::Unbounded,
382                    match rng.rand(3) {
383                        0 => Bound::Excluded(rng.random(A)),
384                        1 => Bound::Included(rng.random(A)),
385                        _ => Bound::Unbounded,
386                    },
387                )
388            };
389
390            let range = (f(), (f(), (f(), (f(),))));
391            let (r0, (r1, (r2, (r3,)))) = range;
392            let expected: u64 = map
393                .iter()
394                .filter_map(|((p0, (p1, (p2, (p3,)))), x)| {
395                    if RangeBounds::contains(&r0, p0)
396                        && RangeBounds::contains(&r1, p1)
397                        && RangeBounds::contains(&r2, p2)
398                        && RangeBounds::contains(&r3, p3)
399                    {
400                        Some(*x)
401                    } else {
402                        None
403                    }
404                })
405                .sum();
406            let result = bit.accumulate(&range);
407            assert_eq!(expected, result);
408
409            let target = rng.random(1..A.end * Q as u64);
410            let (_, inner_ranges) = &range;
411            let (r1, (r2, (r3,))) = inner_ranges;
412            let mut expected_acc = 0;
413            let mut expected_pos = None;
414            for p0 in &bit.compress {
415                let value: u64 = map
416                    .iter()
417                    .filter_map(|((q0, (q1, (q2, (q3,)))), x)| {
418                        if q0 == p0
419                            && RangeBounds::contains(r1, q1)
420                            && RangeBounds::contains(r2, q2)
421                            && RangeBounds::contains(r3, q3)
422                        {
423                            Some(*x)
424                        } else {
425                            None
426                        }
427                    })
428                    .sum();
429                let nacc = expected_acc + value;
430                if nacc < target {
431                    expected_acc = nacc;
432                } else {
433                    expected_pos = Some(p0);
434                    break;
435                }
436            }
437            let (pos, acc) = bit.partition_point_acc(inner_ranges, |&acc| acc < target);
438            assert_eq!((pos, acc), (expected_pos, expected_acc));
439        }
440    }
441}