use super::bvh::MortonEntry;
use rayon::prelude::*;
const RADIX: usize = 256;
const BYTES: usize = 8;
const HISTOGRAM_BLOCK: usize = 1 << 16;
pub(super) fn sort_by_code(
entries: &mut Vec<MortonEntry>,
scratch: &mut Vec<MortonEntry>,
min_len: usize,
) {
if entries.len() < 2 {
return;
}
if entries.len() < min_len {
sort_serial(entries, scratch, 0..BYTES);
return;
}
let counts = histogram(entries, BYTES - 1);
let Some(offsets) = single_bucket_check(&counts) else {
sort_serial(entries, scratch, 0..BYTES - 1);
return;
};
scatter(entries, scratch, BYTES - 1, offsets);
std::mem::swap(entries, scratch);
sort_buckets(entries, &counts);
}
fn histogram(entries: &[MortonEntry], byte: usize) -> [usize; RADIX] {
entries
.par_chunks(HISTOGRAM_BLOCK)
.map(|block| {
let mut counts = [0usize; RADIX];
for entry in block {
counts[digit(entry.code, byte)] += 1;
}
counts
})
.reduce(
|| [0usize; RADIX],
|mut total, part| {
for (total, part) in total.iter_mut().zip(part) {
*total += part;
}
total
},
)
}
fn single_bucket_check(counts: &[usize; RADIX]) -> Option<[usize; RADIX]> {
if counts.iter().filter(|count| **count != 0).count() <= 1 {
return None;
}
let mut offsets = [0usize; RADIX];
let mut running = 0usize;
for (offset, count) in offsets.iter_mut().zip(counts) {
*offset = running;
running += count;
}
Some(offsets)
}
fn scatter(
source: &[MortonEntry],
target: &mut Vec<MortonEntry>,
byte: usize,
mut offsets: [usize; RADIX],
) {
target.clear();
target.resize(source.len(), MortonEntry::PLACEHOLDER);
for entry in source {
let bucket = digit(entry.code, byte);
if let Some(slot) = target.get_mut(offsets[bucket]) {
*slot = *entry;
}
offsets[bucket] += 1;
}
}
fn sort_buckets(entries: &mut [MortonEntry], counts: &[usize; RADIX]) {
let mut buckets: Vec<&mut [MortonEntry]> = Vec::with_capacity(RADIX);
let mut rest = entries;
for count in counts {
let take = (*count).min(rest.len());
let (bucket, tail) = rest.split_at_mut(take);
if bucket.len() > 1 {
buckets.push(bucket);
}
rest = tail;
}
buckets.into_par_iter().for_each(|bucket| {
let mut scratch = Vec::new();
sort_serial(bucket, &mut scratch, 0..BYTES - 1);
});
}
fn sort_serial(
entries: &mut [MortonEntry],
scratch: &mut Vec<MortonEntry>,
bytes: std::ops::Range<usize>,
) {
if entries.len() < 2 {
return;
}
scratch.clear();
scratch.resize(entries.len(), MortonEntry::PLACEHOLDER);
let mut flipped = false;
for byte in bytes {
let counts = {
let source: &[MortonEntry] = if flipped { scratch } else { entries };
let mut counts = [0usize; RADIX];
for entry in source {
counts[digit(entry.code, byte)] += 1;
}
counts
};
let Some(mut offsets) = single_bucket_check(&counts) else {
continue;
};
if flipped {
for entry in scratch.iter() {
let bucket = digit(entry.code, byte);
if let Some(slot) = entries.get_mut(offsets[bucket]) {
*slot = *entry;
}
offsets[bucket] += 1;
}
} else {
for entry in entries.iter() {
let bucket = digit(entry.code, byte);
if let Some(slot) = scratch.get_mut(offsets[bucket]) {
*slot = *entry;
}
offsets[bucket] += 1;
}
}
flipped = !flipped;
}
if flipped {
entries.copy_from_slice(scratch);
}
}
#[inline]
fn digit(code: u64, byte: usize) -> usize {
((code >> (byte * 8)) & 0xff) as usize
}
#[cfg(test)]
#[path = "radix_tests.rs"]
mod tests;