Skip to main content

competitive/data_structure/
binary_indexed_tree_2d.rs

1use super::{Group, Monoid};
2use std::fmt::{self, Debug, Formatter};
3
4pub struct BinaryIndexedTree2D<M>
5where
6    M: Monoid,
7{
8    h: usize,
9    w: usize,
10    bit: Vec<M::T>,
11}
12
13impl<M> Clone for BinaryIndexedTree2D<M>
14where
15    M: Monoid,
16{
17    fn clone(&self) -> Self {
18        Self {
19            h: self.h,
20            w: self.w,
21            bit: self.bit.clone(),
22        }
23    }
24}
25
26impl<M> Debug for BinaryIndexedTree2D<M>
27where
28    M: Monoid<T: Debug>,
29{
30    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
31        f.debug_struct("BinaryIndexedTree2D")
32            .field("h", &self.h)
33            .field("w", &self.w)
34            .field("bit", &self.bit)
35            .finish()
36    }
37}
38
39impl<M> BinaryIndexedTree2D<M>
40where
41    M: Monoid,
42{
43    #[inline]
44    pub fn new(h: usize, w: usize) -> Self {
45        let bit = vec![M::unit(); (h + 1) * (w + 1)];
46        Self { h, w, bit }
47    }
48    #[inline]
49    /// fold [0, i) x [0, j)
50    pub fn accumulate0(&self, i: usize, j: usize) -> M::T {
51        assert!(i <= self.h && j <= self.w);
52        let mut res = M::unit();
53        let mut a = i;
54        let stride = self.w + 1;
55        while a > 0 {
56            let mut b = j;
57            while b > 0 {
58                // SAFETY: the method validates both prefix endpoints, and Fenwick ancestors
59                // remain within the allocated `(h + 1) * (w + 1)` table.
60                M::operate_assign(&mut res, unsafe { self.bit.get_unchecked(a * stride + b) });
61                b -= b & (!b + 1);
62            }
63            a -= a & (!a + 1);
64        }
65        res
66    }
67    #[inline]
68    /// fold [0, i] x [0, j]
69    pub fn accumulate(&self, i: usize, j: usize) -> M::T {
70        self.accumulate0(i + 1, j + 1)
71    }
72    #[inline]
73    pub fn update(&mut self, i: usize, j: usize, x: M::T) {
74        assert!(i < self.h && j < self.w);
75        let mut a = i + 1;
76        let stride = self.w + 1;
77        while a <= self.h {
78            let mut b = j + 1;
79            while b <= self.w {
80                // SAFETY: the method validates the leaf, and Fenwick ancestors stay within the
81                // allocated table.
82                M::operate_assign(unsafe { self.bit.get_unchecked_mut(a * stride + b) }, &x);
83                b += b & (!b + 1);
84            }
85            a += a & (!a + 1);
86        }
87    }
88}
89
90impl<G> BinaryIndexedTree2D<G>
91where
92    G: Group,
93{
94    #[inline]
95    /// 0-indexed [i1, i2) x [j1, j2)
96    pub fn fold(&self, i1: usize, j1: usize, i2: usize, j2: usize) -> G::T {
97        let mut res = self.accumulate0(i1, j1);
98        G::rinv_operate_assign(&mut res, &self.accumulate0(i1, j2));
99        G::rinv_operate_assign(&mut res, &self.accumulate0(i2, j1));
100        G::operate_assign(&mut res, &self.accumulate0(i2, j2));
101        res
102    }
103    #[inline]
104    pub fn get(&self, i: usize, j: usize) -> G::T {
105        self.fold(i, j, i + 1, j + 1)
106    }
107    #[inline]
108    pub fn set(&mut self, i: usize, j: usize, x: G::T) {
109        let y = G::inverse(&self.get(i, j));
110        let z = G::operate(&y, &x);
111        self.update(i, j, z);
112    }
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118    use crate::{
119        algebra::{AdditiveOperation, MaxOperation},
120        tools::Xorshift,
121    };
122
123    const A: u64 = 1_000_000_000;
124    const B: i64 = 1_000_000_000;
125
126    #[test]
127    fn test_binary_indexed_tree_2d() {
128        let mut rng = Xorshift::default();
129        for _ in 0..16 {
130            let h = rng.rand(80) as usize + 1;
131            let w = rng.rand(80) as usize + 1;
132            let q = rng.rand(4_000) as usize + 1_000;
133            let mut bit = BinaryIndexedTree2D::<AdditiveOperation<_>>::new(h, w);
134            let mut arr = vec![vec![0; w]; h];
135            for (i, j, v) in rng.random_iter((..h, ..w, ..A)).take(q) {
136                bit.update(i, j, v);
137                arr[i][j] += v;
138            }
139            for arr in arr.iter_mut() {
140                for j in 0..w - 1 {
141                    arr[j + 1] += arr[j];
142                }
143            }
144            for i in 0..h - 1 {
145                let [a, b] = arr.get_disjoint_mut([i + 1, i]).unwrap();
146                for (a, b) in a.iter_mut().zip(b) {
147                    *a += *b;
148                }
149            }
150            for (i, arr) in arr.iter().enumerate() {
151                for (j, a) in arr.iter().cloned().enumerate() {
152                    assert_eq!(bit.accumulate(i, j), a);
153                }
154            }
155
156            let mut bit = BinaryIndexedTree2D::<MaxOperation<_>>::new(h, w);
157            let mut arr = vec![vec![0; w]; h];
158            for (i, j, v) in rng.random_iter((..h, ..w, ..A)).take(q) {
159                bit.update(i, j, v);
160                arr[i][j] = std::cmp::max(arr[i][j], v);
161            }
162            for arr in arr.iter_mut() {
163                for j in 0..w - 1 {
164                    arr[j + 1] = std::cmp::max(arr[j + 1], arr[j]);
165                }
166            }
167            for i in 0..h - 1 {
168                let [a, b] = arr.get_disjoint_mut([i + 1, i]).unwrap();
169                for (a, b) in a.iter_mut().zip(b) {
170                    *a = std::cmp::max(*a, *b);
171                }
172            }
173            for (i, arr) in arr.iter().enumerate() {
174                for (j, a) in arr.iter().cloned().enumerate() {
175                    assert_eq!(bit.accumulate(i, j), a);
176                }
177            }
178        }
179    }
180
181    #[test]
182    fn test_group_binary_indexed_tree2d() {
183        let mut rng = Xorshift::default();
184        for _ in 0..32 {
185            let h = rng.rand(32) as usize;
186            let w = rng.rand(32) as usize;
187            let mut bit = BinaryIndexedTree2D::<AdditiveOperation<i64>>::new(h, w);
188            let mut values = vec![vec![0; w]; h];
189            for _ in 0..500 {
190                if h != 0 && w != 0 {
191                    let i = rng.rand(h as u64) as usize;
192                    let j = rng.rand(w as u64) as usize;
193                    let value = rng.rand(2 * B as u64) as i64 - B;
194                    if rng.rand(2) == 0 {
195                        bit.update(i, j, value);
196                        values[i][j] += value;
197                    } else {
198                        bit.set(i, j, value);
199                        values[i][j] = value;
200                    }
201                }
202                let i1 = rng.rand(h as u64 + 1) as usize;
203                let i2 = i1 + rng.rand((h - i1) as u64 + 1) as usize;
204                let j1 = rng.rand(w as u64 + 1) as usize;
205                let j2 = j1 + rng.rand((w - j1) as u64 + 1) as usize;
206                assert_eq!(
207                    bit.fold(i1, j1, i2, j2),
208                    values[i1..i2].iter().flat_map(|row| &row[j1..j2]).sum()
209                );
210            }
211        }
212    }
213}