const RADIX_MIN_LEN: usize = 2048;
const RADIX_BITS: u32 = 11;
const RADIX_BUCKETS: usize = 1 << RADIX_BITS;
#[inline]
pub(super) fn sort_key(value: f32) -> u32 {
let bits = value.to_bits();
if bits & 0x8000_0000 != 0 {
!bits
} else {
bits | 0x8000_0000
}
}
#[inline]
pub(super) fn unsort_key(key: u32) -> f32 {
f32::from_bits(if key & 0x8000_0000 != 0 {
key & !0x8000_0000
} else {
!key
})
}
#[derive(Default)]
pub(super) struct RadixScratch<T> {
spare: Vec<T>,
counts: Vec<u32>,
}
pub(super) fn radix_sort<T: Copy + Default>(
items: &mut Vec<T>,
scratch: &mut RadixScratch<T>,
key: impl Fn(&T) -> u32,
) {
let n = items.len();
if n < 2 {
return;
}
let Ok(n32) = u32::try_from(n) else {
items.sort_by_key(key);
return;
};
let RadixScratch { spare, counts } = scratch;
counts.clear();
counts.resize(3 * RADIX_BUCKETS, 0);
for item in items.iter() {
let k = key(item);
for (pass, count) in counts
.as_chunks_mut::<RADIX_BUCKETS>()
.0
.iter_mut()
.enumerate()
{
count[((k >> (RADIX_BITS * pass as u32)) & (RADIX_BUCKETS as u32 - 1)) as usize] += 1;
}
}
spare.resize(n, T::default());
let mut in_spare = false;
for (pass, count) in counts
.as_chunks_mut::<RADIX_BUCKETS>()
.0
.iter_mut()
.enumerate()
{
if count.contains(&n32) {
continue;
}
let mut offset = 0;
for c in count.iter_mut() {
let start = offset;
offset += *c;
*c = start;
}
let shift = RADIX_BITS * pass as u32;
let (src, dst) = if in_spare {
(&*spare, &mut *items)
} else {
(&*items, &mut *spare)
};
for item in src {
let bucket = ((key(item) >> shift) & (RADIX_BUCKETS as u32 - 1)) as usize;
dst[count[bucket] as usize] = *item;
count[bucket] += 1;
}
in_spare = !in_spare;
}
if in_spare {
std::mem::swap(items, spare);
}
}
pub(super) fn sort_values(values: &mut Vec<f32>, scratch: &mut RadixScratch<f32>) {
if values.len() < RADIX_MIN_LEN {
values.sort_unstable_by(f32::total_cmp);
} else {
radix_sort(values, scratch, |&v| sort_key(v));
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn radix_sort_matches_total_order() {
let mut seed = 0x9E37_79B9_7F4A_7C15u64;
let mut next = || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
seed
};
for n in [RADIX_MIN_LEN, RADIX_MIN_LEN + 1, 10_007, 65_536] {
let mut values: Vec<f32> = (0..n)
.map(|i| match i % 11 {
0 => -0.0,
1 => 0.0,
2 => f32::MAX,
3 => f32::MIN,
4 => f32::MIN_POSITIVE,
5 => -f32::MIN_POSITIVE,
_ => (next() as f32 / u64::MAX as f32 - 0.5) * 1e6,
})
.collect();
let mut expected = values.clone();
expected.sort_unstable_by(f32::total_cmp);
sort_values(&mut values, &mut RadixScratch::default());
let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
assert_eq!(bits(&values), bits(&expected));
}
}
}