diff --git a/rust/spark-internal-rs/src/sort.rs b/rust/spark-internal-rs/src/sort.rs index 9b1f947..40dddd9 100644 --- a/rust/spark-internal-rs/src/sort.rs +++ b/rust/spark-internal-rs/src/sort.rs @@ -76,7 +76,9 @@ pub struct Sort32Buffers { /// output indices pub ordering: Vec, /// bucket counts / offsets (length == RADIX_BASE) - pub buckets16: Vec, + pub buckets16lo: Vec, + /// bucket counts / offsets (length == RADIX_BASE) + pub buckets16hi: Vec, /// scratch space for indices pub scratch: Vec, } @@ -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] )); } diff --git a/src/worker.ts b/src/worker.ts index 6ebfe9e..d06ee9f 100644 --- a/src/worker.ts +++ b/src/worker.ts @@ -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]}`, ); }