competitive/data_structure/
dary_prefix_sum_tree.rs1#[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 #[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 #[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 #[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 #[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 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 #[inline]
97 pub fn accumulate(&self, index: usize) -> $value {
98 self.accumulate0(index + 1)
99 }
100
101 #[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 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 #[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 SimdBackend::Avx2 => unsafe { self.partition_point_avx2(value) },
172 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 SimdBackend::Avx2 => unsafe { self.add_avx2(index, value) },
255 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 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 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}