Skip to main content

competitive/data_structure/
static_range_product.rs

1use super::{DisjointSparseTable, SemiGroup};
2use std::fmt::{self, Debug, Formatter};
3
4const DIRECT_SIZE: usize = 4;
5const BLOCK_SCALE: usize = 4;
6
7struct BlockProducts<T> {
8    prefix: Vec<T>,
9    suffix: Vec<T>,
10    products: Vec<T>,
11}
12
13#[derive(Clone)]
14pub struct StaticRangeProduct<S>
15where
16    S: SemiGroup,
17{
18    data: Vec<S::T>,
19    block_shift: usize,
20    prefix: Vec<S::T>,
21    suffix: Vec<S::T>,
22    between: Option<FixedRangeProduct<S>>,
23}
24
25impl<S> Debug for StaticRangeProduct<S>
26where
27    S: SemiGroup<T: Debug>,
28{
29    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
30        f.debug_struct("StaticRangeProduct")
31            .field("data", &self.data)
32            .field("block_size", &(1usize << self.block_shift))
33            .field("prefix", &self.prefix)
34            .field("suffix", &self.suffix)
35            .field("between", &self.between)
36            .finish()
37    }
38}
39
40impl<S> StaticRangeProduct<S>
41where
42    S: SemiGroup,
43{
44    pub fn new(data: Vec<S::T>) -> Self {
45        let n = data.len();
46        if n == 0 {
47            return Self {
48                data,
49                block_shift: 0,
50                prefix: Vec::new(),
51                suffix: Vec::new(),
52                between: None,
53            };
54        }
55        let block_shift = scaled_block_shift(inverse_ackermann(n));
56        let block_size = 1usize << block_shift;
57        let blocks = block_products::<S>(&data, block_size);
58        let between = if blocks.products.len() > 2 {
59            let level = inverse_ackermann(blocks.products.len()).max(1);
60            Some(FixedRangeProduct::new(blocks.products, level))
61        } else {
62            None
63        };
64        Self {
65            data,
66            block_shift,
67            prefix: blocks.prefix,
68            suffix: blocks.suffix,
69            between,
70        }
71    }
72
73    #[inline]
74    pub fn len(&self) -> usize {
75        self.data.len()
76    }
77
78    #[inline]
79    pub fn is_empty(&self) -> bool {
80        self.data.is_empty()
81    }
82
83    #[inline]
84    pub fn get(&self, index: usize) -> &S::T {
85        &self.data[index]
86    }
87
88    #[inline]
89    pub fn fold(&self, l: usize, r: usize) -> S::T {
90        assert!(l < r);
91        assert!(r <= self.data.len());
92        let bl = l >> self.block_shift;
93        let br = (r - 1) >> self.block_shift;
94        if bl == br {
95            return fold_slice::<S>(&self.data, l, r);
96        }
97        let mut res = self.suffix[l].clone();
98        if bl + 1 < br {
99            let mid = self
100                .between
101                .as_ref()
102                .expect("middle block product is not built")
103                .fold(bl + 1, br);
104            res = S::operate(&res, &mid);
105        }
106        S::operate(&res, &self.prefix[r - 1])
107    }
108}
109
110#[derive(Clone)]
111enum FixedRangeProduct<S>
112where
113    S: SemiGroup,
114{
115    Direct {
116        data: Vec<S::T>,
117    },
118    Disjoint {
119        table: DisjointSparseTable<S>,
120    },
121    Recursive {
122        data: Vec<S::T>,
123        block_shift: usize,
124        prefix: Vec<S::T>,
125        suffix: Vec<S::T>,
126        between: Box<FixedRangeProduct<S>>,
127    },
128}
129
130impl<S> Debug for FixedRangeProduct<S>
131where
132    S: SemiGroup<T: Debug>,
133{
134    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
135        match self {
136            Self::Direct { data } => f.debug_struct("Direct").field("data", data).finish(),
137            Self::Disjoint { table } => f.debug_struct("Disjoint").field("table", table).finish(),
138            Self::Recursive {
139                data,
140                block_shift,
141                prefix,
142                suffix,
143                between,
144            } => f
145                .debug_struct("Recursive")
146                .field("data", data)
147                .field("block_size", &(1usize << block_shift))
148                .field("prefix", prefix)
149                .field("suffix", suffix)
150                .field("between", between)
151                .finish(),
152        }
153    }
154}
155
156impl<S> FixedRangeProduct<S>
157where
158    S: SemiGroup,
159{
160    fn new(data: Vec<S::T>, level: usize) -> Self {
161        let n = data.len();
162        if n <= DIRECT_SIZE || level == 0 {
163            return Self::Direct { data };
164        }
165        if level == 1 {
166            return Self::Disjoint {
167                table: DisjointSparseTable::new(data),
168            };
169        }
170        let block_shift = scaled_block_shift(alpha_k(level - 1, n));
171        let block_size = 1usize << block_shift;
172        if block_size <= 1 || block_size >= n {
173            return Self::Direct { data };
174        }
175        let blocks = block_products::<S>(&data, block_size);
176        let between = Box::new(Self::new(blocks.products, level - 1));
177        Self::Recursive {
178            data,
179            block_shift,
180            prefix: blocks.prefix,
181            suffix: blocks.suffix,
182            between,
183        }
184    }
185
186    #[inline]
187    fn fold(&self, l: usize, r: usize) -> S::T {
188        match self {
189            Self::Direct { data } => fold_slice::<S>(data, l, r),
190            Self::Disjoint { table } => table.fold(l, r),
191            Self::Recursive {
192                data,
193                block_shift,
194                prefix,
195                suffix,
196                between,
197            } => {
198                let block_shift = *block_shift;
199                let bl = l >> block_shift;
200                let br = (r - 1) >> block_shift;
201                if bl == br {
202                    return fold_slice::<S>(data, l, r);
203                }
204                let mut res = suffix[l].clone();
205                if bl + 1 < br {
206                    let mid = between.fold(bl + 1, br);
207                    res = S::operate(&res, &mid);
208                }
209                S::operate(&res, &prefix[r - 1])
210            }
211        }
212    }
213}
214
215#[inline]
216fn fold_slice<S>(data: &[S::T], l: usize, r: usize) -> S::T
217where
218    S: SemiGroup,
219{
220    let mut res = data[l].clone();
221    for x in &data[l + 1..r] {
222        res = S::operate(&res, x);
223    }
224    res
225}
226
227fn block_products<S>(data: &[S::T], block_size: usize) -> BlockProducts<S::T>
228where
229    S: SemiGroup,
230{
231    let n = data.len();
232    let mut prefix = data.to_vec();
233    let mut suffix = data.to_vec();
234    let mut products = Vec::with_capacity(n.div_ceil(block_size));
235    for start in (0..n).step_by(block_size) {
236        let end = n.min(start + block_size);
237        for i in start + 1..end {
238            prefix[i] = S::operate(&prefix[i - 1], &data[i]);
239        }
240        for i in (start..end - 1).rev() {
241            suffix[i] = S::operate(&data[i], &suffix[i + 1]);
242        }
243        products.push(prefix[end - 1].clone());
244    }
245    BlockProducts {
246        prefix,
247        suffix,
248        products,
249    }
250}
251
252fn alpha_k(k: usize, n: usize) -> usize {
253    if k == 0 {
254        return n.div_ceil(2);
255    }
256    if n <= 1 {
257        return 0;
258    }
259    let mut x = n;
260    let mut c = 0;
261    while x > 1 {
262        x = alpha_k(k - 1, x);
263        c += 1;
264    }
265    c
266}
267
268fn inverse_ackermann(n: usize) -> usize {
269    let mut k = 0;
270    while alpha_k(k, n) > 3 {
271        k += 1;
272    }
273    k
274}
275
276fn scaled_block_shift(alpha: usize) -> usize {
277    (alpha.max(2) * BLOCK_SCALE)
278        .next_power_of_two()
279        .trailing_zeros() as usize
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285    use crate::{
286        algebra::{AdditiveOperation, ConcatenateOperation, MinOperation},
287        tools::Xorshift,
288    };
289
290    fn assert_all_ranges<S>(data: Vec<S::T>)
291    where
292        S: SemiGroup,
293        S::T: Debug + PartialEq,
294    {
295        let table = StaticRangeProduct::<S>::new(data.clone());
296        assert_eq!(table.len(), data.len());
297        assert_eq!(table.is_empty(), data.is_empty());
298        for (i, x) in data.iter().enumerate() {
299            assert_eq!(table.get(i), x);
300        }
301        for l in 0..data.len() {
302            let mut expected = data[l].clone();
303            assert_eq!(table.fold(l, l + 1), expected);
304            for r in l + 2..=data.len() {
305                expected = S::operate(&expected, &data[r - 1]);
306                assert_eq!(table.fold(l, r), expected);
307            }
308        }
309    }
310
311    #[test]
312    fn test_static_range_product_randomized_exhaustive() {
313        let mut rng = Xorshift::default();
314        let mut sizes = vec![0];
315        sizes.extend((0..96).map(|_| rng.random(0usize..=300)));
316        for n in sizes {
317            let data: Vec<i64> = (0..n).map(|_| rng.random(-1000..=1000)).collect();
318            assert_all_ranges::<AdditiveOperation<i64>>(data.clone());
319            assert_all_ranges::<MinOperation<i64>>(data);
320
321            let n = n.min(80);
322            let data: Vec<Vec<i32>> = (0..n).map(|_| vec![rng.random(-1000..=1000)]).collect();
323            assert_all_ranges::<ConcatenateOperation<i32>>(data);
324        }
325    }
326}