Skip to main content

competitive/data_structure/
radix_heap.rs

1macro_rules! define_radix_heap {
2    ($name:ident, $key:ty, $buckets:expr) => {
3        /// A min-priority queue whose removed keys are monotonically nondecreasing.
4        ///
5        /// Values with equal keys have no specified removal order.
6        #[derive(Clone, Debug)]
7        pub struct $name<T> {
8            buckets: [Vec<($key, T)>; $buckets],
9            last: $key,
10            len: usize,
11        }
12
13        impl<T> $name<T> {
14            pub fn new() -> Self {
15                Self {
16                    buckets: std::array::from_fn(|_| Vec::new()),
17                    last: 0,
18                    len: 0,
19                }
20            }
21
22            pub fn len(&self) -> usize {
23                self.len
24            }
25
26            pub fn is_empty(&self) -> bool {
27                self.len == 0
28            }
29
30            /// Inserts a value whose key is not less than the key most recently removed.
31            ///
32            /// # Panics
33            ///
34            /// Panics if `key` is less than the key most recently removed.
35            pub fn push(&mut self, key: $key, value: T) {
36                assert!(key >= self.last, "key is less than the last removed key");
37                self.buckets[Self::bucket_index(key, self.last)].push((key, value));
38                self.len += 1;
39            }
40
41            pub fn pop(&mut self) -> Option<($key, T)> {
42                if self.len == 0 {
43                    return None;
44                }
45                if self.buckets[0].is_empty() {
46                    let index = (1..self.buckets.len())
47                        .find(|&index| !self.buckets[index].is_empty())
48                        .unwrap();
49                    self.last = self.buckets[index]
50                        .iter()
51                        .map(|&(key, _)| key)
52                        .min()
53                        .unwrap();
54                    let mut values = std::mem::take(&mut self.buckets[index]);
55                    while let Some((key, value)) = values.pop() {
56                        let next = Self::bucket_index(key, self.last);
57                        debug_assert!(next < index);
58                        self.buckets[next].push((key, value));
59                    }
60                    self.buckets[index] = values;
61                }
62                self.len -= 1;
63                self.buckets[0].pop()
64            }
65
66            pub fn clear(&mut self) {
67                for bucket in &mut self.buckets {
68                    bucket.clear();
69                }
70                self.last = 0;
71                self.len = 0;
72            }
73
74            #[inline]
75            fn bucket_index(key: $key, last: $key) -> usize {
76                (<$key>::BITS - (key ^ last).leading_zeros()) as usize
77            }
78        }
79
80        impl<T> Default for $name<T> {
81            fn default() -> Self {
82                Self::new()
83            }
84        }
85    };
86}
87
88define_radix_heap!(RadixHeapU32, u32, 33);
89define_radix_heap!(RadixHeapU64, u64, 65);
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94    use crate::tools::Xorshift;
95    use crate::tools::testutil::{exhaustive_sequences, integer_boundary_values};
96    use std::{cmp::Reverse, collections::BinaryHeap};
97
98    #[test]
99    fn test_radix_heap() {
100        macro_rules! check {
101            ($heap:ident, $key:ty) => {{
102                let mut rng = Xorshift::default();
103                let mut cases: Vec<_> = exhaustive_sequences(0..3, 0..=6).collect();
104                let boundaries = integer_boundary_values!($key);
105                cases.push(boundaries.iter().chain(&boundaries).copied().collect());
106                for values in cases {
107                    let mut actual = $heap::new();
108                    let mut expected: Vec<_> = values
109                        .iter()
110                        .copied()
111                        .enumerate()
112                        .map(|(i, key)| (key, i))
113                        .collect();
114                    rng.shuffle(&mut expected);
115                    for &(key, value) in &expected {
116                        actual.push(key, value);
117                    }
118                    let mut cleared = actual.clone();
119                    cleared.pop();
120                    cleared.clear();
121                    assert_eq!(cleared.len(), 0);
122                    assert!(cleared.is_empty());
123                    assert_eq!(cleared.pop(), None);
124                    let pair = (<$key>::MIN, values.len());
125                    cleared.push(pair.0, pair.1);
126                    assert_eq!(cleared.pop(), Some(pair));
127                    assert_eq!(cleared.pop(), None);
128                    for &(key, value) in &expected {
129                        cleared.push(key, value);
130                    }
131                    expected.sort();
132                    for mut heap in [actual, cleared] {
133                        let mut result = Vec::new();
134                        while let Some(pair) = heap.pop() {
135                            result.push(pair);
136                        }
137                        assert!(result.windows(2).all(|w| w[0].0 <= w[1].0));
138                        result.sort();
139                        assert_eq!(result, expected);
140                    }
141                }
142                let mut actual = $heap::new();
143                let mut expected = BinaryHeap::new();
144                let mut rng = Xorshift::default();
145                for value in 0..4096 {
146                    let key = rng.rand(1_000_000) as $key;
147                    actual.push(key, value);
148                    expected.push(Reverse((key, value)));
149                }
150                for value in 4096..104_096 {
151                    let (key, _) = actual.pop().unwrap();
152                    let Reverse((expected_key, _)) = expected.pop().unwrap();
153                    assert_eq!(key, expected_key);
154                    let key = key.saturating_add(rng.rand(1_000_000) as $key);
155                    actual.push(key, value);
156                    expected.push(Reverse((key, value)));
157                }
158                while let Some((key, _)) = actual.pop() {
159                    let Reverse((expected_key, _)) = expected.pop().unwrap();
160                    assert_eq!(key, expected_key);
161                }
162                assert!(expected.is_empty());
163                assert!(actual.is_empty());
164            }};
165        }
166
167        check!(RadixHeapU32, u32);
168        check!(RadixHeapU64, u64);
169    }
170}