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(®istered);
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) = ⦥
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}