Skip to main content

competitive/algorithm/
sort.rs

1use std::{cmp::Ordering, ptr::copy_nonoverlapping};
2
3pub trait RadixSortKey: Copy {
4    const BYTES: usize;
5    fn radix_byte(self, byte: usize) -> usize;
6}
7
8macro_rules! unsigned_radix_sort_key {
9    ($($t:ty),* $(,)?) => {
10        $(
11            impl RadixSortKey for $t {
12                const BYTES: usize = (<$t>::BITS / 8) as usize;
13
14                fn radix_byte(self, byte: usize) -> usize {
15                    ((self >> (byte * 8)) & 0xff) as usize
16                }
17            }
18        )*
19    };
20}
21
22macro_rules! signed_radix_sort_key {
23    ($($t:ty => $u:ty),* $(,)?) => {
24        $(
25            impl RadixSortKey for $t {
26                const BYTES: usize = (<$t>::BITS / 8) as usize;
27
28                fn radix_byte(self, byte: usize) -> usize {
29                    let key = self as $u ^ (1 << (<$t>::BITS - 1));
30                    ((key >> (byte * 8)) & 0xff) as usize
31                }
32            }
33        )*
34    };
35}
36
37unsigned_radix_sort_key!(u8, u16, u32, u64, u128, usize);
38signed_radix_sort_key!(
39    i8 => u8,
40    i16 => u16,
41    i32 => u32,
42    i64 => u64,
43    i128 => u128,
44    isize => usize,
45);
46
47pub trait SliceSortExt<T> {
48    fn bubble_sort(&mut self)
49    where
50        T: Ord;
51    fn bubble_sort_by<F>(&mut self, compare: F)
52    where
53        F: FnMut(&T, &T) -> Ordering;
54    fn merge_sort(&mut self)
55    where
56        T: Ord;
57    fn merge_sort_by<F>(&mut self, compare: F)
58    where
59        F: FnMut(&T, &T) -> Ordering;
60    fn insertion_sort(&mut self)
61    where
62        T: Ord;
63    fn insertion_sort_by<F>(&mut self, compare: F)
64    where
65        F: FnMut(&T, &T) -> Ordering;
66    fn radix_sort_by_key<K>(&mut self, key: impl FnMut(&T) -> K)
67    where
68        T: Clone,
69        K: RadixSortKey;
70}
71impl<T> SliceSortExt<T> for [T] {
72    fn bubble_sort(&mut self)
73    where
74        T: Ord,
75    {
76        bubble_sort(self, |a, b| a.lt(b));
77    }
78    fn bubble_sort_by<F>(&mut self, mut compare: F)
79    where
80        F: FnMut(&T, &T) -> Ordering,
81    {
82        bubble_sort(self, |a, b| compare(a, b) == Ordering::Less);
83    }
84    fn merge_sort(&mut self)
85    where
86        T: Ord,
87    {
88        merge_sort(self, |a, b| a.lt(b));
89    }
90    fn merge_sort_by<F>(&mut self, mut compare: F)
91    where
92        F: FnMut(&T, &T) -> Ordering,
93    {
94        merge_sort(self, |a, b| compare(a, b) == Ordering::Less);
95    }
96    fn insertion_sort(&mut self)
97    where
98        T: Ord,
99    {
100        insertion_sort(self, |a, b| a.lt(b));
101    }
102    fn insertion_sort_by<F>(&mut self, mut compare: F)
103    where
104        F: FnMut(&T, &T) -> Ordering,
105    {
106        insertion_sort(self, |a, b| compare(a, b) == Ordering::Less);
107    }
108    fn radix_sort_by_key<K>(&mut self, key: impl FnMut(&T) -> K)
109    where
110        T: Clone,
111        K: RadixSortKey,
112    {
113        radix_sort_by_key(self, key);
114    }
115}
116
117fn radix_sort_by_key<T, K>(values: &mut [T], mut key: impl FnMut(&T) -> K)
118where
119    T: Clone,
120    K: RadixSortKey,
121{
122    if values.len() <= 1 {
123        return;
124    }
125    let mut histograms = vec![[0usize; 256]; K::BYTES];
126    for value in values.iter() {
127        let key = key(value);
128        for (byte, histogram) in histograms.iter_mut().enumerate() {
129            histogram[key.radix_byte(byte)] += 1;
130        }
131    }
132    for histogram in histograms.iter_mut() {
133        let mut position = 0;
134        for count in histogram.iter_mut() {
135            let next = position + *count;
136            *count = position;
137            position = next;
138        }
139    }
140    let mut ends = [values.len(); 256];
141    ends[..255].copy_from_slice(&histograms[0][1..]);
142    let mut buffer = Vec::with_capacity(values.len());
143    {
144        let spare = buffer.spare_capacity_mut();
145        let positions = &mut histograms[0];
146        for value in values.iter() {
147            let bucket = key(value).radix_byte(0);
148            let position = &mut positions[bucket];
149            assert!(*position < ends[bucket]);
150            spare[*position].write(value.clone());
151            *position += 1;
152        }
153    }
154    // Every bucket filled its disjoint range completely.
155    unsafe { buffer.set_len(values.len()) };
156    for (byte, positions) in histograms.iter_mut().enumerate().skip(1) {
157        macro_rules! distribute {
158            ($source:expr, $destination:expr) => {{
159                let source = $source;
160                let destination = $destination;
161                for value in source {
162                    let bucket = key(value).radix_byte(byte);
163                    let position = &mut positions[bucket];
164                    destination[*position].clone_from(value);
165                    *position += 1;
166                }
167            }};
168        }
169        if byte % 2 == 0 {
170            distribute!(&*values, &mut buffer);
171        } else {
172            distribute!(&buffer, &mut *values);
173        }
174    }
175    if K::BYTES % 2 == 1 {
176        values.clone_from_slice(&buffer);
177    }
178}
179
180fn bubble_sort<T, F>(v: &mut [T], mut is_less: F)
181where
182    F: FnMut(&T, &T) -> bool,
183{
184    let len = v.len();
185    if len <= 1 {
186        return;
187    }
188    for i in 0..len - 1 {
189        for j in 0..len - i - 1 {
190            unsafe {
191                if is_less(v.get_unchecked(j + 1), v.get_unchecked(j)) {
192                    v.swap(j, j + 1);
193                }
194            }
195        }
196    }
197}
198
199unsafe fn merge<T, F>(v: &mut [T], mut mid: usize, buf: *mut T, is_less: &mut F)
200where
201    F: FnMut(&T, &T) -> bool,
202{
203    unsafe {
204        let len = v.len();
205        let v = v.as_mut_ptr();
206        let (v_mid, v_end) = (v.add(mid), v.add(len));
207
208        copy_nonoverlapping(v, buf, mid);
209        let mut start = buf;
210        let end = buf.add(mid);
211        let mut dest = v;
212
213        let left = &mut start;
214        let mut right = v_mid;
215        while *left < end && right < v_end {
216            let to_copy = if is_less(&*right, &**left) {
217                get_and_increment(&mut right)
218            } else {
219                mid -= 1;
220                get_and_increment(left)
221            };
222            copy_nonoverlapping(to_copy, get_and_increment(&mut dest), 1);
223        }
224
225        // let len = end.sub_ptr(start);
226        copy_nonoverlapping(start, dest, mid);
227    }
228
229    unsafe fn get_and_increment<T>(ptr: &mut *mut T) -> *mut T {
230        let old = *ptr;
231        *ptr = unsafe { ptr.offset(1) };
232        old
233    }
234}
235
236fn merge_sort<T, F>(v: &mut [T], mut is_less: F)
237where
238    F: FnMut(&T, &T) -> bool,
239{
240    let len = v.len();
241    if len <= 1 {
242        return;
243    }
244    let mut buf = Vec::with_capacity(len / 2);
245    let mut runs: Vec<Run> = vec![];
246    let mut end = len;
247    while end > 0 {
248        let start = end - 1;
249        let mut left = Run {
250            start,
251            len: end - start,
252        };
253        end = start;
254
255        while let Some(right) = runs.pop_if(|right| left.start == 0 || right.len <= left.len) {
256            unsafe {
257                merge(
258                    &mut v[left.start..right.start + right.len],
259                    left.len,
260                    buf.as_mut_ptr(),
261                    &mut is_less,
262                );
263            }
264            left = Run {
265                start: left.start,
266                len: left.len + right.len,
267            };
268        }
269        runs.push(left);
270    }
271
272    debug_assert!(runs.len() == 1 && runs[0].start == 0 && runs[0].len == len);
273
274    #[derive(Clone, Copy)]
275    struct Run {
276        start: usize,
277        len: usize,
278    }
279}
280
281fn insertion_sort<T, F>(v: &mut [T], mut is_less: F)
282where
283    F: FnMut(&T, &T) -> bool,
284{
285    for i in 1..v.len() {
286        let x = &v[i];
287        let p = v[..i].partition_point(|y| is_less(y, x));
288        v[p..=i].rotate_right(1);
289    }
290}
291
292#[cfg(test)]
293mod tests {
294    use super::*;
295    use crate::tools::{
296        Xorshift,
297        testutil::{exhaustive_sequences, sample_usize},
298    };
299
300    #[test]
301    fn test_comparison_sorts() {
302        let mut rng = Xorshift::default();
303        let mut cases: Vec<_> = exhaustive_sequences(-1..=1, 0..=8).collect();
304        for n in sample_usize(&mut rng, 16, 0..=3000, 100) {
305            let bound = rng.random(0..=1000i32);
306            let values: Vec<_> = rng.random_iter(-bound..=bound).take(n).collect();
307            cases.push(vec![0; n]);
308            cases.push((0..n as i32).collect());
309            cases.push((0..n as i32).rev().collect());
310            cases.push(values);
311        }
312        for values in cases {
313            let mut expected = values.clone();
314            expected.sort();
315            let mut actual = values.clone();
316            actual.bubble_sort();
317            assert_eq!(actual, expected);
318            let mut actual = values.clone();
319            actual.merge_sort();
320            assert_eq!(actual, expected);
321            let mut actual = values;
322            actual.insertion_sort();
323            assert_eq!(actual, expected);
324        }
325    }
326
327    #[test]
328    fn test_large_comparison_sorts() {
329        let mut rng = Xorshift::default();
330        for n in sample_usize(&mut rng, 16, 0..=100_000, 10) {
331            let values: Vec<i32> = rng.random_iter(..).take(n).collect();
332            let mut expected = values.clone();
333            expected.sort();
334            let mut actual = values.clone();
335            actual.merge_sort();
336            assert_eq!(actual, expected);
337            let mut actual = values;
338            actual.insertion_sort();
339            assert_eq!(actual, expected);
340        }
341    }
342
343    #[test]
344    fn test_radix_sort() {
345        let mut rng = Xorshift::default();
346        macro_rules! test_types {
347            ($($t:ty),* $(,)?) => {
348                $(
349                    for _ in 0..20 {
350                        let n = rng.random(0..1000);
351                        let mut actual: Vec<($t, usize)> =
352                            rng.random_iter(..).take(n).zip(0..).collect();
353                        let mut expected = actual.clone();
354                        actual.radix_sort_by_key(|&(key, _)| key);
355                        expected.sort_by_key(|&(key, _)| key);
356                        assert_eq!(actual, expected);
357                    }
358                )*
359            };
360        }
361        test_types!(u8, u16, u32, u64, u128, usize);
362        test_types!(i8, i16, i32, i64, i128, isize);
363        for _ in 0..20 {
364            let n = rng.random(0..1000);
365            let mut actual: Vec<_> = rng
366                .random_iter(0u32..64)
367                .take(n)
368                .zip(0..)
369                .map(|(key, index)| (key, index.to_string()))
370                .collect();
371            let mut expected = actual.clone();
372            actual.radix_sort_by_key(|value| value.0);
373            expected.sort_by_key(|value| value.0);
374            assert_eq!(actual, expected);
375        }
376    }
377}