mirror of
https://github.com/storytold/spark.git
synced 2026-10-09 00:09:53 +00:00
Implement float32 sorting option (#129)
* Add commented benchmarkSort to compare JS vs Wasm sorting. * Add 32-bit float sort via SparkViewpooint.sort32, using 2-pass radix-65536 sort. Turn on sort32 by default for examples/editor. Implemented Rust sort and JS sort, updated benchmarking code to include float16 and float32 sort. * Add sort32 to docs/spark-viewpoint.md.
This commit is contained in:
committed by
GitHub
parent
2ab9ee1e78
commit
0bfb6c8963
@@ -23,6 +23,7 @@ const viewpoint = spark.newViewpoint({
|
||||
sortCoorient?: boolean;
|
||||
depthBias?: number;
|
||||
sort360?: boolean;
|
||||
sort32?: boolean;
|
||||
});
|
||||
```
|
||||
|
||||
@@ -44,6 +45,7 @@ const viewpoint = spark.newViewpoint({
|
||||
| **sortCoorient** | View direction dot product threshold for re-sorting splats. For `sortRadial: true` it defaults to 0.99 while `sortRadial: false` uses 0.999 because it is more sensitive to view direction. (default: `0.99` if `sortRadial` else `0.999`)
|
||||
| **depthBias** | Constant added to Z-depth to bias values into the positive range for `sortRadial: false`, but also used for culling splats "well behind" the viewpoint origin (default: `1.0`)
|
||||
| **sort360** | Set this to true if rendering a 360 to disable "behind the viewpoint" culling during sorting. This is set automatically when rendering 360 envMaps using the `SparkRenderer.renderEnvMap()` utility function. (default: `false`)
|
||||
| **sort32** | Set this to true to sort with float32 precision with two-pass sort. (default: `false`)
|
||||
|
||||
## `dispose()`
|
||||
|
||||
|
||||
@@ -480,6 +480,8 @@
|
||||
stats.dom.style.display = value ? "block" : "none";
|
||||
});
|
||||
gui.add(spark.defaultView, "sortRadial").name("Radial sort").listen();
|
||||
spark.defaultView.sort32 = true;
|
||||
gui.add(spark.defaultView, "sort32").name("Float32 sort").listen();
|
||||
gui.add(grid, "opacity", 0, 1, 0.01).name("Grid opacity").listen();
|
||||
gui.add({
|
||||
logFocalDistance: 0.0,
|
||||
|
||||
@@ -4,7 +4,7 @@ use js_sys::{Float32Array, Uint16Array, Uint32Array};
|
||||
use wasm_bindgen::prelude::*;
|
||||
|
||||
mod sort;
|
||||
use sort::{sort_internal, SortBuffers};
|
||||
use sort::{sort_internal, SortBuffers, sort32_internal, Sort32Buffers};
|
||||
|
||||
mod raycast;
|
||||
use raycast::{raycast_ellipsoids, raycast_spheres};
|
||||
@@ -13,6 +13,7 @@ const RAYCAST_BUFFER_COUNT: u32 = 65536;
|
||||
|
||||
thread_local! {
|
||||
static SORT_BUFFERS: RefCell<SortBuffers> = RefCell::new(SortBuffers::default());
|
||||
static SORT32_BUFFERS: RefCell<Sort32Buffers> = RefCell::new(Sort32Buffers::default());
|
||||
static RAYCAST_BUFFER: RefCell<Vec<u32>> = RefCell::new(vec![0; RAYCAST_BUFFER_COUNT as usize * 4]);
|
||||
}
|
||||
|
||||
@@ -45,6 +46,35 @@ pub fn sort_splats(
|
||||
active_splats
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn sort32_splats(
|
||||
num_splats: u32, readback: Uint32Array, ordering: Uint32Array,
|
||||
) -> u32 {
|
||||
let max_splats = readback.length() as usize;
|
||||
|
||||
let active_splats = SORT32_BUFFERS.with_borrow_mut(|buffers| {
|
||||
buffers.ensure_size(max_splats);
|
||||
let sub_readback = readback.subarray(0, num_splats);
|
||||
sub_readback.copy_to(&mut buffers.readback[..num_splats as usize]);
|
||||
|
||||
let active_splats = match sort32_internal(buffers, max_splats, num_splats as usize) {
|
||||
Ok(active_splats) => active_splats,
|
||||
Err(err) => {
|
||||
wasm_bindgen::throw_str(&format!("{}", err));
|
||||
}
|
||||
};
|
||||
|
||||
if active_splats > 0 {
|
||||
// Copy out ordering result
|
||||
let subarray = &buffers.ordering[..active_splats as usize];
|
||||
ordering.subarray(0, active_splats).copy_from(&subarray);
|
||||
}
|
||||
active_splats
|
||||
});
|
||||
|
||||
active_splats
|
||||
}
|
||||
|
||||
#[wasm_bindgen]
|
||||
pub fn raycast_splats(
|
||||
origin_x: f32, origin_y: f32, origin_z: f32,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use anyhow::anyhow;
|
||||
|
||||
const DEPTH_INFINITY: u32 = 0x7c00;
|
||||
const DEPTH_SIZE: usize = DEPTH_INFINITY as usize + 1;
|
||||
const DEPTH_INFINITY_F16: u32 = 0x7c00;
|
||||
const DEPTH_SIZE_F16: usize = DEPTH_INFINITY_F16 as usize + 1;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct SortBuffers {
|
||||
@@ -18,8 +18,8 @@ impl SortBuffers {
|
||||
if self.ordering.len() < max_splats {
|
||||
self.ordering.resize(max_splats, 0);
|
||||
}
|
||||
if self.buckets.len() < DEPTH_SIZE {
|
||||
self.buckets.resize(DEPTH_SIZE, 0);
|
||||
if self.buckets.len() < DEPTH_SIZE_F16 {
|
||||
self.buckets.resize(DEPTH_SIZE_F16, 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -30,11 +30,11 @@ pub fn sort_internal(buffers: &mut SortBuffers, num_splats: usize) -> anyhow::Re
|
||||
|
||||
// Set the bucket counts to zero
|
||||
buckets.clear();
|
||||
buckets.resize(DEPTH_SIZE, 0);
|
||||
buckets.resize(DEPTH_SIZE_F16, 0);
|
||||
|
||||
// Count the number of splats in each bucket
|
||||
for &metric in readback.iter() {
|
||||
if (metric as u32) < DEPTH_INFINITY {
|
||||
if (metric as u32) < DEPTH_INFINITY_F16 {
|
||||
buckets[metric as usize] += 1;
|
||||
}
|
||||
}
|
||||
@@ -49,7 +49,7 @@ pub fn sort_internal(buffers: &mut SortBuffers, num_splats: usize) -> anyhow::Re
|
||||
|
||||
// Write out splat indices at the right location using bucket offsets
|
||||
for (index, &metric) in readback.iter().enumerate() {
|
||||
if (metric as u32) < DEPTH_INFINITY {
|
||||
if (metric as u32) < DEPTH_INFINITY_F16 {
|
||||
ordering[buckets[metric as usize] as usize] = index as u32;
|
||||
buckets[metric as usize] += 1;
|
||||
}
|
||||
@@ -65,3 +65,111 @@ pub fn sort_internal(buffers: &mut SortBuffers, num_splats: usize) -> anyhow::Re
|
||||
}
|
||||
Ok(active_splats)
|
||||
}
|
||||
|
||||
const DEPTH_INFINITY_F32: u32 = 0x7f800000;
|
||||
const RADIX_BASE: usize = 1 << 16; // 65536
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct Sort32Buffers {
|
||||
/// raw f32 bit‑patterns (one per splat)
|
||||
pub readback: Vec<u32>,
|
||||
/// output indices
|
||||
pub ordering: Vec<u32>,
|
||||
/// bucket counts / offsets (length == RADIX_BASE)
|
||||
pub buckets16: Vec<u32>,
|
||||
/// scratch space for indices
|
||||
pub scratch: Vec<u32>,
|
||||
}
|
||||
|
||||
impl Sort32Buffers {
|
||||
/// ensure all internal buffers are large enough for up to `max_splats`
|
||||
pub fn ensure_size(&mut self, max_splats: usize) {
|
||||
if self.readback.len() < max_splats {
|
||||
self.readback.resize(max_splats, 0);
|
||||
}
|
||||
if self.ordering.len() < max_splats {
|
||||
self.ordering.resize(max_splats, 0);
|
||||
}
|
||||
if self.scratch.len() < max_splats {
|
||||
self.scratch.resize(max_splats, 0);
|
||||
}
|
||||
if self.buckets16.len() < RADIX_BASE {
|
||||
self.buckets16.resize(RADIX_BASE, 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Two‑pass radix sort (base 2¹⁶) of 32‑bit float bit‑patterns,
|
||||
/// descending order (largest keys first). Mirrors the JS `sort32Splats`.
|
||||
pub fn sort32_internal(
|
||||
buffers: &mut Sort32Buffers,
|
||||
max_splats: usize,
|
||||
num_splats: usize,
|
||||
) -> anyhow::Result<u32> {
|
||||
// make sure our buffers can hold `max_splats`
|
||||
buffers.ensure_size(max_splats);
|
||||
|
||||
let Sort32Buffers { readback, ordering, buckets16, scratch } = buffers;
|
||||
let keys = &readback[..num_splats];
|
||||
|
||||
// ——— Pass #1: bucket by inv(low 16 bits) ———
|
||||
buckets16.fill(0);
|
||||
for &key in keys.iter() {
|
||||
if key < DEPTH_INFINITY_F32 {
|
||||
let inv = !key;
|
||||
buckets16[(inv & 0xFFFF) as usize] += 1;
|
||||
}
|
||||
}
|
||||
// exclusive prefix‑sum → starting offsets
|
||||
let mut total: u32 = 0;
|
||||
for slot in buckets16.iter_mut() {
|
||||
let cnt = *slot;
|
||||
*slot = total;
|
||||
total = total.wrapping_add(cnt);
|
||||
}
|
||||
let active_splats = total;
|
||||
|
||||
// scatter into scratch by low bits of inv
|
||||
for (i, &key) in keys.iter().enumerate() {
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
// ——— 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() {
|
||||
let cnt = *slot;
|
||||
*slot = sum;
|
||||
sum = sum.wrapping_add(cnt);
|
||||
}
|
||||
// scatter into final ordering by high bits of inv
|
||||
for &idx in scratch.iter().take(active_splats as usize) {
|
||||
let key = keys[idx as usize];
|
||||
let inv = !key;
|
||||
let hi = (inv >> 16) as usize;
|
||||
ordering[buckets16[hi] as usize] = idx;
|
||||
buckets16[hi] += 1;
|
||||
}
|
||||
|
||||
// sanity‑check: last bucket should have consumed all entries
|
||||
if buckets16[RADIX_BASE - 1] != active_splats {
|
||||
return Err(anyhow!(
|
||||
"Expected {} active splats but got {}",
|
||||
active_splats,
|
||||
buckets16[RADIX_BASE - 1]
|
||||
));
|
||||
}
|
||||
|
||||
Ok(active_splats)
|
||||
}
|
||||
+60
-10
@@ -18,6 +18,7 @@ import {
|
||||
dyno,
|
||||
dynoBlock,
|
||||
dynoConst,
|
||||
floatBitsToUint,
|
||||
mul,
|
||||
packHalf2x16,
|
||||
readPackedSplat,
|
||||
@@ -117,6 +118,11 @@ export type SparkViewpointOptions = {
|
||||
* @default false
|
||||
*/
|
||||
sort360?: boolean;
|
||||
/*
|
||||
* Set this to true to sort with float32 precision with two-pass sort.
|
||||
* @default true
|
||||
*/
|
||||
sort32?: boolean;
|
||||
};
|
||||
|
||||
// A SparkViewpoint is created from and tied to a SparkRenderer, and represents
|
||||
@@ -149,6 +155,7 @@ export class SparkViewpoint {
|
||||
sortCoorient?: boolean;
|
||||
depthBias?: number;
|
||||
sort360?: boolean;
|
||||
sort32?: boolean;
|
||||
|
||||
display: {
|
||||
accumulator: SplatAccumulator;
|
||||
@@ -164,7 +171,8 @@ export class SparkViewpoint {
|
||||
} | null = null;
|
||||
private sortingCheck = false;
|
||||
|
||||
private readback: Uint16Array = new Uint16Array(0);
|
||||
private readback16: Uint16Array = new Uint16Array(0);
|
||||
private readback32: Uint32Array = new Uint32Array(0);
|
||||
private orderingFreelist: FreeList<Uint32Array, number>;
|
||||
|
||||
constructor(options: SparkViewpointOptions & { spark: SparkRenderer }) {
|
||||
@@ -209,6 +217,7 @@ export class SparkViewpoint {
|
||||
this.sortCoorient = options.sortCoorient;
|
||||
this.depthBias = options.depthBias;
|
||||
this.sort360 = options.sort360;
|
||||
this.sort32 = options.sort32;
|
||||
|
||||
this.orderingFreelist = new FreeList({
|
||||
allocate: (maxSplats) => new Uint32Array(maxSplats),
|
||||
@@ -557,6 +566,7 @@ export class SparkViewpoint {
|
||||
const {
|
||||
reader,
|
||||
doubleSortReader,
|
||||
sort32Reader,
|
||||
dynoSortRadial,
|
||||
dynoOrigin,
|
||||
dynoDirection,
|
||||
@@ -564,8 +574,16 @@ export class SparkViewpoint {
|
||||
dynoSort360,
|
||||
dynoSplats,
|
||||
} = SparkViewpoint.makeSorter();
|
||||
const halfMaxSplats = Math.ceil(maxSplats / 2);
|
||||
this.readback = reader.ensureBuffer(halfMaxSplats, this.readback);
|
||||
const sort32 = this.sort32 ?? false;
|
||||
let readback: Uint16Array | Uint32Array;
|
||||
if (sort32) {
|
||||
this.readback32 = reader.ensureBuffer(maxSplats, this.readback32);
|
||||
readback = this.readback32;
|
||||
} else {
|
||||
const halfMaxSplats = Math.ceil(maxSplats / 2);
|
||||
this.readback16 = reader.ensureBuffer(halfMaxSplats, this.readback16);
|
||||
readback = this.readback16;
|
||||
}
|
||||
|
||||
const worldToOrigin = accumulator.toWorld.clone().invert();
|
||||
const viewToOrigin = viewToWorld.clone().premultiply(worldToOrigin);
|
||||
@@ -581,25 +599,33 @@ export class SparkViewpoint {
|
||||
dynoSort360.value = this.sort360 ?? false;
|
||||
dynoSplats.packedSplats = accumulator.splats;
|
||||
|
||||
const sortReader = sort32 ? sort32Reader : doubleSortReader;
|
||||
const count = sort32 ? numSplats : Math.ceil(numSplats / 2);
|
||||
await reader.renderReadback({
|
||||
renderer: this.spark.renderer,
|
||||
reader: doubleSortReader,
|
||||
count: Math.ceil(numSplats / 2),
|
||||
readback: this.readback,
|
||||
reader: sortReader,
|
||||
count,
|
||||
readback,
|
||||
});
|
||||
|
||||
const result = (await withWorker(async (worker) => {
|
||||
return worker.call("sortDoubleSplats", {
|
||||
const rpcName = sort32 ? "sort32Splats" : "sortDoubleSplats";
|
||||
return worker.call(rpcName, {
|
||||
maxSplats,
|
||||
numSplats,
|
||||
readback: this.readback,
|
||||
readback,
|
||||
ordering,
|
||||
});
|
||||
})) as {
|
||||
readback: Uint16Array;
|
||||
readback: Uint16Array | Uint32Array;
|
||||
ordering: Uint32Array;
|
||||
activeSplats: number;
|
||||
};
|
||||
this.readback = result.readback;
|
||||
if (sort32) {
|
||||
this.readback32 = result.readback as Uint32Array;
|
||||
} else {
|
||||
this.readback16 = result.readback as Uint16Array;
|
||||
}
|
||||
ordering = result.ordering;
|
||||
activeSplats = result.activeSplats;
|
||||
}
|
||||
@@ -669,6 +695,7 @@ export class SparkViewpoint {
|
||||
dynoSplats: DynoPackedSplats;
|
||||
reader: Readback;
|
||||
doubleSortReader: DynoBlock<{ index: "int" }, { rgba8: "vec4" }>;
|
||||
sort32Reader: DynoBlock<{ index: "int" }, { rgba8: "vec4" }>;
|
||||
} | null = null;
|
||||
|
||||
private static makeSorter() {
|
||||
@@ -716,6 +743,28 @@ export class SparkViewpoint {
|
||||
},
|
||||
);
|
||||
|
||||
const sort32Reader = dynoBlock(
|
||||
{ index: "int" },
|
||||
{ rgba8: "vec4" },
|
||||
({ index }) => {
|
||||
if (!index) {
|
||||
throw new Error("No index");
|
||||
}
|
||||
const sortParams = {
|
||||
sortRadial: dynoSortRadial,
|
||||
sortOrigin: dynoOrigin,
|
||||
sortDirection: dynoDirection,
|
||||
sortDepthBias: dynoDepthBias,
|
||||
sort360: dynoSort360,
|
||||
};
|
||||
|
||||
const gsplat = readPackedSplat(dynoSplats, index);
|
||||
const metric = computeSortMetric({ gsplat, ...sortParams });
|
||||
const rgba8 = uintToRgba8(floatBitsToUint(metric));
|
||||
return { rgba8 };
|
||||
},
|
||||
);
|
||||
|
||||
SparkViewpoint.dynos = {
|
||||
dynoSortRadial,
|
||||
dynoOrigin,
|
||||
@@ -725,6 +774,7 @@ export class SparkViewpoint {
|
||||
dynoSplats,
|
||||
reader,
|
||||
doubleSortReader,
|
||||
sort32Reader,
|
||||
};
|
||||
}
|
||||
return SparkViewpoint.dynos;
|
||||
|
||||
+223
-35
@@ -1,4 +1,4 @@
|
||||
import init_wasm, { sort_splats } from "spark-internal-rs";
|
||||
import init_wasm, { sort_splats, sort32_splats } from "spark-internal-rs";
|
||||
import type { PcSogsJson, TranscodeSpzInput } from "./SplatLoader";
|
||||
import { unpackAntiSplat } from "./antisplat";
|
||||
import { WASM_SPLAT_SORT } from "./defines";
|
||||
@@ -18,6 +18,7 @@ import {
|
||||
setPackedSplatQuat,
|
||||
setPackedSplatRgb,
|
||||
setPackedSplatScales,
|
||||
toHalf,
|
||||
} from "./utils";
|
||||
|
||||
// WebWorker for Spark's background CPU tasks, such as Gsplat file decoding
|
||||
@@ -134,11 +135,6 @@ async function onMessage(event: MessageEvent) {
|
||||
readback: Uint16Array;
|
||||
ordering: Uint32Array;
|
||||
};
|
||||
result = {
|
||||
id,
|
||||
readback,
|
||||
ordering,
|
||||
};
|
||||
if (WASM_SPLAT_SORT) {
|
||||
result = {
|
||||
id,
|
||||
@@ -155,6 +151,31 @@ async function onMessage(event: MessageEvent) {
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "sort32Splats": {
|
||||
const { maxSplats, numSplats, readback, ordering } = args as {
|
||||
maxSplats: number;
|
||||
numSplats: number;
|
||||
readback: Uint32Array;
|
||||
ordering: Uint32Array;
|
||||
};
|
||||
// Benchmark sort
|
||||
// benchmarkSort(numSplats, readback, ordering);
|
||||
if (WASM_SPLAT_SORT) {
|
||||
result = {
|
||||
id,
|
||||
readback,
|
||||
ordering,
|
||||
activeSplats: sort32_splats(numSplats, readback, ordering),
|
||||
};
|
||||
} else {
|
||||
result = {
|
||||
id,
|
||||
readback,
|
||||
...sort32Splats({ maxSplats, numSplats, readback, ordering }),
|
||||
};
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "transcodeSpz": {
|
||||
const input = args as TranscodeSpzInput;
|
||||
const spzBytes = await transcodeSpz(input);
|
||||
@@ -171,6 +192,7 @@ async function onMessage(event: MessageEvent) {
|
||||
}
|
||||
} catch (e) {
|
||||
error = e;
|
||||
console.error(error);
|
||||
}
|
||||
|
||||
// Send the result or error back to the main thread, making sure to transfer any ArrayBuffers
|
||||
@@ -180,6 +202,82 @@ async function onMessage(event: MessageEvent) {
|
||||
);
|
||||
}
|
||||
|
||||
function benchmarkSort(
|
||||
numSplats: number,
|
||||
readback32: Uint32Array,
|
||||
ordering: Uint32Array,
|
||||
) {
|
||||
if (numSplats > 0) {
|
||||
console.log("Running sort benchmark");
|
||||
const readbackF32 = new Float32Array(readback32.buffer);
|
||||
const readback16 = new Uint16Array(readback32.length);
|
||||
for (let i = 0; i < numSplats; ++i) {
|
||||
readback16[i] = toHalf(readbackF32[i]);
|
||||
}
|
||||
|
||||
const WARMUP = 10;
|
||||
for (let i = 0; i < WARMUP; ++i) {
|
||||
const activeSplats = sort_splats(numSplats, readback16, ordering);
|
||||
const activeSplats32 = sort32_splats(numSplats, readback32, ordering);
|
||||
const results = sortDoubleSplats({
|
||||
numSplats,
|
||||
readback: readback16,
|
||||
ordering,
|
||||
});
|
||||
const results32 = sort32Splats({
|
||||
maxSplats: numSplats,
|
||||
numSplats,
|
||||
readback: readback32,
|
||||
ordering,
|
||||
});
|
||||
}
|
||||
|
||||
const TIMING_SAMPLES = 1000;
|
||||
let start: number;
|
||||
|
||||
start = performance.now();
|
||||
for (let i = 0; i < TIMING_SAMPLES; ++i) {
|
||||
const activeSplats = sort_splats(numSplats, readback16, ordering);
|
||||
}
|
||||
const wasmTime = (performance.now() - start) / TIMING_SAMPLES;
|
||||
|
||||
start = performance.now();
|
||||
for (let i = 0; i < TIMING_SAMPLES; ++i) {
|
||||
const results = sortDoubleSplats({
|
||||
numSplats,
|
||||
readback: readback16,
|
||||
ordering,
|
||||
});
|
||||
}
|
||||
const jsTime = (performance.now() - start) / TIMING_SAMPLES;
|
||||
|
||||
console.log(
|
||||
`JS: ${jsTime} ms, WASM: ${wasmTime} ms, numSplats: ${numSplats}`,
|
||||
);
|
||||
|
||||
start = performance.now();
|
||||
for (let i = 0; i < TIMING_SAMPLES; ++i) {
|
||||
const activeSplats32 = sort32_splats(numSplats, readback32, ordering);
|
||||
}
|
||||
const wasm32Time = (performance.now() - start) / TIMING_SAMPLES;
|
||||
|
||||
start = performance.now();
|
||||
for (let i = 0; i < TIMING_SAMPLES; ++i) {
|
||||
const results = sort32Splats({
|
||||
maxSplats: numSplats,
|
||||
numSplats,
|
||||
readback: readback32,
|
||||
ordering,
|
||||
});
|
||||
}
|
||||
const js32Time = (performance.now() - start) / TIMING_SAMPLES;
|
||||
|
||||
console.log(
|
||||
`JS32: ${js32Time} ms, WASM32: ${wasm32Time} ms, numSplats: ${numSplats}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async function unpackPly({
|
||||
packedArray,
|
||||
fileBytes,
|
||||
@@ -308,9 +406,9 @@ function unpackSpz(fileBytes: Uint8Array): {
|
||||
}
|
||||
|
||||
// Array of buckets for sorting float16 distances with range [0, DEPTH_INFINITY].
|
||||
const DEPTH_INFINITY = 0x7c00;
|
||||
const DEPTH_SIZE = DEPTH_INFINITY + 1;
|
||||
let depthArray: Uint32Array | null = null;
|
||||
const DEPTH_INFINITY_F16 = 0x7c00;
|
||||
const DEPTH_SIZE_16 = DEPTH_INFINITY_F16 + 1;
|
||||
let depthArray16: Uint32Array | null = null;
|
||||
|
||||
function sortSplats({
|
||||
totalSplats,
|
||||
@@ -323,10 +421,10 @@ function sortSplats({
|
||||
// Sort totalSplats Gsplats, each with 4 bytes of readback, and outputs Uint32Array
|
||||
// of indices from most distant to nearest. Each 4 bytes encode a float16 distance
|
||||
// and unused high bytes.
|
||||
if (!depthArray) {
|
||||
depthArray = new Uint32Array(DEPTH_SIZE);
|
||||
if (!depthArray16) {
|
||||
depthArray16 = new Uint32Array(DEPTH_SIZE_16);
|
||||
}
|
||||
depthArray.fill(0);
|
||||
depthArray16.fill(0);
|
||||
|
||||
const readbackUint32 = readback.map((layer) => new Uint32Array(layer.buffer));
|
||||
const layerSize = readbackUint32[0].length;
|
||||
@@ -338,17 +436,17 @@ function sortSplats({
|
||||
const layerSplats = Math.min(readbackLayer.length, totalSplats - layerBase);
|
||||
for (let i = 0; i < layerSplats; ++i) {
|
||||
const pri = readbackLayer[i] & 0x7fff;
|
||||
if (pri < DEPTH_INFINITY) {
|
||||
depthArray[pri] += 1;
|
||||
if (pri < DEPTH_INFINITY_F16) {
|
||||
depthArray16[pri] += 1;
|
||||
}
|
||||
}
|
||||
layerBase += layerSplats;
|
||||
}
|
||||
|
||||
let activeSplats = 0;
|
||||
for (let j = 0; j < DEPTH_SIZE; ++j) {
|
||||
const nextIndex = activeSplats + depthArray[j];
|
||||
depthArray[j] = activeSplats;
|
||||
for (let j = 0; j < DEPTH_SIZE_16; ++j) {
|
||||
const nextIndex = activeSplats + depthArray16[j];
|
||||
depthArray16[j] = activeSplats;
|
||||
activeSplats = nextIndex;
|
||||
}
|
||||
|
||||
@@ -358,16 +456,16 @@ function sortSplats({
|
||||
const layerSplats = Math.min(readbackLayer.length, totalSplats - layerBase);
|
||||
for (let i = 0; i < layerSplats; ++i) {
|
||||
const pri = readbackLayer[i] & 0x7fff;
|
||||
if (pri < DEPTH_INFINITY) {
|
||||
ordering[depthArray[pri]] = layerBase + i;
|
||||
depthArray[pri] += 1;
|
||||
if (pri < DEPTH_INFINITY_F16) {
|
||||
ordering[depthArray16[pri]] = layerBase + i;
|
||||
depthArray16[pri] += 1;
|
||||
}
|
||||
}
|
||||
layerBase += layerSplats;
|
||||
}
|
||||
if (depthArray[DEPTH_SIZE - 1] !== activeSplats) {
|
||||
if (depthArray16[DEPTH_SIZE_16 - 1] !== activeSplats) {
|
||||
throw new Error(
|
||||
`Expected ${activeSplats} active splats but got ${depthArray[DEPTH_SIZE - 1]}`,
|
||||
`Expected ${activeSplats} active splats but got ${depthArray16[DEPTH_SIZE_16 - 1]}`,
|
||||
);
|
||||
}
|
||||
|
||||
@@ -385,16 +483,16 @@ function sortDoubleSplats({
|
||||
ordering: Uint32Array;
|
||||
} {
|
||||
// Ensure depthArray is allocated and zeroed out for our buckets.
|
||||
if (!depthArray) {
|
||||
depthArray = new Uint32Array(DEPTH_SIZE);
|
||||
if (!depthArray16) {
|
||||
depthArray16 = new Uint32Array(DEPTH_SIZE_16);
|
||||
}
|
||||
depthArray.fill(0);
|
||||
depthArray16.fill(0);
|
||||
|
||||
// Count the number of splats in each bucket (cull Gsplats at infinity).
|
||||
for (let i = 0; i < numSplats; ++i) {
|
||||
const pri = readback[i];
|
||||
if (pri < DEPTH_INFINITY) {
|
||||
depthArray[pri] += 1;
|
||||
if (pri < DEPTH_INFINITY_F16) {
|
||||
depthArray16[pri] += 1;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -402,9 +500,9 @@ function sortDoubleSplats({
|
||||
// total number of active (non-infinity) splats, going in reverse order
|
||||
// because we want most distant Gsplats to be first in the output array.
|
||||
let activeSplats = 0;
|
||||
for (let j = DEPTH_INFINITY - 1; j >= 0; --j) {
|
||||
const nextIndex = activeSplats + depthArray[j];
|
||||
depthArray[j] = activeSplats;
|
||||
for (let j = DEPTH_INFINITY_F16 - 1; j >= 0; --j) {
|
||||
const nextIndex = activeSplats + depthArray16[j];
|
||||
depthArray16[j] = activeSplats;
|
||||
activeSplats = nextIndex;
|
||||
}
|
||||
|
||||
@@ -412,16 +510,106 @@ function sortDoubleSplats({
|
||||
// bucket order.
|
||||
for (let i = 0; i < numSplats; ++i) {
|
||||
const pri = readback[i];
|
||||
if (pri < DEPTH_INFINITY) {
|
||||
ordering[depthArray[pri]] = i;
|
||||
depthArray[pri] += 1;
|
||||
if (pri < DEPTH_INFINITY_F16) {
|
||||
ordering[depthArray16[pri]] = i;
|
||||
depthArray16[pri] += 1;
|
||||
}
|
||||
}
|
||||
// Sanity check that the end of the closest bucket is the same as
|
||||
// our total count of active splats (not at infinity).
|
||||
if (depthArray[0] !== activeSplats) {
|
||||
if (depthArray16[0] !== activeSplats) {
|
||||
throw new Error(
|
||||
`Expected ${activeSplats} active splats but got ${depthArray[0]}`,
|
||||
`Expected ${activeSplats} active splats but got ${depthArray16[0]}`,
|
||||
);
|
||||
}
|
||||
|
||||
return { activeSplats, ordering };
|
||||
}
|
||||
|
||||
const DEPTH_INFINITY_F32 = 0x7f800000;
|
||||
let bucket16: Uint32Array | null = null;
|
||||
let scratchSplats: Uint32Array | null = null;
|
||||
|
||||
// two-pass radix sort (base 65536) of 32-bit keys in readback,
|
||||
// but placing largest values first.
|
||||
function sort32Splats({
|
||||
maxSplats,
|
||||
numSplats,
|
||||
readback, // Uint32Array of bit‑patterns
|
||||
ordering, // Uint32Array to fill with sorted indices
|
||||
}: {
|
||||
maxSplats: number;
|
||||
numSplats: number;
|
||||
readback: Uint32Array;
|
||||
ordering: Uint32Array;
|
||||
}): { activeSplats: number; ordering: Uint32Array } {
|
||||
const BASE = 1 << 16; // 65536
|
||||
|
||||
// allocate once
|
||||
if (!bucket16) {
|
||||
bucket16 = new Uint32Array(BASE);
|
||||
}
|
||||
if (!scratchSplats || scratchSplats.length < maxSplats) {
|
||||
scratchSplats = new Uint32Array(maxSplats);
|
||||
}
|
||||
|
||||
//
|
||||
// ——— Pass #1: bucket by inv(lo 16 bits) ———
|
||||
//
|
||||
bucket16.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;
|
||||
}
|
||||
}
|
||||
// exclusive prefix‑sum → starting offsets
|
||||
let total = 0;
|
||||
for (let b = 0; b < BASE; ++b) {
|
||||
const c = bucket16[b];
|
||||
bucket16[b] = total;
|
||||
total += c;
|
||||
}
|
||||
const activeSplats = total;
|
||||
|
||||
// scatter into scratch by low bits of inv
|
||||
for (let i = 0; i < numSplats; ++i) {
|
||||
const key = readback[i];
|
||||
if (key < DEPTH_INFINITY_F32) {
|
||||
const inv = ~key >>> 0;
|
||||
scratchSplats[bucket16[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;
|
||||
sum += c;
|
||||
}
|
||||
|
||||
// scatter into final ordering by high bits of inv
|
||||
for (let k = 0; k < activeSplats; ++k) {
|
||||
const idx = scratchSplats[k];
|
||||
const inv = ~readback[idx] >>> 0;
|
||||
ordering[bucket16[inv >>> 16]++] = idx;
|
||||
}
|
||||
|
||||
// sanity‑check: the last bucket should have eaten all entries
|
||||
if (bucket16[BASE - 1] !== activeSplats) {
|
||||
throw new Error(
|
||||
`Expected ${activeSplats} active splats but got ${bucket16[BASE - 1]}`,
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user