Skip to main content

competitive/data_structure/
dary_heap.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 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    // Track the value separately to avoid dependent array loads.
63    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            /// Unconditionally replaces the greatest value, or inserts into an empty heap.
177            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                    // SAFETY: callers only pass occupied heap indices.
246                    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                    // SAFETY: callers only pass occupied heap indices. `push` allocates the
257                    // destination block before increasing `len`.
258                    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                    // SAFETY: the loop condition guarantees a child block. Real children occupy
275                    // its prefix and padding is the type minimum, so the first maximum is always
276                    // a real child.
277                    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                    // SAFETY: automatic construction only selects supported instruction sets;
434                    // explicit construction is private and restricted to tests and benchmarks.
435                    SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
436                    // SAFETY: same as above.
437                    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                    // SAFETY: automatic construction only selects supported instruction sets;
510                    // explicit construction is private and restricted to tests and benchmarks.
511                    SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
512                    // SAFETY: same as above.
513                    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                    // Scalar selection avoids SIMD setup overhead below this crossover.
525                    SimdBackend::Avx2 if self.len < 1 << 18 => self.sift_down_scalar(hole, value),
526                    // SAFETY: automatic construction only selects supported instruction sets;
527                    // explicit construction is private and restricted to tests and benchmarks.
528                    SimdBackend::Avx2 => unsafe { self.sift_down_avx2(hole, value) },
529                    // Scalar selection avoids SIMD setup overhead below this crossover.
530                    SimdBackend::Avx512 if self.len < 1 << 15 => self.sift_down_scalar(hole, value),
531                    // SAFETY: same as above.
532                    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}