Skip to main content

competitive/data_structure/
dary_segment_tree.rs

1#[cfg(target_arch = "x86_64")]
2use super::simd;
3use super::{RangeBoundsExt, SimdBackend, simd_backend};
4use std::ops::RangeBounds;
5
6#[repr(C, align(64))]
7#[derive(Clone, Debug)]
8struct Block<T, const B: usize>([T; B]);
9
10macro_rules! define_dary_segment_tree {
11    (
12        $name:ident,
13        $doc:literal,
14        $value:ty,
15        $branch:expr,
16        $unit:expr,
17        $operation:ident,
18        $backend:expr,
19        $sum:literal,
20        $reduce_avx2:ident,
21        $reduce_range_avx2:ident,
22        $reduce_avx512:ident,
23        $reduce_range_avx512:ident
24    ) => {
25        #[doc = $doc]
26        #[derive(Clone, Debug)]
27        pub struct $name {
28            levels: Vec<Vec<Block<$value, $branch>>>,
29            len: usize,
30            #[cfg(target_arch = "x86_64")]
31            backend: SimdBackend,
32        }
33
34        impl $name {
35            pub fn new(len: usize) -> Self {
36                let mut levels = Vec::new();
37                let mut current = len.max(1);
38                loop {
39                    levels.push(vec![Block([$unit; $branch]); current.div_ceil($branch)]);
40                    if current == 1 {
41                        break;
42                    }
43                    current = current.div_ceil($branch);
44                }
45                Self {
46                    levels,
47                    len,
48                    #[cfg(target_arch = "x86_64")]
49                    backend: $backend,
50                }
51            }
52
53            pub fn from_vec(values: Vec<$value>) -> Self {
54                Self::build(values, $backend)
55            }
56
57            #[inline]
58            pub fn len(&self) -> usize {
59                self.len
60            }
61
62            #[inline]
63            pub fn is_empty(&self) -> bool {
64                self.len == 0
65            }
66
67            #[inline]
68            pub fn set(&mut self, index: usize, value: $value) {
69                assert!(index < self.len);
70                self.set_value(index, value);
71            }
72
73            #[inline]
74            pub fn clear(&mut self, index: usize) {
75                self.set(index, $unit);
76            }
77
78            #[inline]
79            pub fn update(&mut self, index: usize, value: $value) {
80                assert!(index < self.len);
81                let current = self.levels[0][index / $branch].0[index % $branch];
82                self.set_value(index, current.$operation(value));
83            }
84
85            #[inline]
86            pub fn get(&self, index: usize) -> $value {
87                assert!(index < self.len);
88                self.levels[0][index / $branch].0[index % $branch]
89            }
90
91            #[inline]
92            pub fn fold<R>(&self, range: R) -> $value
93            where
94                R: RangeBounds<usize>,
95            {
96                let range = range.to_range_bounded(0, self.len).expect("invalid range");
97                #[cfg(target_arch = "x86_64")]
98                return match self.backend {
99                    SimdBackend::Scalar => {
100                        self.fold_by(range.start, range.end, Self::reduce_range_scalar)
101                    }
102                    // SAFETY: construction fixes every block width and selects a supported backend.
103                    SimdBackend::Avx2 => unsafe { self.fold_avx2(range.start, range.end) },
104                    // SAFETY: same as above.
105                    SimdBackend::Avx512 => unsafe { self.fold_avx512(range.start, range.end) },
106                };
107                #[cfg(not(target_arch = "x86_64"))]
108                self.fold_by(range.start, range.end, Self::reduce_range_scalar)
109            }
110
111            #[inline]
112            pub fn fold_all(&self) -> $value {
113                self.levels.last().unwrap()[0].0[0]
114            }
115
116            fn build(values: Vec<$value>, backend: SimdBackend) -> Self {
117                let _ = &backend;
118                let len = values.len();
119                let mut current = if values.is_empty() {
120                    vec![$unit]
121                } else {
122                    values
123                };
124                let mut levels = Vec::new();
125                loop {
126                    let blocks: Vec<_> = current
127                        .chunks($branch)
128                        .map(|chunk| {
129                            let mut values = [$unit; $branch];
130                            values[..chunk.len()].copy_from_slice(chunk);
131                            Block(values)
132                        })
133                        .collect();
134                    if current.len() == 1 {
135                        levels.push(blocks);
136                        break;
137                    }
138                    current = blocks
139                        .iter()
140                        .map(|block| Self::reduce_scalar(&block.0))
141                        .collect();
142                    levels.push(blocks);
143                }
144                Self {
145                    levels,
146                    len,
147                    #[cfg(target_arch = "x86_64")]
148                    backend,
149                }
150            }
151
152            #[inline]
153            fn set_value(&mut self, index: usize, value: $value) {
154                if self.levels[0][index / $branch].0[index % $branch] == value {
155                    return;
156                }
157                if $sum {
158                    let delta =
159                        value.wrapping_sub(self.levels[0][index / $branch].0[index % $branch]);
160                    let mut index = index;
161                    for level in &mut self.levels {
162                        let current = &mut level[index / $branch].0[index % $branch];
163                        *current = current.wrapping_add(delta);
164                        index /= $branch;
165                    }
166                    return;
167                }
168                #[cfg(target_arch = "x86_64")]
169                match self.backend {
170                    SimdBackend::Scalar => self.set_by(index, value, Self::reduce_scalar),
171                    // SAFETY: construction selects a supported backend.
172                    SimdBackend::Avx2 => unsafe { self.set_avx2(index, value) },
173                    // SAFETY: same as above.
174                    SimdBackend::Avx512 => unsafe { self.set_avx512(index, value) },
175                }
176                #[cfg(not(target_arch = "x86_64"))]
177                self.set_by(index, value, Self::reduce_scalar);
178            }
179
180            #[inline(always)]
181            fn reduce_scalar(values: &[$value; $branch]) -> $value {
182                let mut result = values[0];
183                for &value in &values[1..] {
184                    result = result.$operation(value);
185                }
186                result
187            }
188
189            #[inline(always)]
190            fn reduce_range_scalar(values: &[$value; $branch], start: usize, end: usize) -> $value {
191                values[start..end]
192                    .iter()
193                    .copied()
194                    .reduce(<$value>::$operation)
195                    .unwrap_or($unit)
196            }
197
198            #[inline(always)]
199            fn set_by<F>(&mut self, mut index: usize, value: $value, mut reduce: F)
200            where
201                F: FnMut(&[$value; $branch]) -> $value,
202            {
203                self.levels[0][index / $branch].0[index % $branch] = value;
204                for level in 0..self.levels.len() - 1 {
205                    let block = index / $branch;
206                    let aggregate = reduce(&self.levels[level][block].0);
207                    index = block;
208                    let parent = &mut self.levels[level + 1][index / $branch].0[index % $branch];
209                    if *parent == aggregate {
210                        break;
211                    }
212                    *parent = aggregate;
213                }
214            }
215
216            #[inline(always)]
217            fn fold_by<F>(&self, mut left: usize, mut right: usize, mut reduce: F) -> $value
218            where
219                F: FnMut(&[$value; $branch], usize, usize) -> $value,
220            {
221                let mut result: $value = $unit;
222                for level in &self.levels {
223                    if left >= right {
224                        break;
225                    }
226                    let first = left / $branch;
227                    let last = (right - 1) / $branch;
228                    if first == last {
229                        return result.$operation(reduce(
230                            &level[first].0,
231                            left % $branch,
232                            (right - 1) % $branch + 1,
233                        ));
234                    }
235                    if left % $branch != 0 {
236                        result =
237                            result.$operation(reduce(&level[first].0, left % $branch, $branch));
238                        left = (first + 1) * $branch;
239                    }
240                    if right % $branch != 0 {
241                        result = result.$operation(reduce(&level[last].0, 0, right % $branch));
242                        right = last * $branch;
243                    }
244                    left /= $branch;
245                    right /= $branch;
246                }
247                result
248            }
249
250            #[cfg(target_arch = "x86_64")]
251            #[target_feature(enable = "avx2")]
252            unsafe fn set_avx2(&mut self, index: usize, value: $value) {
253                self.set_by(index, value, |values| unsafe { simd::$reduce_avx2(values) });
254            }
255
256            #[cfg(target_arch = "x86_64")]
257            #[target_feature(enable = "avx2")]
258            unsafe fn fold_avx2(&self, left: usize, right: usize) -> $value {
259                self.fold_by(left, right, |values, start, end| unsafe {
260                    simd::$reduce_range_avx2(values, start, end)
261                })
262            }
263
264            #[cfg(target_arch = "x86_64")]
265            #[target_feature(enable = "avx512f")]
266            unsafe fn set_avx512(&mut self, index: usize, value: $value) {
267                self.set_by(index, value, |values| unsafe {
268                    simd::$reduce_avx512(values)
269                });
270            }
271
272            #[cfg(target_arch = "x86_64")]
273            #[target_feature(enable = "avx512f")]
274            unsafe fn fold_avx512(&self, left: usize, right: usize) -> $value {
275                self.fold_by(left, right, |values, start, end| unsafe {
276                    simd::$reduce_range_avx512(values, start, end)
277                })
278            }
279        }
280    };
281}
282
283define_dary_segment_tree!(
284    DarySegmentTreeMinI32,
285    "A cache-line-oriented d-ary point-update segment tree for range minima over `i32`.",
286    i32,
287    16,
288    i32::MAX,
289    min,
290    simd_backend(),
291    false,
292    minimum_i32x16_avx2,
293    minimum_range_i32x16_avx2,
294    minimum_i32x16_avx512,
295    minimum_range_i32x16_avx512
296);
297define_dary_segment_tree!(
298    DarySegmentTreeMaxI32,
299    "A cache-line-oriented d-ary point-update segment tree for range maxima over `i32`.",
300    i32,
301    16,
302    i32::MIN,
303    max,
304    simd_backend(),
305    false,
306    maximum_i32x16_avx2,
307    maximum_range_i32x16_avx2,
308    maximum_i32x16_avx512,
309    maximum_range_i32x16_avx512
310);
311define_dary_segment_tree!(
312    DarySegmentTreeMinI64,
313    "A cache-line-oriented d-ary point-update segment tree for range minima over `i64`.",
314    i64,
315    8,
316    i64::MAX,
317    min,
318    simd_backend(),
319    false,
320    minimum_i64x8_avx2,
321    minimum_range_i64x8_avx2,
322    minimum_i64x8_avx512,
323    minimum_range_i64x8_avx512
324);
325define_dary_segment_tree!(
326    DarySegmentTreeMaxI64,
327    "A cache-line-oriented d-ary point-update segment tree for range maxima over `i64`.",
328    i64,
329    8,
330    i64::MIN,
331    max,
332    simd_backend(),
333    false,
334    maximum_i64x8_avx2,
335    maximum_range_i64x8_avx2,
336    maximum_i64x8_avx512,
337    maximum_range_i64x8_avx512
338);
339define_dary_segment_tree!(
340    DarySegmentTreeAddI32,
341    "A cache-line-oriented d-ary point-update segment tree for wrapping range sums over `i32`.",
342    i32,
343    16,
344    0,
345    wrapping_add,
346    simd_backend(),
347    true,
348    sum_i32x16_avx2,
349    sum_range_i32x16_avx2,
350    sum_i32x16_avx512,
351    sum_range_i32x16_avx512
352);
353define_dary_segment_tree!(
354    DarySegmentTreeAddI64,
355    "A cache-line-oriented d-ary point-update segment tree for wrapping range sums over `i64`.",
356    i64,
357    8,
358    0,
359    wrapping_add,
360    simd_backend(),
361    true,
362    sum_i64x8_avx2,
363    sum_range_i64x8_avx2,
364    sum_i64x8_avx512,
365    sum_range_i64x8_avx512
366);
367
368#[cfg(test)]
369mod tests {
370    use super::*;
371    use crate::tools::Xorshift;
372    #[cfg(target_arch = "x86_64")]
373    use crate::tools::avx512_supported;
374
375    #[cfg(target_arch = "x86_64")]
376    fn backends() -> Vec<SimdBackend> {
377        let mut result = vec![SimdBackend::Scalar];
378        if is_x86_feature_detected!("avx2") {
379            result.push(SimdBackend::Avx2);
380        }
381        if avx512_supported() {
382            result.push(SimdBackend::Avx512);
383        }
384        result
385    }
386
387    #[cfg(not(target_arch = "x86_64"))]
388    fn backends() -> Vec<SimdBackend> {
389        vec![SimdBackend::Scalar]
390    }
391
392    #[test]
393    fn test_dary_segment_tree() {
394        let mut rng = Xorshift::default();
395        macro_rules! check {
396            ($value:ty, $minimum:ty, $maximum:ty, $sum:ty) => {{
397                for len in [
398                    0, 1, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65, 255, 256, 257, 4097,
399                ] {
400                    let mut values: Vec<$value> =
401                        (0..len).map(|_| rng.rand64() as $value).collect();
402                    if let Some(value) = values.get_mut(0) {
403                        *value = <$value>::MIN;
404                    }
405                    if let Some(value) = values.get_mut(1) {
406                        *value = 0;
407                    }
408                    if let Some(value) = values.get_mut(2) {
409                        *value = <$value>::MAX;
410                    }
411                    for (backend, from_values) in backends()
412                        .into_iter()
413                        .flat_map(|backend| [(backend, false), (backend, true)])
414                    {
415                        let mut minimum = if from_values {
416                            <$minimum>::build(values.clone(), backend)
417                        } else {
418                            <$minimum>::new(len)
419                        };
420                        let mut maximum = if from_values {
421                            <$maximum>::build(values.clone(), backend)
422                        } else {
423                            <$maximum>::new(len)
424                        };
425                        let mut sum = if from_values {
426                            <$sum>::build(values.clone(), backend)
427                        } else {
428                            <$sum>::new(len)
429                        };
430                        let mut expected_minimum = if from_values {
431                            values.clone()
432                        } else {
433                            vec![<$value>::MAX; len]
434                        };
435                        let mut expected_maximum = if from_values {
436                            values.clone()
437                        } else {
438                            vec![<$value>::MIN; len]
439                        };
440                        let mut expected_sum = if from_values {
441                            values.clone()
442                        } else {
443                            vec![0; len]
444                        };
445                        assert_eq!(minimum.len(), len);
446                        assert_eq!(maximum.len(), len);
447                        assert_eq!(sum.len(), len);
448                        assert_eq!(minimum.is_empty(), len == 0);
449                        assert_eq!(maximum.is_empty(), len == 0);
450                        assert_eq!(sum.is_empty(), len == 0);
451                        for _ in 0..500 {
452                            if len != 0 {
453                                let index = rng.rand(len as u64) as usize;
454                                let value = rng.rand64() as $value;
455                                match rng.rand(5) {
456                                    0 => {
457                                        minimum.set(index, value);
458                                        maximum.set(index, value);
459                                        sum.set(index, value);
460                                        expected_minimum[index] = value;
461                                        expected_maximum[index] = value;
462                                        expected_sum[index] = value;
463                                    }
464                                    1 => {
465                                        minimum.update(index, value);
466                                        maximum.update(index, value);
467                                        sum.update(index, value);
468                                        expected_minimum[index] =
469                                            expected_minimum[index].min(value);
470                                        expected_maximum[index] =
471                                            expected_maximum[index].max(value);
472                                        expected_sum[index] =
473                                            expected_sum[index].wrapping_add(value);
474                                    }
475                                    2 => {
476                                        minimum.clear(index);
477                                        maximum.clear(index);
478                                        sum.clear(index);
479                                        expected_minimum[index] = <$value>::MAX;
480                                        expected_maximum[index] = <$value>::MIN;
481                                        expected_sum[index] = 0;
482                                    }
483                                    3 => {
484                                        assert_eq!(minimum.get(index), expected_minimum[index]);
485                                        assert_eq!(maximum.get(index), expected_maximum[index]);
486                                        assert_eq!(sum.get(index), expected_sum[index]);
487                                    }
488                                    _ => {
489                                        let left = rng.rand(len as u64 + 1) as usize;
490                                        let right =
491                                            left + rng.rand((len - left) as u64 + 1) as usize;
492                                        assert_eq!(
493                                            minimum.fold(left..right),
494                                            expected_minimum[left..right]
495                                                .iter()
496                                                .copied()
497                                                .min()
498                                                .unwrap_or(<$value>::MAX)
499                                        );
500                                        assert_eq!(
501                                            maximum.fold(left..right),
502                                            expected_maximum[left..right]
503                                                .iter()
504                                                .copied()
505                                                .max()
506                                                .unwrap_or(<$value>::MIN)
507                                        );
508                                        assert_eq!(
509                                            sum.fold(left..right),
510                                            expected_sum[left..right]
511                                                .iter()
512                                                .copied()
513                                                .fold(0, <$value>::wrapping_add)
514                                        );
515                                    }
516                                }
517                            }
518                            assert_eq!(
519                                minimum.fold_all(),
520                                expected_minimum
521                                    .iter()
522                                    .copied()
523                                    .min()
524                                    .unwrap_or(<$value>::MAX)
525                            );
526                            assert_eq!(
527                                maximum.fold_all(),
528                                expected_maximum
529                                    .iter()
530                                    .copied()
531                                    .max()
532                                    .unwrap_or(<$value>::MIN)
533                            );
534                            assert_eq!(
535                                sum.fold_all(),
536                                expected_sum.iter().copied().fold(0, <$value>::wrapping_add)
537                            );
538                        }
539                    }
540                }
541            }};
542        }
543
544        check!(
545            i32,
546            DarySegmentTreeMinI32,
547            DarySegmentTreeMaxI32,
548            DarySegmentTreeAddI32
549        );
550        check!(
551            i64,
552            DarySegmentTreeMinI64,
553            DarySegmentTreeMaxI64,
554            DarySegmentTreeAddI64
555        );
556    }
557}