competitive/data_structure/
binary_indexed_tree_2d.rs1use 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 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 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 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 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 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}