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 SimdBackend::Avx2 => unsafe { self.fold_avx2(range.start, range.end) },
104 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 SimdBackend::Avx2 => unsafe { self.set_avx2(index, value) },
173 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}