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 HeapBlock<T, const D: usize>([T; D]);
8
9impl<T: Copy, const D: usize> HeapBlock<T, D> {
10 #[inline(always)]
11 fn filled(value: T) -> Self {
12 Self([value; D])
13 }
14
15 #[inline(always)]
16 fn get(&self, index: usize) -> T {
17 unsafe { *self.0.get_unchecked(index) }
18 }
19
20 #[inline(always)]
21 fn set(&mut self, index: usize, value: T) {
22 unsafe {
23 *self.0.get_unchecked_mut(index) = value;
24 }
25 }
26}
27
28#[repr(C, align(64))]
29#[derive(Clone, Debug)]
30struct U128HeapBlock {
31 low: [u64; 4],
32 high: [u64; 4],
33}
34
35impl U128HeapBlock {
36 #[inline(always)]
37 fn filled(value: u128) -> Self {
38 Self {
39 low: [value as u64; 4],
40 high: [(value >> 64) as u64; 4],
41 }
42 }
43
44 #[inline(always)]
45 fn get(&self, index: usize) -> u128 {
46 unsafe {
47 (*self.high.get_unchecked(index) as u128) << 64 | *self.low.get_unchecked(index) as u128
48 }
49 }
50
51 #[inline(always)]
52 fn set(&mut self, index: usize, value: u128) {
53 unsafe {
54 *self.low.get_unchecked_mut(index) = value as u64;
55 *self.high.get_unchecked_mut(index) = (value >> 64) as u64;
56 }
57 }
58}
59
60#[inline(always)]
61fn max_index<T: Copy + Ord, const D: usize>(values: &[T; D]) -> usize {
62 let mut maximum = values[0];
64 let mut result = 0;
65 for (index, &value) in values.iter().enumerate().skip(1) {
66 if value > maximum {
67 maximum = value;
68 result = index;
69 }
70 }
71 result
72}
73
74#[inline(always)]
75fn max_index_u128(values: &U128HeapBlock) -> usize {
76 let mut result = 0;
77 for index in 1..4 {
78 if values.high[index] > values.high[result]
79 || (values.high[index] == values.high[result] && values.low[index] > values.low[result])
80 {
81 result = index;
82 }
83 }
84 result
85}
86
87macro_rules! define_dary_heap {
88 (
89 $name:ident,
90 $doc:literal,
91 $value:ty,
92 $storage:ty,
93 $branch:expr,
94 $block:ty
95 , encode = $encode:expr
96 , decode = $decode:expr
97 $(, $field:ident: $field_type:ty = $field_value:expr)*
98 $(,)?
99 ) => {
100 #[doc = $doc]
101 #[derive(Clone, Debug)]
102 pub struct $name {
103 root: $storage,
104 blocks: Vec<$block>,
105 len: usize,
106 $(#[cfg(target_arch = "x86_64")] $field: $field_type,)*
107 }
108
109 impl $name {
110 pub fn new() -> Self {
111 Self::with_capacity(0)
112 }
113
114 pub fn with_capacity(capacity: usize) -> Self {
115 Self::empty(capacity $(, $field_value)*)
116 }
117
118 #[inline]
119 pub fn len(&self) -> usize {
120 self.len
121 }
122
123 #[inline]
124 pub fn is_empty(&self) -> bool {
125 self.len == 0
126 }
127
128 #[inline]
129 pub fn peek(&self) -> Option<$value> {
130 (self.len != 0).then(|| Self::decode(self.root))
131 }
132
133 pub fn push(&mut self, value: $value) {
134 let value = Self::encode(value);
135 if self.len == 0 {
136 self.root = value;
137 self.len = 1;
138 return;
139 }
140 let mut hole = self.len;
141 let block = (hole - 1) / $branch;
142 if block == self.blocks.len() {
143 self.blocks.push(<$block>::filled(<$storage>::MIN));
144 }
145 self.len += 1;
146 while hole != 0 {
147 let parent = (hole - 1) / $branch;
148 let parent_key = self.key(parent);
149 if parent_key >= value {
150 break;
151 }
152 self.set_key(hole, parent_key);
153 hole = parent;
154 }
155 self.set_key(hole, value);
156 }
157
158 pub fn pop(&mut self) -> Option<$value> {
159 if self.len == 0 {
160 return None;
161 }
162 let result = self.root;
163 if self.len == 1 {
164 self.root = <$storage>::MIN;
165 self.len = 0;
166 return Some(Self::decode(result));
167 }
168 let last = self.len - 1;
169 let value = self.key(last);
170 self.set_key(last, <$storage>::MIN);
171 self.len = last;
172 self.sift_down_after_pop(0, value);
173 Some(Self::decode(result))
174 }
175
176 pub fn replace(&mut self, value: $value) -> Option<$value> {
178 if self.len == 0 {
179 self.push(value);
180 return None;
181 }
182 let result = self.root;
183 self.sift_down(0, Self::encode(value));
184 Some(Self::decode(result))
185 }
186
187 pub fn clear(&mut self) {
188 self.root = <$storage>::MIN;
189 self.blocks.clear();
190 self.len = 0;
191 }
192
193 pub fn into_sorted_vec(mut self) -> Vec<$value> {
194 let mut values = Vec::with_capacity(self.len);
195 while let Some(value) = self.pop() {
196 values.push(value);
197 }
198 values.reverse();
199 values
200 }
201
202 fn empty(capacity: usize $(, $field: $field_type)*) -> Self {
203 $({ let _ = &$field; })*
204 Self {
205 root: <$storage>::MIN,
206 blocks: Vec::with_capacity(capacity.saturating_sub(1).div_ceil($branch)),
207 len: 0,
208 $(#[cfg(target_arch = "x86_64")] $field,)*
209 }
210 }
211
212 fn build(values: Vec<$value> $(, $field: $field_type)*) -> Self {
213 let len = values.len();
214 let mut heap = Self::empty(len $(, $field)*);
215 heap.len = len;
216 if let Some((&root, values)) = values.split_first() {
217 heap.root = Self::encode(root);
218 heap.blocks.resize(
219 len.saturating_sub(1).div_ceil($branch),
220 <$block>::filled(<$storage>::MIN),
221 );
222 for (index, &value) in values.iter().enumerate() {
223 heap.blocks[index / $branch].set(index % $branch, Self::encode(value));
224 }
225 heap.heapify();
226 }
227 heap
228 }
229
230 #[inline(always)]
231 fn encode(value: $value) -> $storage {
232 ($encode)(value)
233 }
234
235 #[inline(always)]
236 fn decode(value: $storage) -> $value {
237 ($decode)(value)
238 }
239
240 #[inline(always)]
241 fn key(&self, index: usize) -> $storage {
242 if index == 0 {
243 self.root
244 } else {
245 unsafe { self.blocks.get_unchecked((index - 1) / $branch) }
247 .get((index - 1) % $branch)
248 }
249 }
250
251 #[inline(always)]
252 fn set_key(&mut self, index: usize, value: $storage) {
253 if index == 0 {
254 self.root = value;
255 } else {
256 unsafe { self.blocks.get_unchecked_mut((index - 1) / $branch) }
259 .set((index - 1) % $branch, value);
260 }
261 }
262
263 #[inline(always)]
264 fn sift_down_by<F>(&mut self, mut hole: usize, value: $storage, mut max_index: F)
265 where
266 F: FnMut(&$block) -> usize,
267 {
268 if self.len <= 1 {
269 self.set_key(hole, value);
270 return;
271 }
272 let last_parent = (self.len - 2) / $branch;
273 while hole <= last_parent {
274 let block = unsafe { self.blocks.get_unchecked(hole) };
278 let lane = max_index(block);
279 let child_key = block.get(lane);
280 if child_key <= value {
281 break;
282 }
283 self.set_key(hole, child_key);
284 hole = hole * $branch + lane + 1;
285 }
286 self.set_key(hole, value);
287 }
288
289 fn heapify(&mut self) {
290 if self.len <= 1 {
291 return;
292 }
293 for parent in (0..=(self.len - 2) / $branch).rev() {
294 let value = self.key(parent);
295 self.sift_down(parent, value);
296 }
297 }
298 }
299
300 impl Default for $name {
301 fn default() -> Self {
302 Self::new()
303 }
304 }
305
306 impl From<Vec<$value>> for $name {
307 fn from(values: Vec<$value>) -> Self {
308 Self::build(values $(, $field_value)*)
309 }
310 }
311
312 impl Extend<$value> for $name {
313 fn extend<I>(&mut self, iter: I)
314 where
315 I: IntoIterator<Item = $value>,
316 {
317 for value in iter {
318 self.push(value);
319 }
320 }
321 }
322
323 impl FromIterator<$value> for $name {
324 fn from_iter<I>(iter: I) -> Self
325 where
326 I: IntoIterator<Item = $value>,
327 {
328 let values: Vec<_> = iter.into_iter().collect();
329 Self::from(values)
330 }
331 }
332 };
333}
334
335define_dary_heap!(
336 DaryHeapU32,
337 "A cache-line-oriented 16-ary max-heap for medium-to-large 32-bit heaps. `BinaryHeap` can be faster for small heaps and monotone replacements.",
338 u32,
339 u32,
340 16,
341 HeapBlock<u32, 16>,
342 encode = |value| value,
343 decode = |value| value,
344 backend: SimdBackend = simd_backend(),
345);
346define_dary_heap!(
347 DaryHeapI32,
348 "A cache-line-oriented 16-ary max-heap for medium-to-large 32-bit heaps. `BinaryHeap` can be faster for small heaps and monotone replacements.",
349 i32,
350 u32,
351 16,
352 HeapBlock<u32, 16>,
353 encode = |value: i32| value as u32 ^ (1 << 31),
354 decode = |value: u32| (value ^ (1 << 31)) as i32,
355 backend: SimdBackend = simd_backend(),
356);
357define_dary_heap!(
358 DaryHeapU64,
359 "A cache-line-oriented 8-ary max-heap for large 64-bit heaps. `BinaryHeap` can be faster for small heaps and monotone replacements.",
360 u64,
361 u64,
362 8,
363 HeapBlock<u64, 8>,
364 encode = |value| value,
365 decode = |value| value,
366 backend: SimdBackend = simd_backend(),
367);
368define_dary_heap!(
369 DaryHeapI64,
370 "A cache-line-oriented 8-ary max-heap for large 64-bit heaps. `BinaryHeap` can be faster for small heaps and monotone replacements.",
371 i64,
372 u64,
373 8,
374 HeapBlock<u64, 8>,
375 encode = |value: i64| value as u64 ^ (1 << 63),
376 decode = |value: u64| (value ^ (1 << 63)) as i64,
377 backend: SimdBackend = simd_backend(),
378);
379define_dary_heap!(
380 DaryHeapU128,
381 "A cache-line-oriented 4-ary max-heap for large full-width 128-bit heaps. `BinaryHeap` can be faster for small heaps, monotone replacements, and heavily repeated keys.",
382 u128,
383 u128,
384 4,
385 U128HeapBlock,
386 encode = |value| value,
387 decode = |value| value,
388 backend: SimdBackend = simd_backend(),
389);
390define_dary_heap!(
391 DaryHeapI128,
392 "A cache-line-oriented 4-ary max-heap for large full-width 128-bit heaps. `BinaryHeap` can be faster for small heaps, monotone replacements, and heavily repeated keys.",
393 i128,
394 u128,
395 4,
396 U128HeapBlock,
397 encode = |value: i128| value as u128 ^ (1 << 127),
398 decode = |value: u128| (value ^ (1 << 127)) as i128,
399 backend: SimdBackend = simd_backend(),
400);
401
402macro_rules! impl_simd_heap {
403 (
404 $name:ident,
405 $value:ty,
406 $branch:expr,
407 $max_avx2:ident,
408 $max_avx512:ident
409 ) => {
410 impl $name {
411 #[inline(always)]
412 fn sift_down_scalar(&mut self, hole: usize, value: $value) {
413 self.sift_down_by(hole, value, |block| max_index(&block.0))
414 }
415
416 #[cfg(target_arch = "x86_64")]
417 #[target_feature(enable = "avx2")]
418 unsafe fn sift_down_avx2(&mut self, hole: usize, value: $value) {
419 self.sift_down_by(hole, value, |block| unsafe { simd::$max_avx2(&block.0) })
420 }
421
422 #[cfg(target_arch = "x86_64")]
423 #[target_feature(enable = "avx512f")]
424 unsafe fn sift_down_avx512(&mut self, hole: usize, value: $value) {
425 self.sift_down_by(hole, value, |block| unsafe { simd::$max_avx512(&block.0) })
426 }
427
428 #[inline]
429 fn sift_down(&mut self, hole: usize, value: $value) {
430 #[cfg(target_arch = "x86_64")]
431 match self.backend {
432 SimdBackend::Scalar => self.sift_down_scalar(hole, value),
433 SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
436 SimdBackend::Avx512 => unsafe { self.sift_down_avx512(hole, value) },
438 }
439 #[cfg(not(target_arch = "x86_64"))]
440 self.sift_down_scalar(hole, value);
441 }
442
443 #[inline(always)]
444 fn sift_down_after_pop(&mut self, hole: usize, value: $value) {
445 self.sift_down(hole, value);
446 }
447 }
448 };
449}
450
451impl_simd_heap!(
452 DaryHeapU32,
453 u32,
454 16,
455 max_index_u32x16_avx2,
456 max_index_u32x16_avx512
457);
458impl_simd_heap!(
459 DaryHeapI32,
460 u32,
461 16,
462 max_index_u32x16_avx2,
463 max_index_u32x16_avx512
464);
465impl_simd_heap!(
466 DaryHeapU64,
467 u64,
468 8,
469 max_index_u64x8_avx2,
470 max_index_u64x8_avx512
471);
472impl_simd_heap!(
473 DaryHeapI64,
474 u64,
475 8,
476 max_index_u64x8_avx2,
477 max_index_u64x8_avx512
478);
479
480macro_rules! impl_u128_heap {
481 ($name:ident) => {
482 impl $name {
483 #[inline(always)]
484 fn sift_down_scalar(&mut self, hole: usize, value: u128) {
485 self.sift_down_by(hole, value, max_index_u128)
486 }
487
488 #[cfg(target_arch = "x86_64")]
489 #[target_feature(enable = "avx2")]
490 unsafe fn sift_down_avx2(&mut self, hole: usize, value: u128) {
491 self.sift_down_by(hole, value, |block| unsafe {
492 simd::max_index_u128x4_avx2(&block.low, &block.high)
493 })
494 }
495
496 #[cfg(target_arch = "x86_64")]
497 #[target_feature(enable = "avx2,avx512f,avx512vl")]
498 unsafe fn sift_down_avx512(&mut self, hole: usize, value: u128) {
499 self.sift_down_by(hole, value, |block| unsafe {
500 simd::max_index_u128x4_avx512(&block.low, &block.high)
501 })
502 }
503
504 #[inline]
505 fn sift_down(&mut self, hole: usize, value: u128) {
506 #[cfg(target_arch = "x86_64")]
507 match self.backend {
508 SimdBackend::Scalar => self.sift_down_scalar(hole, value),
509 SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
512 SimdBackend::Avx512 => unsafe { self.sift_down_avx512(hole, value) },
514 }
515 #[cfg(not(target_arch = "x86_64"))]
516 self.sift_down_scalar(hole, value);
517 }
518
519 #[inline]
520 fn sift_down_after_pop(&mut self, hole: usize, value: u128) {
521 #[cfg(target_arch = "x86_64")]
522 match self.backend {
523 SimdBackend::Scalar => self.sift_down_scalar(hole, value),
524 SimdBackend::Avx2 if self.len < 1 << 18 => self.sift_down_scalar(hole, value),
526 SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
529 SimdBackend::Avx512 if self.len < 1 << 15 => self.sift_down_scalar(hole, value),
531 SimdBackend::Avx512 => unsafe { self.sift_down_avx512(hole, value) },
533 }
534 #[cfg(not(target_arch = "x86_64"))]
535 self.sift_down_scalar(hole, value);
536 }
537 }
538 };
539}
540
541impl_u128_heap!(DaryHeapU128);
542impl_u128_heap!(DaryHeapI128);
543
544#[cfg(test)]
545mod tests {
546 use super::*;
547 use crate::tools::Xorshift;
548 #[cfg(target_arch = "x86_64")]
549 use crate::tools::avx512_supported;
550 use crate::tools::testutil::exhaustive_sequences;
551 use std::collections::BinaryHeap;
552
553 #[cfg(target_arch = "x86_64")]
554 fn backends() -> Vec<SimdBackend> {
555 let mut result = vec![SimdBackend::Scalar];
556 if is_x86_feature_detected!("avx2") {
557 result.push(SimdBackend::Avx2);
558 }
559 if avx512_supported() {
560 result.push(SimdBackend::Avx512);
561 }
562 result
563 }
564
565 #[cfg(not(target_arch = "x86_64"))]
566 fn backends() -> Vec<SimdBackend> {
567 vec![SimdBackend::Scalar]
568 }
569
570 #[test]
571 fn test_dary_heap() {
572 let mut rng = Xorshift::default();
573 for len in [0, 1, 7, 8, 15, 16, 17, 255, 256, 257, 4095, 4096, 4097] {
574 let mut values: Vec<_> = (0..len).map(|_| rng.rand64() as u32).collect();
575 values.extend([0, u32::MAX, u32::MAX]);
576 for backend in backends() {
577 let mut actual = DaryHeapU32::build(values.clone(), backend);
578 let mut expected = BinaryHeap::from(values.clone());
579 while !expected.is_empty() {
580 assert_eq!(actual.pop(), expected.pop());
581 }
582 assert_eq!(actual.pop(), None);
583 }
584 }
585
586 for backend in backends() {
587 let mut actual = DaryHeapU32::empty(10_000, backend);
588 let mut expected = BinaryHeap::new();
589 for _ in 0..20_000 {
590 match rng.rand(3) {
591 0 => {
592 let value = rng.rand64() as u32;
593 actual.push(value);
594 expected.push(value);
595 }
596 1 => assert_eq!(actual.pop(), expected.pop()),
597 _ => {
598 let value = rng.rand64() as u32;
599 let old = expected.pop();
600 expected.push(value);
601 assert_eq!(actual.replace(value), old);
602 }
603 }
604 assert_eq!(actual.peek(), expected.peek().copied());
605 assert_eq!(actual.len(), expected.len());
606 }
607 }
608
609 for backend in backends() {
610 let values: Vec<_> = (0..4097)
611 .map(|_| rng.rand64() as i32)
612 .chain([i32::MIN, 0, i32::MAX, i32::MAX])
613 .collect();
614 let mut actual = DaryHeapI32::build(values.clone(), backend);
615 let mut expected = BinaryHeap::from(values);
616 while !expected.is_empty() {
617 assert_eq!(actual.pop(), expected.pop());
618 }
619 }
620
621 macro_rules! check_heap {
622 ($heap:ty, $value:ty, $values:expr $(, $backend:expr)?) => {{
623 let values: Vec<$value> = ($values).collect();
624 let mut actual = <$heap>::build(values.clone() $(, $backend)?);
625 let mut expected = BinaryHeap::from(values);
626 while !expected.is_empty() {
627 assert_eq!(actual.pop(), expected.pop());
628 }
629 }};
630 }
631
632 for backend in backends() {
633 check_heap!(
634 DaryHeapU64,
635 u64,
636 (0..1025).map(|_| rng.rand64()).chain([0, u64::MAX]),
637 backend
638 );
639 check_heap!(
640 DaryHeapI64,
641 i64,
642 (0..1025)
643 .map(|_| rng.rand64() as i64)
644 .chain([i64::MIN, i64::MAX]),
645 backend
646 );
647 }
648 for backend in backends() {
649 check_heap!(
650 DaryHeapU128,
651 u128,
652 (0..1025)
653 .map(|_| (rng.rand64() as u128) << 64 | rng.rand64() as u128)
654 .chain([0, u128::MAX]),
655 backend
656 );
657 check_heap!(
658 DaryHeapI128,
659 i128,
660 (0..1025)
661 .map(|_| ((rng.rand64() as u128) << 64 | rng.rand64() as u128) as i128)
662 .chain([i128::MIN, i128::MAX]),
663 backend
664 );
665 }
666
667 for values in exhaustive_sequences([0u32, 1, u32::MAX], 0..=6) {
668 for backend in backends() {
669 let mut actual = DaryHeapU32::build(values.clone(), backend);
670 let mut cleared = actual.clone();
671 cleared.clear();
672 assert_eq!(cleared.len(), 0);
673 assert!(cleared.is_empty());
674 assert_eq!(cleared.peek(), None);
675 assert_eq!(cleared.pop(), None);
676 cleared.extend(values.iter().copied());
677 let mut expected = BinaryHeap::from(values.clone());
678 while let Some(value) = expected.pop() {
679 assert_eq!(actual.pop(), Some(value));
680 assert_eq!(cleared.pop(), Some(value));
681 }
682 assert_eq!(actual.pop(), None);
683 assert_eq!(cleared.pop(), None);
684 assert!(actual.is_empty());
685 let value: u32 = rng.random(..);
686 assert_eq!(actual.replace(value), None);
687 assert_eq!(actual.peek(), Some(value));
688 }
689 }
690 }
691}