Skip to main content

competitive/data_structure/
dary_prefix_sum_tree.rs

1#[cfg(target_arch = "x86_64")]
2use super::simd;
3use super::{SimdBackend, simd_backend};
4
5#[repr(C, align(64))]
6#[derive(Clone, Debug)]
7struct PrefixBlock<T, const B: usize>([T; B]);
8
9macro_rules! define_dary_prefix_sum_tree {
10    (
11        $name:ident,
12        $value:ty,
13        $branch:expr,
14        $add_avx2:ident,
15        $first_gt_avx2:ident,
16        $add_avx512:ident,
17        $first_gt_avx512:ident
18    ) => {
19        /// A cache-line-oriented d-ary tree for point updates, prefix sums, and prefix searches.
20        ///
21        /// `BinaryIndexedTree` has a denser layout and can be preferable for small workloads.
22        /// Use this type for repeated updates and prefix searches, especially with SIMD support.
23        #[derive(Clone, Debug)]
24        pub struct $name {
25            levels: Vec<Vec<PrefixBlock<$value, $branch>>>,
26            len: usize,
27            total: $value,
28            partition_valid: bool,
29            #[cfg(target_arch = "x86_64")]
30            backend: SimdBackend,
31        }
32
33        impl $name {
34            pub fn new(len: usize) -> Self {
35                Self::zeroed(len, simd_backend())
36            }
37
38            pub fn from_slice(values: &[$value]) -> Self {
39                Self::build(values, simd_backend())
40            }
41
42            #[inline]
43            pub fn len(&self) -> usize {
44                self.len
45            }
46
47            #[inline]
48            pub fn is_empty(&self) -> bool {
49                self.len == 0
50            }
51
52            /// Adds `value` at `index`. Arithmetic is wrapping.
53            #[inline]
54            pub fn update(&mut self, index: usize, value: $value) {
55                assert!(index < self.len);
56                self.add(index, value);
57                self.partition_valid &= self.total.checked_add(value).is_some();
58                self.total = self.total.wrapping_add(value);
59            }
60
61            /// Replaces the value at `index`. Arithmetic is wrapping.
62            #[inline]
63            pub fn set(&mut self, index: usize, value: $value) {
64                let previous = self.get(index);
65                self.add(index, value.wrapping_sub(previous));
66                self.partition_valid &= self
67                    .total
68                    .checked_sub(previous)
69                    .and_then(|total| total.checked_add(value))
70                    .is_some();
71                self.total = self.total.wrapping_sub(previous).wrapping_add(value);
72            }
73
74            /// Returns the wrapping sum of `0..end`.
75            #[inline]
76            pub fn accumulate0(&self, mut end: usize) -> $value {
77                assert!(end <= self.len);
78                if end == self.len {
79                    return self.total;
80                }
81                let mut result: $value = 0;
82                for level in &self.levels {
83                    let block = end / $branch;
84                    let lane = end % $branch;
85                    if lane != 0 {
86                        // SAFETY: `end < len`; mapping it upward remains within every level.
87                        let value = unsafe { level.get_unchecked(block).0.get_unchecked(lane - 1) };
88                        result = result.wrapping_add(*value);
89                    }
90                    end = block;
91                }
92                result
93            }
94
95            /// Returns the wrapping sum of `0..=index`.
96            #[inline]
97            pub fn accumulate(&self, index: usize) -> $value {
98                self.accumulate0(index + 1)
99            }
100
101            /// Returns the wrapping sum of `left..right`.
102            #[inline]
103            pub fn fold(&self, left: usize, right: usize) -> $value {
104                assert!(left <= right && right <= self.len);
105                if right == self.len {
106                    return self.total.wrapping_sub(self.accumulate0(left));
107                }
108                let mut left = left;
109                let mut right = right;
110                let mut result: $value = 0;
111                // Both endpoints are below len here; ancestors remain in the stored levels.
112                for level in &self.levels {
113                    if left == right {
114                        break;
115                    }
116                    if right % $branch != 0 {
117                        result = result.wrapping_add(unsafe {
118                            *level
119                                .get_unchecked(right / $branch)
120                                .0
121                                .get_unchecked(right % $branch - 1)
122                        });
123                    }
124                    if left % $branch != 0 {
125                        result = result.wrapping_sub(unsafe {
126                            *level
127                                .get_unchecked(left / $branch)
128                                .0
129                                .get_unchecked(left % $branch - 1)
130                        });
131                    }
132                    left /= $branch;
133                    right /= $branch;
134                }
135                result
136            }
137
138            #[inline]
139            pub fn get(&self, index: usize) -> $value {
140                assert!(index < self.len);
141                let prefix = &self.levels[0][index / $branch].0;
142                let lane = index % $branch;
143                if lane == 0 {
144                    prefix[0]
145                } else {
146                    prefix[lane].wrapping_sub(prefix[lane - 1])
147                }
148            }
149
150            #[inline]
151            pub fn fold_all(&self) -> $value {
152                self.total
153            }
154
155            /// Returns the number of leading values whose inclusive prefix sum is at most `value`.
156            ///
157            /// # Panics
158            ///
159            /// Panics if a prefix sum has overflowed.
160            #[inline]
161            pub fn partition_point_acc(&self, value: $value) -> usize {
162                assert!(self.partition_valid, "prefix sum overflowed");
163                if value >= self.total {
164                    return self.len;
165                }
166                #[cfg(target_arch = "x86_64")]
167                return match self.backend {
168                    SimdBackend::Scalar => self.partition_point_scalar(value),
169                    // SAFETY: `simd_backend` only selects supported instruction sets. Tests and
170                    // standalone benchmarks pass supported backends to the private constructor.
171                    SimdBackend::Avx2 => unsafe { self.partition_point_avx2(value) },
172                    // SAFETY: same as above.
173                    SimdBackend::Avx512 => unsafe { self.partition_point_avx512(value) },
174                };
175                #[cfg(not(target_arch = "x86_64"))]
176                self.partition_point_scalar(value)
177            }
178
179            fn zeroed(len: usize, backend: SimdBackend) -> Self {
180                let _ = &backend;
181                let mut levels = Vec::new();
182                let mut level_len = len;
183                while level_len != 0 {
184                    level_len = level_len.div_ceil($branch);
185                    levels.push(vec![PrefixBlock([0; $branch]); level_len]);
186                    if level_len == 1 {
187                        break;
188                    }
189                }
190                Self {
191                    levels,
192                    len,
193                    total: 0,
194                    partition_valid: true,
195                    #[cfg(target_arch = "x86_64")]
196                    backend,
197                }
198            }
199
200            fn build(values: &[$value], backend: SimdBackend) -> Self {
201                let _ = &backend;
202                let mut levels = Vec::new();
203                let mut partition_valid = true;
204                let mut current = Vec::with_capacity(values.len().div_ceil($branch));
205                let mut blocks = Vec::with_capacity(current.capacity());
206                for chunk in values.chunks($branch) {
207                    let mut prefix = [0; $branch];
208                    let mut sum: $value = 0;
209                    for (index, &value) in chunk.iter().enumerate() {
210                        partition_valid &= sum.checked_add(value).is_some();
211                        sum = sum.wrapping_add(value);
212                        prefix[index] = sum;
213                    }
214                    prefix[chunk.len()..].fill(sum);
215                    blocks.push(PrefixBlock(prefix));
216                    current.push(sum);
217                }
218                if !blocks.is_empty() {
219                    levels.push(blocks);
220                }
221                while current.len() > 1 {
222                    let mut blocks = Vec::with_capacity(current.len().div_ceil($branch));
223                    for chunk in current.chunks($branch) {
224                        let mut prefix = [0; $branch];
225                        let mut sum: $value = 0;
226                        for (index, &value) in chunk.iter().enumerate() {
227                            partition_valid &= sum.checked_add(value).is_some();
228                            sum = sum.wrapping_add(value);
229                            prefix[index] = sum;
230                        }
231                        prefix[chunk.len()..].fill(sum);
232                        blocks.push(PrefixBlock(prefix));
233                    }
234                    current = blocks.iter().map(|block| block.0[$branch - 1]).collect();
235                    levels.push(blocks);
236                }
237                Self {
238                    levels,
239                    len: values.len(),
240                    total: current.first().copied().unwrap_or(0),
241                    partition_valid,
242                    #[cfg(target_arch = "x86_64")]
243                    backend,
244                }
245            }
246
247            #[inline]
248            fn add(&mut self, index: usize, value: $value) {
249                #[cfg(target_arch = "x86_64")]
250                match self.backend {
251                    SimdBackend::Scalar => self.add_scalar(index, value),
252                    // SAFETY: `simd_backend` only selects supported instruction sets. Tests and
253                    // standalone benchmarks pass supported backends to the private constructor.
254                    SimdBackend::Avx2 => unsafe { self.add_avx2(index, value) },
255                    // SAFETY: same as above.
256                    SimdBackend::Avx512 => unsafe { self.add_avx512(index, value) },
257                }
258                #[cfg(not(target_arch = "x86_64"))]
259                self.add_scalar(index, value);
260            }
261
262            #[inline]
263            fn add_scalar(&mut self, mut index: usize, value: $value) {
264                for level in &mut self.levels {
265                    let block = index / $branch;
266                    let lane = index % $branch;
267                    // SAFETY: the public methods validate the leaf index; construction fixes the
268                    // mapping from every child block to its parent.
269                    let prefix = &mut unsafe { level.get_unchecked_mut(block) }.0;
270                    for prefix in &mut prefix[lane..] {
271                        *prefix = prefix.wrapping_add(value);
272                    }
273                    index = block;
274                }
275            }
276
277            #[cfg(target_arch = "x86_64")]
278            #[target_feature(enable = "avx2")]
279            unsafe fn add_avx2(&mut self, mut index: usize, value: $value) {
280                for level in &mut self.levels {
281                    let block = index / $branch;
282                    let lane = index % $branch;
283                    let prefix = &mut unsafe { level.get_unchecked_mut(block) }.0;
284                    unsafe { simd::$add_avx2(prefix, lane, value) };
285                    index = block;
286                }
287            }
288
289            #[cfg(target_arch = "x86_64")]
290            #[target_feature(enable = "avx512f")]
291            unsafe fn add_avx512(&mut self, mut index: usize, value: $value) {
292                for level in &mut self.levels {
293                    let block = index / $branch;
294                    let lane = index % $branch;
295                    let prefix = &mut unsafe { level.get_unchecked_mut(block) }.0;
296                    unsafe { simd::$add_avx512(prefix, lane, value) };
297                    index = block;
298                }
299            }
300
301            #[inline(always)]
302            fn partition_point_by<F>(&self, mut value: $value, mut first_gt: F) -> usize
303            where
304                F: FnMut(&[$value; $branch], $value) -> usize,
305            {
306                let mut node = 0;
307                for level in self.levels.iter().rev() {
308                    // SAFETY: `value < total` and non-overflowing prefixes select a real child at
309                    // every level.
310                    let prefix = &unsafe { level.get_unchecked(node) }.0;
311                    let lane = first_gt(prefix, value);
312                    if lane != 0 {
313                        value = value.wrapping_sub(unsafe { *prefix.get_unchecked(lane - 1) });
314                    }
315                    node = node * $branch + lane;
316                }
317                node.min(self.len)
318            }
319
320            #[inline(always)]
321            fn partition_point_scalar(&self, value: $value) -> usize {
322                self.partition_point_by(value, |prefix, value| {
323                    prefix.partition_point(|&sum| sum <= value)
324                })
325            }
326
327            #[cfg(target_arch = "x86_64")]
328            #[target_feature(enable = "avx2")]
329            unsafe fn partition_point_avx2(&self, value: $value) -> usize {
330                self.partition_point_by(value, |prefix, value| unsafe {
331                    simd::$first_gt_avx2(prefix, value)
332                })
333            }
334
335            #[cfg(target_arch = "x86_64")]
336            #[target_feature(enable = "avx512f")]
337            unsafe fn partition_point_avx512(&self, value: $value) -> usize {
338                self.partition_point_by(value, |prefix, value| unsafe {
339                    simd::$first_gt_avx512(prefix, value)
340                })
341            }
342        }
343    };
344}
345
346define_dary_prefix_sum_tree!(
347    DaryPrefixSumTreeU32,
348    u32,
349    16,
350    add_suffix_u32x16_avx2,
351    first_gt_u32x16_avx2,
352    add_suffix_u32x16_avx512,
353    first_gt_u32x16_avx512
354);
355define_dary_prefix_sum_tree!(
356    DaryPrefixSumTreeU64,
357    u64,
358    8,
359    add_suffix_u64x8_avx2,
360    first_gt_u64x8_avx2,
361    add_suffix_u64x8_avx512,
362    first_gt_u64x8_avx512
363);
364
365#[cfg(test)]
366mod tests {
367    use super::*;
368    use crate::tools::Xorshift;
369    #[cfg(target_arch = "x86_64")]
370    use crate::tools::avx512_supported;
371
372    #[cfg(target_arch = "x86_64")]
373    fn backends() -> Vec<SimdBackend> {
374        let mut result = vec![SimdBackend::Scalar];
375        if is_x86_feature_detected!("avx2") {
376            result.push(SimdBackend::Avx2);
377        }
378        if avx512_supported() {
379            result.push(SimdBackend::Avx512);
380        }
381        result
382    }
383
384    #[cfg(not(target_arch = "x86_64"))]
385    fn backends() -> Vec<SimdBackend> {
386        vec![SimdBackend::Scalar]
387    }
388
389    #[test]
390    fn test_dary_prefix_sum_tree() {
391        let mut rng = Xorshift::default();
392        for len in [0, 1, 7, 8, 15, 16, 17, 255, 256, 257, 4095, 4096, 4097] {
393            let values: Vec<_> = (0..len).map(|_| rng.rand(8) as u32).collect();
394            for backend in backends() {
395                let mut actual = DaryPrefixSumTreeU32::build(&values, backend);
396                let mut expected = values.clone();
397                for step in 0..500 {
398                    if len != 0 {
399                        let index = rng.rand(len as u64) as usize;
400                        if step % 3 == 0 {
401                            let value = rng.rand(32) as u32;
402                            actual.set(index, value);
403                            expected[index] = value;
404                        } else {
405                            let value = rng.rand(8) as u32;
406                            actual.update(index, value);
407                            expected[index] += value;
408                        }
409                    }
410                    for end in [0, len / 2, len] {
411                        assert_eq!(actual.accumulate0(end), expected[..end].iter().sum());
412                    }
413                    if len != 0 {
414                        let left = rng.rand(len as u64) as usize;
415                        let right = left + rng.rand((len - left + 1) as u64) as usize;
416                        assert_eq!(actual.fold(left, right), expected[left..right].iter().sum());
417                        assert_eq!(actual.get(left), expected[left]);
418                    }
419                    let mut sum = 0;
420                    let prefix: Vec<_> = expected
421                        .iter()
422                        .map(|&value| {
423                            sum += value;
424                            sum
425                        })
426                        .collect();
427                    for value in [0, sum / 2, sum] {
428                        assert_eq!(
429                            actual.partition_point_acc(value),
430                            prefix.partition_point(|&prefix| prefix <= value)
431                        );
432                    }
433                }
434            }
435        }
436
437        let values: Vec<_> = (0..513).map(|_| rng.rand(16)).collect();
438        for backend in backends() {
439            let mut actual = DaryPrefixSumTreeU64::build(&values, backend);
440            let mut expected = values.clone();
441            for step in 0..1000 {
442                let index = rng.rand(expected.len() as u64) as usize;
443                let value = rng.rand64();
444                if step % 2 == 0 {
445                    actual.set(index, value);
446                    expected[index] = value;
447                } else {
448                    actual.update(index, value);
449                    expected[index] = expected[index].wrapping_add(value);
450                }
451                assert_eq!(actual.get(index), expected[index]);
452                let end = rng.rand(expected.len() as u64 + 1) as usize;
453                let start = rng.rand(end as u64 + 1) as usize;
454                assert_eq!(
455                    actual.fold(start, end),
456                    expected[start..end]
457                        .iter()
458                        .copied()
459                        .fold(0, u64::wrapping_add)
460                );
461                assert_eq!(
462                    actual.accumulate0(end),
463                    expected[..end].iter().copied().fold(0, u64::wrapping_add)
464                );
465            }
466            assert_eq!(
467                actual.fold_all(),
468                expected.iter().copied().fold(0, u64::wrapping_add)
469            );
470        }
471    }
472}