mirror of
https://github.com/storytold/spark.git
synced 2026-10-09 00:09:53 +00:00
Combine tallying of buckets for both passes of float32 sorting (#132)
This commit is contained in:
@@ -76,7 +76,9 @@ pub struct Sort32Buffers {
|
||||
/// output indices
|
||||
pub ordering: Vec<u32>,
|
||||
/// bucket counts / offsets (length == RADIX_BASE)
|
||||
pub buckets16: Vec<u32>,
|
||||
pub buckets16lo: Vec<u32>,
|
||||
/// bucket counts / offsets (length == RADIX_BASE)
|
||||
pub buckets16hi: Vec<u32>,
|
||||
/// scratch space for indices
|
||||
pub scratch: Vec<u32>,
|
||||
}
|
||||
@@ -93,8 +95,11 @@ impl Sort32Buffers {
|
||||
if self.scratch.len() < max_splats {
|
||||
self.scratch.resize(max_splats, 0);
|
||||
}
|
||||
if self.buckets16.len() < RADIX_BASE {
|
||||
self.buckets16.resize(RADIX_BASE, 0);
|
||||
if self.buckets16lo.len() < RADIX_BASE {
|
||||
self.buckets16lo.resize(RADIX_BASE, 0);
|
||||
}
|
||||
if self.buckets16hi.len() < RADIX_BASE {
|
||||
self.buckets16hi.resize(RADIX_BASE, 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -109,20 +114,24 @@ pub fn sort32_internal(
|
||||
// make sure our buffers can hold `max_splats`
|
||||
buffers.ensure_size(max_splats);
|
||||
|
||||
let Sort32Buffers { readback, ordering, buckets16, scratch } = buffers;
|
||||
let Sort32Buffers { readback, ordering, buckets16lo, buckets16hi, scratch } = buffers;
|
||||
let keys = &readback[..num_splats];
|
||||
|
||||
// ——— Pass #1: bucket by inv(low 16 bits) ———
|
||||
buckets16.fill(0);
|
||||
// tally low and high buckets
|
||||
buckets16lo.fill(0);
|
||||
buckets16hi.fill(0);
|
||||
for &key in keys.iter() {
|
||||
if key < DEPTH_INFINITY_F32 {
|
||||
let inv = !key;
|
||||
buckets16[(inv & 0xFFFF) as usize] += 1;
|
||||
buckets16lo[(inv & 0xFFFF) as usize] += 1;
|
||||
buckets16hi[(inv >> 16) as usize] += 1;
|
||||
}
|
||||
}
|
||||
|
||||
// ——— Pass #1: bucket by inv(low 16 bits) ———
|
||||
// exclusive prefix‑sum → starting offsets
|
||||
let mut total: u32 = 0;
|
||||
for slot in buckets16.iter_mut() {
|
||||
for slot in buckets16lo.iter_mut() {
|
||||
let cnt = *slot;
|
||||
*slot = total;
|
||||
total = total.wrapping_add(cnt);
|
||||
@@ -134,21 +143,15 @@ pub fn sort32_internal(
|
||||
if key < DEPTH_INFINITY_F32 {
|
||||
let inv = !key;
|
||||
let lo = (inv & 0xFFFF) as usize;
|
||||
scratch[buckets16[lo] as usize] = i as u32;
|
||||
buckets16[lo] += 1;
|
||||
scratch[buckets16lo[lo] as usize] = i as u32;
|
||||
buckets16lo[lo] += 1;
|
||||
}
|
||||
}
|
||||
|
||||
// ——— Pass #2: bucket by inv(high 16 bits) ———
|
||||
buckets16.fill(0);
|
||||
for &idx in scratch.iter().take(active_splats as usize) {
|
||||
let key = keys[idx as usize];
|
||||
let inv = !key;
|
||||
buckets16[(inv >> 16) as usize] += 1;
|
||||
}
|
||||
// exclusive prefix‑sum again
|
||||
let mut sum: u32 = 0;
|
||||
for slot in buckets16.iter_mut() {
|
||||
for slot in buckets16hi.iter_mut() {
|
||||
let cnt = *slot;
|
||||
*slot = sum;
|
||||
sum = sum.wrapping_add(cnt);
|
||||
@@ -158,16 +161,16 @@ pub fn sort32_internal(
|
||||
let key = keys[idx as usize];
|
||||
let inv = !key;
|
||||
let hi = (inv >> 16) as usize;
|
||||
ordering[buckets16[hi] as usize] = idx;
|
||||
buckets16[hi] += 1;
|
||||
ordering[buckets16hi[hi] as usize] = idx;
|
||||
buckets16hi[hi] += 1;
|
||||
}
|
||||
|
||||
// sanity‑check: last bucket should have consumed all entries
|
||||
if buckets16[RADIX_BASE - 1] != active_splats {
|
||||
if buckets16hi[RADIX_BASE - 1] != active_splats {
|
||||
return Err(anyhow!(
|
||||
"Expected {} active splats but got {}",
|
||||
active_splats,
|
||||
buckets16[RADIX_BASE - 1]
|
||||
buckets16hi[RADIX_BASE - 1]
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
+24
-22
@@ -527,7 +527,8 @@ function sortDoubleSplats({
|
||||
}
|
||||
|
||||
const DEPTH_INFINITY_F32 = 0x7f800000;
|
||||
let bucket16: Uint32Array | null = null;
|
||||
let bucket16lo: Uint32Array | null = null;
|
||||
let bucket16hi: Uint32Array | null = null;
|
||||
let scratchSplats: Uint32Array | null = null;
|
||||
|
||||
// two-pass radix sort (base 65536) of 32-bit keys in readback,
|
||||
@@ -546,29 +547,36 @@ function sort32Splats({
|
||||
const BASE = 1 << 16; // 65536
|
||||
|
||||
// allocate once
|
||||
if (!bucket16) {
|
||||
bucket16 = new Uint32Array(BASE);
|
||||
if (!bucket16lo) {
|
||||
bucket16lo = new Uint32Array(BASE);
|
||||
}
|
||||
if (!bucket16hi) {
|
||||
bucket16hi = new Uint32Array(BASE);
|
||||
}
|
||||
if (!scratchSplats || scratchSplats.length < maxSplats) {
|
||||
scratchSplats = new Uint32Array(maxSplats);
|
||||
}
|
||||
|
||||
//
|
||||
// ——— Pass #1: bucket by inv(lo 16 bits) ———
|
||||
//
|
||||
bucket16.fill(0);
|
||||
// tally low and high buckets
|
||||
bucket16lo.fill(0);
|
||||
bucket16hi.fill(0);
|
||||
for (let i = 0; i < numSplats; ++i) {
|
||||
const key = readback[i];
|
||||
if (key < DEPTH_INFINITY_F32) {
|
||||
const inv = ~key >>> 0;
|
||||
bucket16[inv & 0xffff] += 1;
|
||||
bucket16lo[inv & 0xffff] += 1;
|
||||
bucket16hi[inv >>> 16] += 1;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// ——— Pass #1: bucket by inv(lo 16 bits) ———
|
||||
//
|
||||
// exclusive prefix‑sum → starting offsets
|
||||
let total = 0;
|
||||
for (let b = 0; b < BASE; ++b) {
|
||||
const c = bucket16[b];
|
||||
bucket16[b] = total;
|
||||
const c = bucket16lo[b];
|
||||
bucket16lo[b] = total;
|
||||
total += c;
|
||||
}
|
||||
const activeSplats = total;
|
||||
@@ -578,24 +586,18 @@ function sort32Splats({
|
||||
const key = readback[i];
|
||||
if (key < DEPTH_INFINITY_F32) {
|
||||
const inv = ~key >>> 0;
|
||||
scratchSplats[bucket16[inv & 0xffff]++] = i;
|
||||
scratchSplats[bucket16lo[inv & 0xffff]++] = i;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// ——— Pass #2: bucket by inv(hi 16 bits) ———
|
||||
//
|
||||
bucket16.fill(0);
|
||||
for (let k = 0; k < activeSplats; ++k) {
|
||||
const idx = scratchSplats[k];
|
||||
const inv = ~readback[idx] >>> 0;
|
||||
bucket16[inv >>> 16] += 1;
|
||||
}
|
||||
// exclusive prefix‑sum again
|
||||
let sum = 0;
|
||||
for (let b = 0; b < BASE; ++b) {
|
||||
const c = bucket16[b];
|
||||
bucket16[b] = sum;
|
||||
const c = bucket16hi[b];
|
||||
bucket16hi[b] = sum;
|
||||
sum += c;
|
||||
}
|
||||
|
||||
@@ -603,13 +605,13 @@ function sort32Splats({
|
||||
for (let k = 0; k < activeSplats; ++k) {
|
||||
const idx = scratchSplats[k];
|
||||
const inv = ~readback[idx] >>> 0;
|
||||
ordering[bucket16[inv >>> 16]++] = idx;
|
||||
ordering[bucket16hi[inv >>> 16]++] = idx;
|
||||
}
|
||||
|
||||
// sanity‑check: the last bucket should have eaten all entries
|
||||
if (bucket16[BASE - 1] !== activeSplats) {
|
||||
if (bucket16hi[BASE - 1] !== activeSplats) {
|
||||
throw new Error(
|
||||
`Expected ${activeSplats} active splats but got ${bucket16[BASE - 1]}`,
|
||||
`Expected ${activeSplats} active splats but got ${bucket16hi[BASE - 1]}`,
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user