competitive/data_structure/
binary_trie.rs1use super::{AbelianMonoid, LazyMapMonoid};
2use std::{
3 mem::replace,
4 ops::{Bound, RangeBounds},
5};
6
7struct Node<M>
8where
9 M: LazyMapMonoid,
10{
11 child: [usize; 2],
12 parent: usize,
13 agg: M::Agg,
14 lazy: M::Act,
15}
16
17impl<M> Node<M>
18where
19 M: LazyMapMonoid,
20{
21 fn new(parent: usize) -> Self {
22 Self {
23 child: [usize::MAX; 2],
24 parent,
25 agg: M::agg_unit(),
26 lazy: M::act_unit(),
27 }
28 }
29}
30
31pub struct BinaryTrie<M>
32where
33 M: LazyMapMonoid,
34{
35 bit_len: usize,
36 max_key: u64,
37 len: usize,
38 xor_mask: u64,
39 nodes: Vec<Node<M>>,
40}
41
42impl<M> BinaryTrie<M>
43where
44 M: LazyMapMonoid,
45{
46 pub fn new(bit_len: usize) -> Self {
47 Self::with_capacity(bit_len, 0)
48 }
49
50 pub fn with_capacity(bit_len: usize, capacity: usize) -> Self {
51 assert!(bit_len <= 64);
52 let max_key = if bit_len == 64 {
53 u64::MAX
54 } else {
55 (1u64 << bit_len) - 1
56 };
57 let mut nodes = Vec::with_capacity(
58 capacity
59 .saturating_mul(bit_len.saturating_add(1))
60 .saturating_add(1),
61 );
62 nodes.push(Node::new(usize::MAX));
63 Self {
64 bit_len,
65 max_key,
66 len: 0,
67 xor_mask: 0,
68 nodes,
69 }
70 }
71
72 pub fn len(&self) -> usize {
73 self.len
74 }
75
76 pub fn is_empty(&self) -> bool {
77 self.len() == 0
78 }
79
80 pub fn clear(&mut self) {
81 self.len = 0;
82 self.xor_mask = 0;
83 self.nodes.clear();
84 self.nodes.push(Node::new(usize::MAX));
85 }
86
87 pub fn set(&mut self, key: u64, value: M::Agg) {
88 self.modify_or_insert(key, |x| *x = value);
89 }
90
91 pub fn modify_or_insert(&mut self, key: u64, f: impl FnOnce(&mut M::Agg)) {
92 assert!(key <= self.max_key);
93 if self.bit_len == 0 {
94 if self.is_empty() {
95 self.len = 1;
96 }
97 f(&mut self.nodes[0].agg);
98 return;
99 }
100
101 let key = key ^ self.xor_mask;
102 let mut inserted = false;
103 let mut node = 0;
104 for d in (0..self.bit_len).rev() {
105 self.push_at(node, d + 1);
106 let bit = ((key >> d) & 1) as usize;
107 if self.nodes[node].child[bit] == usize::MAX {
108 inserted = true;
109 let next = self.nodes.len();
110 self.nodes[node].child[bit] = next;
111 self.nodes.push(Node::new(node));
112 }
113 node = self.nodes[node].child[bit];
114 }
115
116 if inserted {
117 self.len += 1;
118 }
119 self.nodes[node].lazy = M::act_unit();
120 f(&mut self.nodes[node].agg);
121 self.recalc_up(node);
122 }
123
124 pub fn get(&mut self, key: u64) -> Option<M::Agg> {
125 assert!(key <= self.max_key);
126 if self.is_empty() {
127 return None;
128 }
129 if self.bit_len == 0 {
130 return Some(self.nodes[0].agg.clone());
131 }
132
133 let key = key ^ self.xor_mask;
134 let mut node = 0;
135 for d in (0..self.bit_len).rev() {
136 let bit = ((key >> d) & 1) as usize;
137 let next = self.nodes[node].child[bit];
138 if next == usize::MAX {
139 return None;
140 }
141 self.push_at(node, d + 1);
142 node = next;
143 }
144 Some(self.nodes[node].agg.clone())
145 }
146
147 pub fn update<R>(&mut self, range: R, act: M::Act)
148 where
149 R: RangeBounds<u64>,
150 {
151 let Some(range) = self.range_to_bounds(range) else {
152 return;
153 };
154 if self.is_empty() {
155 return;
156 }
157
158 let (ql, qr) = range;
159 if ql == 0 && qr == self.max_key {
160 self.apply_at(0, self.bit_len, &act);
161 return;
162 }
163
164 let mut l = ql;
165 loop {
166 let depth = (l.trailing_zeros() as usize)
167 .min(self.bit_len)
168 .min(63 - (qr - l + 1).leading_zeros() as usize);
169 let r = l | ((1u64 << depth) - 1);
170
171 let mut node = 0;
172 for d in (depth..self.bit_len).rev() {
173 self.push_at(node, d + 1);
174 node = self.nodes[node].child[(((l ^ self.xor_mask) >> d) & 1) as usize];
175 if node == usize::MAX {
176 break;
177 }
178 }
179 if node != usize::MAX {
180 self.apply_at(node, depth, &act);
181 self.recalc_up(node);
182 }
183 if r == qr {
184 break;
185 }
186 l = r + 1;
187 }
188 }
189
190 pub fn fold<R>(&mut self, range: R) -> M::Agg
191 where
192 R: RangeBounds<u64>,
193 {
194 let Some(range) = self.range_to_bounds(range) else {
195 return M::agg_unit();
196 };
197
198 let (ql, qr) = range;
199 if ql == 0 && qr == self.max_key {
200 return self.nodes[0].agg.clone();
201 }
202
203 let mut res = M::agg_unit();
204 let mut l = ql;
205 loop {
206 let depth = (l.trailing_zeros() as usize)
207 .min(self.bit_len)
208 .min(63 - (qr - l + 1).leading_zeros() as usize);
209 let r = l | ((1u64 << depth) - 1);
210
211 let mut node = 0;
212 for d in (depth..self.bit_len).rev() {
213 self.push_at(node, d + 1);
214 node = self.nodes[node].child[(((l ^ self.xor_mask) >> d) & 1) as usize];
215 if node == usize::MAX {
216 break;
217 }
218 }
219 if node != usize::MAX {
220 res = M::agg_operate(&res, &self.nodes[node].agg);
221 }
222 if r == qr {
223 break;
224 }
225 l = r + 1;
226 }
227 res
228 }
229
230 fn apply_at(&mut self, node: usize, depth: usize, act: &M::Act) {
231 if M::is_act_unit(act) {
232 return;
233 }
234 if let Some(agg) = M::act_agg(&self.nodes[node].agg, act) {
235 self.nodes[node].agg = agg;
236 if depth > 0 {
237 M::act_operate_assign(&mut self.nodes[node].lazy, act);
238 }
239 } else if depth == 0 {
240 panic!("act failed on leaf");
241 } else {
242 self.push_at(node, depth);
243 for child in self.nodes[node].child {
244 if child != usize::MAX {
245 self.apply_at(child, depth - 1, act);
246 }
247 }
248 self.recalc_at(node);
249 }
250 }
251
252 fn push_at(&mut self, node: usize, depth: usize) {
253 let act = replace(&mut self.nodes[node].lazy, M::act_unit());
254 if M::is_act_unit(&act) {
255 return;
256 }
257 let child = self.nodes[node].child;
258 for child in child {
259 if child != usize::MAX {
260 self.apply_at(child, depth - 1, &act);
261 }
262 }
263 }
264
265 fn recalc_at(&mut self, node: usize) {
266 let mut agg = M::agg_unit();
267 for child in self.nodes[node].child {
268 if child != usize::MAX {
269 agg = M::agg_operate(&agg, &self.nodes[child].agg);
270 }
271 }
272 self.nodes[node].agg = agg;
273 }
274
275 fn recalc_up(&mut self, mut node: usize) {
276 while self.nodes[node].parent != usize::MAX {
277 node = self.nodes[node].parent;
278 self.recalc_at(node);
279 }
280 }
281
282 fn range_to_bounds<R>(&self, range: R) -> Option<(u64, u64)>
283 where
284 R: RangeBounds<u64>,
285 {
286 let start = match range.start_bound() {
287 Bound::Included(&x) => {
288 assert!(x <= self.max_key || (self.bit_len < 64 && x == self.max_key + 1));
289 if x <= self.max_key { Some(x) } else { None }
290 }
291 Bound::Excluded(&x) => {
292 assert!(x <= self.max_key);
293 (x < self.max_key).then_some(x + 1)
294 }
295 Bound::Unbounded => Some(0),
296 };
297 let end = match range.end_bound() {
298 Bound::Included(&x) => {
299 assert!(x <= self.max_key);
300 Some(x)
301 }
302 Bound::Excluded(&x) => {
303 if x == 0 {
304 None
305 } else {
306 assert!(self.bit_len == 64 || x <= self.max_key + 1);
307 Some((x - 1).min(self.max_key))
308 }
309 }
310 Bound::Unbounded => Some(self.max_key),
311 };
312 if let (Some(start), Some(end)) = (start, end) {
313 (start <= end).then_some((start, end))
314 } else {
315 None
316 }
317 }
318}
319
320impl<M> BinaryTrie<M>
321where
322 M: LazyMapMonoid,
323 M::AggMonoid: AbelianMonoid,
324{
325 pub fn xor_all(&mut self, mask: u64) {
326 assert!(mask <= self.max_key);
327 self.xor_mask ^= mask;
328 }
329}
330
331#[cfg(test)]
332mod tests {
333 use super::*;
334 use crate::{
335 algebra::{
336 AdditiveOperation, Associative, FlattenAct, LazyMapMonoid, Magma, RangeSumRangeAdd,
337 Unital,
338 },
339 tools::Xorshift,
340 };
341 use std::{
342 collections::BTreeMap,
343 marker::PhantomData,
344 ops::{Bound, RangeBounds},
345 };
346
347 #[test]
348 fn binary_trie_range_sum_randomized() {
349 const A: i64 = 100;
350 const Q: usize = 4_000;
351 let mut rng = Xorshift::default();
352
353 for bit_len in [0, 1, 6, 64] {
354 let mut trie = BinaryTrie::<RangeSumRangeAdd<i64>>::new(bit_len);
355 let mut map = BTreeMap::new();
356 let universe = (bit_len < 64).then(|| 1u64 << bit_len);
357 let max_key = universe.map_or(u64::MAX, |n| n - 1);
358
359 for _ in 0..Q {
360 let key = match universe {
361 Some(n) => rng.random(0..n),
362 None => rng.random(..),
363 };
364 let mut x = match universe {
365 Some(_) => rng.random(0..=max_key),
366 None => rng.random(..),
367 };
368 let mut y = match universe {
369 Some(_) => rng.random(0..=max_key),
370 None => rng.random(..),
371 };
372 if x > y {
373 std::mem::swap(&mut x, &mut y);
374 }
375 let range = (
376 match rng.random(0..3) {
377 0 => Bound::Excluded(x),
378 1 => Bound::Included(x),
379 _ => Bound::Unbounded,
380 },
381 match rng.random(0..3) {
382 0 => Bound::Excluded(y),
383 1 => Bound::Included(y),
384 _ => Bound::Unbounded,
385 },
386 );
387 match rng.random(0..10) {
388 0 => {
389 let value = (rng.random(-A..=A), rng.random(0i64..=5));
390 trie.set(key, value);
391 map.insert(key, value);
392 }
393 1 => {
394 let dx = rng.random(-A..=A);
395 let dy = rng.random(0i64..=3);
396 trie.modify_or_insert(key, |value| {
397 value.0 += dx;
398 value.1 += dy;
399 });
400 let value = map.entry(key).or_insert((0, 0));
401 value.0 += dx;
402 value.1 += dy;
403 }
404 2 => {
405 assert_eq!(trie.get(key), map.get(&key).copied());
406 }
407 3 => {
408 let add = rng.random(-A..=A);
409 trie.update(range, add);
410 for (key, value) in map.iter_mut() {
411 if range.contains(key) {
412 value.0 += add * value.1;
413 }
414 }
415 }
416 4 => {
417 assert_eq!(
418 trie.fold(range),
419 map.iter()
420 .filter(|(key, _)| range.contains(key))
421 .fold((0, 0), |(sx, sy), (_, &(x, y))| (sx + x, sy + y))
422 );
423 }
424 5 => {
425 let add = rng.random(-A..=A);
426 trie.update(.., add);
427 for value in map.values_mut() {
428 value.0 += add * value.1;
429 }
430 }
431 6 => {
432 assert_eq!(
433 trie.fold(..),
434 map.values()
435 .fold((0, 0), |(sx, sy), &(x, y)| (sx + x, sy + y))
436 );
437 }
438 7 => {
439 let add = rng.random(-A..=A);
440 trie.update(..=max_key, add);
441 for value in map.values_mut() {
442 value.0 += add * value.1;
443 }
444 }
445 8 => {
446 assert_eq!(
447 trie.fold(max_key..=max_key),
448 map.get(&max_key).copied().unwrap_or((0, 0))
449 );
450 }
451 _ => {
452 let mask = match universe {
453 Some(n) => rng.random(0..n),
454 None => rng.random(..),
455 };
456 trie.xor_all(mask);
457 map = map
458 .into_iter()
459 .map(|(key, value)| (key ^ mask, value))
460 .collect();
461 }
462 }
463 assert_eq!(trie.len(), map.len());
464 assert_eq!(trie.is_empty(), map.is_empty());
465 }
466
467 trie.clear();
468 map.clear();
469 assert_eq!(trie.fold(..), (0, 0));
470 assert!(trie.is_empty());
471 }
472 }
473
474 struct Concat;
475
476 impl Magma for Concat {
477 type T = Vec<i32>;
478
479 fn operate(x: &Self::T, y: &Self::T) -> Self::T {
480 let mut res = x.clone();
481 res.extend(y);
482 res
483 }
484 }
485
486 impl Associative for Concat {}
487
488 impl Unital for Concat {
489 fn unit() -> Self::T {
490 Vec::new()
491 }
492 }
493
494 struct DescendAdd {
495 _marker: PhantomData<fn()>,
496 }
497
498 impl LazyMapMonoid for DescendAdd {
499 type Key = i32;
500 type Agg = Vec<i32>;
501 type Act = i32;
502 type AggMonoid = Concat;
503 type ActMonoid = AdditiveOperation<i32>;
504 type KeyAct = FlattenAct<AdditiveOperation<i32>>;
505
506 fn single_agg(key: &Self::Key) -> Self::Agg {
507 vec![*key]
508 }
509
510 fn act_agg(x: &Self::Agg, a: &Self::Act) -> Option<Self::Agg> {
511 Self::is_act_unit(a)
512 .then_some(x.clone())
513 .or_else(|| (x.len() <= 1).then(|| x.iter().map(|x| x + a).collect()))
514 }
515 }
516
517 #[test]
518 fn binary_trie_non_commutative_descending_lazy_randomized() {
519 const B: usize = 5;
520 const Q: usize = 2_000;
521 let mut rng = Xorshift::default();
522 let mut trie = BinaryTrie::<DescendAdd>::new(B);
523 let mut map = BTreeMap::<u64, Vec<i32>>::new();
524 let universe = 1u64 << B;
525
526 for _ in 0..Q {
527 let key = rng.random(0..universe);
528 let l = rng.random(0..=universe);
529 let r = rng.random(l..=universe);
530 match rng.random(0..5) {
531 0 => {
532 let value = vec![rng.random(-100..=100)];
533 trie.set(key, value.clone());
534 map.insert(key, value);
535 }
536 1 => {
537 let value = rng.random(-100..=100);
538 trie.modify_or_insert(key, |bucket| {
539 if bucket.is_empty() {
540 bucket.push(value);
541 } else {
542 bucket[0] += value;
543 }
544 });
545 map.entry(key)
546 .and_modify(|bucket| bucket[0] += value)
547 .or_insert_with(|| vec![value]);
548 }
549 2 => {
550 let add = rng.random(-100..=100);
551 trie.update(l..r, add);
552 for (_, bucket) in map.range_mut(l..r) {
553 for value in bucket {
554 *value += add;
555 }
556 }
557 }
558 3 => {
559 assert_eq!(trie.get(key), map.get(&key).cloned());
560 }
561 _ => {
562 let expected = map
563 .range(l..r)
564 .flat_map(|(_, value)| value.iter().copied())
565 .collect::<Vec<_>>();
566 assert_eq!(trie.fold(l..r), expected);
567 }
568 }
569 assert_eq!(trie.len(), map.len());
570 }
571 }
572}