Skip to main content

competitive/data_structure/
binary_indexed_tree.rs

1use super::{AbelianGroup, Group, Monoid};
2use std::fmt::{self, Debug, Formatter};
3
4pub struct BinaryIndexedTree<M>
5where
6    M: Monoid,
7{
8    n: usize,
9    bit: Vec<M::T>,
10}
11
12impl<M> Clone for BinaryIndexedTree<M>
13where
14    M: Monoid,
15{
16    fn clone(&self) -> Self {
17        Self {
18            n: self.n,
19            bit: self.bit.clone(),
20        }
21    }
22}
23
24impl<M> Debug for BinaryIndexedTree<M>
25where
26    M: Monoid<T: Debug>,
27{
28    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
29        f.debug_struct("BinaryIndexedTree")
30            .field("n", &self.n)
31            .field("bit", &self.bit)
32            .finish()
33    }
34}
35
36impl<M> BinaryIndexedTree<M>
37where
38    M: Monoid,
39{
40    #[inline]
41    pub fn new(n: usize) -> Self {
42        let bit = vec![M::unit(); n + 1];
43        Self { n, bit }
44    }
45    #[inline]
46    pub fn from_slice(slice: &[M::T]) -> Self {
47        let n = slice.len();
48        let mut bit = vec![M::unit(); n + 1];
49        for (i, x) in slice.iter().enumerate() {
50            let k = i + 1;
51            M::operate_assign(&mut bit[k], x);
52            let j = k + (k & (!k + 1));
53            if j <= n {
54                bit[j] = M::operate(&bit[j], &bit[k]);
55            }
56        }
57        Self { n, bit }
58    }
59    #[inline]
60    /// fold [0, k)
61    pub fn accumulate0(&self, mut k: usize) -> M::T {
62        debug_assert!(k <= self.n);
63        let mut res = M::unit();
64        while k > 0 {
65            res = M::operate(&res, &self.bit[k]);
66            k -= k & (!k + 1);
67        }
68        res
69    }
70    #[inline]
71    /// fold [0, k]
72    pub fn accumulate(&self, k: usize) -> M::T {
73        self.accumulate0(k + 1)
74    }
75    #[inline]
76    pub fn update(&mut self, k: usize, x: M::T) {
77        debug_assert!(k < self.n);
78        let mut k = k + 1;
79        while k <= self.n {
80            self.bit[k] = M::operate(&self.bit[k], &x);
81            k += k & (!k + 1);
82        }
83    }
84    #[inline]
85    pub fn partition_point_acc<P>(&self, mut pred: P) -> usize
86    where
87        P: FnMut(&M::T) -> bool,
88    {
89        let n = self.n;
90        let mut acc = M::unit();
91        let mut pos = 0;
92        let mut k = n.next_power_of_two();
93        while k > 0 {
94            if k + pos <= n {
95                let nacc = M::operate(&acc, &self.bit[k + pos]);
96                if pred(&nacc) {
97                    pos += k;
98                    acc = nacc;
99                }
100            }
101            k >>= 1;
102        }
103        pos
104    }
105}
106
107impl<G: Group> BinaryIndexedTree<G> {
108    #[inline]
109    pub fn fold(&self, l: usize, r: usize) -> G::T {
110        debug_assert!(l <= self.n && r <= self.n);
111        G::operate(&G::inverse(&self.accumulate0(l)), &self.accumulate0(r))
112    }
113    #[inline]
114    pub fn fold_abelian(&self, mut l: usize, mut r: usize) -> G::T
115    where
116        G: AbelianGroup,
117    {
118        debug_assert!(l <= self.n && r <= self.n);
119        if l == r {
120            return G::unit();
121        }
122        // The prefix above the highest differing bit cancels in an Abelian group.
123        let common = l & !(usize::MAX >> (l ^ r).leading_zeros());
124        let mut left = G::unit();
125        let mut right = G::unit();
126        while l != common {
127            G::operate_assign(&mut left, &self.bit[l]);
128            l &= l - 1;
129        }
130        while r != common {
131            G::operate_assign(&mut right, &self.bit[r]);
132            r &= r - 1;
133        }
134        G::rinv_operate(&right, &left)
135    }
136    #[inline]
137    pub fn get(&self, k: usize) -> G::T {
138        self.fold(k, k + 1)
139    }
140    #[inline]
141    pub fn set(&mut self, k: usize, x: G::T) {
142        self.update(k, G::operate(&G::inverse(&self.get(k)), &x));
143    }
144}
145
146#[cfg(test)]
147mod tests {
148    use super::*;
149    use crate::{
150        algebra::{AdditiveOperation, MaxOperation},
151        tools::Xorshift,
152    };
153
154    const N: usize = 10_000;
155    const Q: usize = 100_000;
156    const A: u64 = 1_000_000_000;
157    const B: i64 = 1_000_000_000;
158
159    #[test]
160    fn test_binary_indexed_tree() {
161        let mut rng = Xorshift::default();
162        let mut arr: Vec<_> = rng.random_iter(..A).take(N).collect();
163        let mut bit = BinaryIndexedTree::<AdditiveOperation<_>>::from_slice(&arr);
164        for (k, v) in rng.random_iter((..N, ..A)).take(Q) {
165            bit.update(k, v);
166            arr[k] += v;
167        }
168        for i in 0..N - 1 {
169            arr[i + 1] += arr[i];
170        }
171        for (i, a) in arr.iter().cloned().enumerate() {
172            assert_eq!(bit.accumulate(i), a);
173        }
174
175        let mut arr: Vec<_> = rng.random_iter(..A).take(N).collect();
176        let mut bit = BinaryIndexedTree::<MaxOperation<_>>::from_slice(&arr);
177        for (k, v) in rng.random_iter((..N, ..A)).take(Q) {
178            bit.update(k, v);
179            arr[k] = std::cmp::max(arr[k], v);
180        }
181        for i in 0..N - 1 {
182            arr[i + 1] = std::cmp::max(arr[i], arr[i + 1]);
183        }
184        for (i, a) in arr.iter().cloned().enumerate() {
185            assert_eq!(bit.accumulate(i), a);
186        }
187    }
188
189    #[test]
190    fn test_group_binary_indexed_tree() {
191        const N: usize = 2_000;
192        let mut rng = Xorshift::default();
193        let mut arr: Vec<_> = rng.random_iter(-B..B).take(N).collect();
194        let mut bit = BinaryIndexedTree::<AdditiveOperation<_>>::from_slice(&arr);
195        for (k, v) in rng.random_iter((..N, -B..B)).take(Q) {
196            bit.set(k, v);
197            arr[k] = v;
198        }
199        for i in 0..N - 1 {
200            arr[i + 1] += arr[i];
201        }
202        for i in 0..=N {
203            for j in i..=N {
204                let expected =
205                    if j == 0 { 0 } else { arr[j - 1] } - if i == 0 { 0 } else { arr[i - 1] };
206                assert_eq!(bit.fold(i, j), expected);
207                assert_eq!(bit.fold_abelian(i, j), expected);
208            }
209        }
210    }
211
212    #[test]
213    fn test_binary_indexed_tree_partition_point_acc() {
214        let mut rng = Xorshift::default();
215        let mut arr: Vec<_> = rng.random_iter(1..B).take(N).collect();
216        let mut bit = BinaryIndexedTree::<AdditiveOperation<_>>::from_slice(&arr);
217        for (k, v) in rng.random_iter((..N, 1..B)).take(Q) {
218            bit.set(k, v);
219            arr[k] = v;
220        }
221        for i in 0..N - 1 {
222            arr[i + 1] += arr[i];
223        }
224        for x in rng.random_iter(1..B * N as i64).take(Q) {
225            assert_eq!(
226                bit.partition_point_acc(|&a| a < x),
227                arr.partition_point(|&a| a < x)
228            );
229        }
230    }
231}