Skip to main content

competitive/data_structure/
disjoint_sparse_table.rs

1use super::SemiGroup;
2use std::{
3    fmt::{self, Debug, Formatter},
4    ops::Index,
5};
6
7pub struct DisjointSparseTable<S>
8where
9    S: SemiGroup,
10{
11    table: Vec<S::T>,
12    offsets: Vec<usize>,
13    len: usize,
14}
15
16impl<S> Clone for DisjointSparseTable<S>
17where
18    S: SemiGroup,
19{
20    fn clone(&self) -> Self {
21        Self {
22            table: self.table.clone(),
23            offsets: self.offsets.clone(),
24            len: self.len,
25        }
26    }
27}
28
29impl<S> Debug for DisjointSparseTable<S>
30where
31    S: SemiGroup<T: Debug>,
32{
33    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
34        f.debug_struct("DisjointSparseTable")
35            .field("table", &self.table)
36            .field("offsets", &self.offsets)
37            .field("len", &self.len)
38            .finish()
39    }
40}
41
42impl<S> DisjointSparseTable<S>
43where
44    S: SemiGroup,
45{
46    pub fn new(v: Vec<S::T>) -> Self {
47        let n = v.len();
48        let mut levels = 1;
49        let mut k = 2;
50        while k < n {
51            levels += 1;
52            k *= 2;
53        }
54
55        let mut table = Vec::with_capacity(n * levels);
56        table.extend(v);
57        let mut offsets = vec![0];
58        let mut k = 2;
59        while k < n {
60            let offset = table.len();
61            offsets.push(offset);
62            table.extend_from_within(0..n);
63            for i in (0..n).step_by(k * 2) {
64                for j in (i..n.min(i + k) - 1).rev() {
65                    let j = offset + j;
66                    let x = S::operate(&table[j], &table[j + 1]);
67                    table[j] = x;
68                }
69                for j in i + k + 1..n.min(i + k * 2) {
70                    let j = offset + j;
71                    let x = S::operate(&table[j - 1], &table[j]);
72                    table[j] = x;
73                }
74            }
75            k *= 2;
76        }
77        Self {
78            table,
79            offsets,
80            len: n,
81        }
82    }
83    #[inline]
84    pub fn height(&self) -> usize {
85        self.len
86    }
87    #[inline]
88    fn most_significant_bit_place(x: usize) -> Option<usize> {
89        const C: u32 = usize::MAX.count_ones();
90        ((C - x.leading_zeros()) as usize).checked_sub(1)
91    }
92    #[inline]
93    pub fn fold_close(&self, l: usize, r: usize) -> S::T {
94        debug_assert!(l < self.height());
95        debug_assert!(r < self.height());
96        debug_assert!(l <= r);
97        if let Some(x) = Self::most_significant_bit_place(l ^ r) {
98            let offset = self.offsets[x];
99            S::operate(&self.table[offset + l], &self.table[offset + r])
100        } else {
101            self.table[l].clone()
102        }
103    }
104    #[inline]
105    pub fn fold(&self, l: usize, r: usize) -> S::T {
106        debug_assert!(l < r);
107        self.fold_close(l, r - 1)
108    }
109}
110
111impl<S> Index<usize> for DisjointSparseTable<S>
112where
113    S: SemiGroup,
114{
115    type Output = S::T;
116    #[inline]
117    fn index(&self, index: usize) -> &Self::Output {
118        &self.table[index]
119    }
120}
121
122#[cfg(test)]
123mod tests {
124    use super::*;
125    use crate::{
126        algebra::{AdditiveOperation, ConcatenateOperation, MinOperation},
127        tools::Xorshift,
128    };
129    use std::fmt::Debug;
130
131    fn assert_all_ranges<S>(data: Vec<S::T>)
132    where
133        S: SemiGroup,
134        S::T: Debug + PartialEq,
135    {
136        let table = DisjointSparseTable::<S>::new(data.clone());
137        assert_eq!(table.height(), data.len());
138        for (i, x) in data.iter().enumerate() {
139            assert_eq!(&table[i], x);
140        }
141        for l in 0..data.len() {
142            let mut expected = data[l].clone();
143            assert_eq!(table.fold(l, l + 1), expected);
144            assert_eq!(table.fold_close(l, l), expected);
145            for r in l + 2..=data.len() {
146                expected = S::operate(&expected, &data[r - 1]);
147                assert_eq!(table.fold(l, r), expected);
148                assert_eq!(table.fold_close(l, r - 1), expected);
149            }
150        }
151    }
152
153    #[test]
154    fn test_disjoint_sparse_table_randomized_exhaustive() {
155        let mut rng = Xorshift::default();
156        let mut sizes = vec![0];
157        sizes.extend((0..96).map(|_| rng.random(0usize..=300)));
158        for n in sizes {
159            let data: Vec<i64> = (0..n).map(|_| rng.random(-1000..=1000)).collect();
160            assert_all_ranges::<AdditiveOperation<i64>>(data.clone());
161            assert_all_ranges::<MinOperation<i64>>(data);
162
163            let n = n.min(80);
164            let data: Vec<Vec<i32>> = (0..n).map(|_| vec![rng.random(-1000..=1000)]).collect();
165            assert_all_ranges::<ConcatenateOperation<i32>>(data);
166        }
167    }
168}