Skip to main content

competitive/data_structure/
binary_trie.rs

1use 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}