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