Initial commit

This commit is contained in:
Diego Marcos
2025-05-23 08:52:14 -07:00
committed by Diego Marcos Segura
commit 3332f609a9
129 changed files with 64575 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
target/
forge-internal-rs/pkg/
+224
View File
@@ -0,0 +1,224 @@
# This file is automatically @generated by Cargo.
# It is not intended for manual editing.
version = 3
[[package]]
name = "anyhow"
version = "1.0.98"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e16d2d3311acee920a9eb8d33b8cbc1787ce4a264e85f964c2404b969bdcd487"
[[package]]
name = "bincode"
version = "1.3.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad"
dependencies = [
"serde",
]
[[package]]
name = "bumpalo"
version = "3.17.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1628fb46dfa0b37568d12e5edd512553eccf6a22a78e8bde00bb4aed84d5bdbf"
[[package]]
name = "cfg-if"
version = "1.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "baf1de4339761588bc0619e3cbc0120ee582ebb74b53b4efbf79117bd2da40fd"
[[package]]
name = "crunchy"
version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "43da5946c66ffcc7745f48db692ffbb10a83bfe0afd96235c5c2a4fb23994929"
[[package]]
name = "forge-internal-rs"
version = "0.1.0"
dependencies = [
"anyhow",
"half",
"js-sys",
"wasm-bindgen",
"wlg",
]
[[package]]
name = "half"
version = "2.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "459196ed295495a68f7d7fe1d84f6c4b7ff0e21fe3017b2f283c6fac3ad803c9"
dependencies = [
"cfg-if",
"crunchy",
]
[[package]]
name = "js-sys"
version = "0.3.77"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1cfaf33c695fc6e08064efbc1f72ec937429614f25eef83af942d0e227c3a28f"
dependencies = [
"once_cell",
"wasm-bindgen",
]
[[package]]
name = "log"
version = "0.4.27"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "13dc2df351e3202783a1fe0d44375f7295ffb4049267b0f3018346dc122a1d94"
[[package]]
name = "once_cell"
version = "1.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d"
[[package]]
name = "proc-macro2"
version = "1.0.95"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "02b3e5e68a3a1a02aad3ec490a98007cbc13c37cbe84a3cd7b8e406d76e7f778"
dependencies = [
"unicode-ident",
]
[[package]]
name = "quote"
version = "1.0.40"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1885c039570dc00dcb4ff087a89e185fd56bae234ddc7f056a945bf36467248d"
dependencies = [
"proc-macro2",
]
[[package]]
name = "rustversion"
version = "1.0.20"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "eded382c5f5f786b989652c49544c4877d9f015cc22e145a5ea8ea66c2921cd2"
[[package]]
name = "ruzstd"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c581601827da5c717bfae77d7b187e54293d23d8fb6b700b4b5e9b5828a13cc3"
dependencies = [
"twox-hash",
]
[[package]]
name = "serde"
version = "1.0.219"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5f0e2c6ed6606019b4e29e69dbaba95b11854410e5347d525002456dbbb786b6"
dependencies = [
"serde_derive",
]
[[package]]
name = "serde_derive"
version = "1.0.219"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5b0276cf7f2c73365f7157c8123c21cd9a50fbbd844757af28ca1f5925fc2a00"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "syn"
version = "2.0.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b09a44accad81e1ba1cd74a32461ba89dee89095ba17b32f5d03683b1b1fc2a0"
dependencies = [
"proc-macro2",
"quote",
"unicode-ident",
]
[[package]]
name = "twox-hash"
version = "2.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e7b17f197b3050ba473acf9181f7b1d3b66d1cf7356c6cc57886662276e65908"
[[package]]
name = "unicode-ident"
version = "1.0.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5a5f39404a5da50712a4c1eecf25e90dd62b613502b7e925fd4e4d19b5c96512"
[[package]]
name = "wasm-bindgen"
version = "0.2.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1edc8929d7499fc4e8f0be2262a241556cfc54a0bea223790e71446f2aab1ef5"
dependencies = [
"cfg-if",
"once_cell",
"rustversion",
"wasm-bindgen-macro",
]
[[package]]
name = "wasm-bindgen-backend"
version = "0.2.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2f0a0651a5c2bc21487bde11ee802ccaf4c51935d0d3d42a6101f98161700bc6"
dependencies = [
"bumpalo",
"log",
"proc-macro2",
"quote",
"syn",
"wasm-bindgen-shared",
]
[[package]]
name = "wasm-bindgen-macro"
version = "0.2.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7fe63fc6d09ed3792bd0897b314f53de8e16568c2b3f7982f468c0bf9bd0b407"
dependencies = [
"quote",
"wasm-bindgen-macro-support",
]
[[package]]
name = "wasm-bindgen-macro-support"
version = "0.2.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8ae87ea40c9f689fc23f209965b6fb8a99ad69aeeb0231408be24920604395de"
dependencies = [
"proc-macro2",
"quote",
"syn",
"wasm-bindgen-backend",
"wasm-bindgen-shared",
]
[[package]]
name = "wasm-bindgen-shared"
version = "0.2.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1a05d73b933a847d6cccdda8f838a22ff101ad9bf93e33684f39c1f5f0eece3d"
dependencies = [
"unicode-ident",
]
[[package]]
name = "wlg"
version = "0.1.0"
dependencies = [
"anyhow",
"bincode",
"half",
"ruzstd",
"serde",
]
+22
View File
@@ -0,0 +1,22 @@
[workspace]
members = [
"forge-internal-rs",
"wlg",
]
resolver = "2"
[workspace.package]
rust-version = "1.82"
edition = "2021"
license = "Proprietary"
authors = ["World Labs Technologies"]
repository = "https://github.com/forge-gfx/forge"
[workspace.dependencies]
anyhow = "1.0.98"
bincode = "1.3.3"
half = "2.6.0"
js-sys = "0.3.77"
ruzstd = "0.8.0"
serde = { version = "1.0.219", features = ["derive"] }
wasm-bindgen = "0.2.100"
+24
View File
@@ -0,0 +1,24 @@
# Exit on any error
$ErrorActionPreference = "Stop"
# Resolve the script's directory and change to it
Set-Location -Path (Split-Path -Parent $MyInvocation.MyCommand.Path)
# Check if rustup is installed
if (-not (Get-Command rustup -ErrorAction SilentlyContinue)) {
Write-Host "Rust tool 'rustup' not found! Please install Rust to build."
Write-Host "Visit: https://www.rust-lang.org/tools/install"
exit 1
}
# Ensure wasm32-unknown-unknown target is installed
rustup target add wasm32-unknown-unknown
# Check if wasm-pack is installed, and install if not
if (-not (Get-Command wasm-pack -ErrorAction SilentlyContinue)) {
cargo install wasm-pack
}
# Change directory and build using wasm-pack
Set-Location -Path "./forge-internal-rs"
wasm-pack build --target web
+22
View File
@@ -0,0 +1,22 @@
#!/bin/bash
# Check if "rustup" tool is installed
if ! command -v rustup &> /dev/null; then
echo "Rust tool 'rustup' not found! Please install Rust to build."
echo "Visit Rust installation page: https://www.rust-lang.org/tools/install"
echo "- Likely install one-liner: curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh"
exit 1
fi
cd $(dirname "$0")
# Make sure Rust wasm target is installed
rustup target add wasm32-unknown-unknown
# Make sure wasm-pack is installed
cargo install wasm-pack
cd forge-internal-rs
# Build the project
wasm-pack build --target web
+18
View File
@@ -0,0 +1,18 @@
const { execSync } = await import("node:child_process");
const { platform } = await import("node:os");
const isWindows = platform() === "win32";
try {
if (isWindows) {
const output = execSync(
"powershell.exe -ExecutionPolicy Bypass -File ./rust/build_rust_wasm.ps1",
{ stdio: "inherit" },
);
} else {
execSync("rust/build_rust_wasm.sh", { stdio: "inherit" });
}
} catch (err) {
console.error("Failed to build RUST WASM:", err.message);
process.exit(1);
}
+18
View File
@@ -0,0 +1,18 @@
[package]
name = "forge-internal-rs"
version = "0.1.0"
rust-version.workspace = true
edition.workspace = true
license.workspace = true
authors.workspace = true
repository.workspace = true
[lib]
crate-type = ["cdylib"]
[dependencies]
anyhow.workspace = true
half.workspace = true
js-sys.workspace = true
wasm-bindgen.workspace = true
wlg = { path = "../wlg" }
+35
View File
@@ -0,0 +1,35 @@
# forge-internal-rs
Rust WebAssembly functions for forge-internal.
## Installing build tools
First, we need to install Rust. Though it is possible to install it using Homebrew, we recommend installing `rustup` using the approach on the Rust homepage:
https://www.rust-lang.org/tools/install
It will most likely involve simply running:
```
curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh
```
Once you have `rustup` and Rust, we need to install dependencies for building Rust Wasm.
```
rustup target add wasm32-unknown-unknown
cargo install wasm-pack
```
## Building
Run the following script inside `forge-internal/rust`:
```
./build_rust_wasm.sh
```
You can also build it manually by running these commands:
```
cd forge-internal-rs
wasm-pack build --target web
```
The generated files will be in the `pkg/` subdirectory, which is already symlinked in `forge-internal/package.json`.
+154
View File
@@ -0,0 +1,154 @@
use std::cell::RefCell;
use js_sys::{Array, ArrayBuffer, Float32Array, Object, Reflect, Uint16Array, Uint32Array, Uint8Array};
use wasm_bindgen::prelude::*;
use wlg::decode12_packed;
mod sort;
use sort::{old_sort_internal, OldSortBuffers, sort_internal, SortBuffers};
mod raycast;
use raycast::{raycast_ellipsoids, raycast_spheres};
const RAYCAST_BUFFER_COUNT: u32 = 65536;
thread_local! {
static SORT_BUFFERS: RefCell<SortBuffers> = RefCell::new(SortBuffers::default());
static OLD_SORT_BUFFERS: RefCell<OldSortBuffers> = RefCell::new(OldSortBuffers::default());
static RAYCAST_BUFFER: RefCell<Vec<u32>> = RefCell::new(vec![0; RAYCAST_BUFFER_COUNT as usize * 4]);
}
#[wasm_bindgen]
pub fn decode_wlg(bytes: ArrayBuffer, tex_width: u32, tex_height: u32) -> Object {
let mut bytes = Uint8Array::new(&bytes).to_vec();
let (_settings, packed_splats) = match decode12_packed(&mut bytes) {
Ok(result) => result,
Err(err) => {
wasm_bindgen::throw_str(&format!("{}", err));
}
};
let num_splats = packed_splats.0.len() / 4;
let max_splats = if tex_width != 0 && tex_height != 0 {
let width = tex_width as usize;
let height = num_splats.div_ceil(width).min(tex_height as usize);
let depth = num_splats.div_ceil(width * height);
width * height * depth
} else {
num_splats
};
let packed = Uint32Array::new_with_length(max_splats as u32 * 4);
let packed_slice = packed.subarray(0, packed_splats.0.len() as u32);
packed_slice.copy_from(&packed_splats.0);
let result = Object::new();
Reflect::set(&result, &JsValue::from_str("numSplats"), &JsValue::from_f64(num_splats as f64)).unwrap();
Reflect::set(&result, &JsValue::from_str("packedSplats"), &JsValue::from(packed)).unwrap();
result
}
#[wasm_bindgen]
pub fn old_sort_splats(
max_splats: u32, total_splats: u32, readback: Array, ordering: Uint32Array,
) -> u32 {
let max_splats = max_splats as usize;
let total_splats = total_splats as usize;
let num_layers = readback.length() as usize;
let layer_size = readback.get(0).dyn_into::<Uint8Array>().unwrap().length() as usize / 4;
let active_splats = OLD_SORT_BUFFERS.with_borrow_mut(|buffers| {
buffers.ensure_size(max_splats);
// Copy the readback data layers into a contiguous buffer
let mut layer_base = 0;
for layer in 0..num_layers {
let layer_count = layer_size.min(total_splats - layer_base);
if layer_count > 0 {
let layer_buffer = readback.get(layer as u32).dyn_into::<Uint8Array>().unwrap().buffer();
let layer_uint32 = Uint32Array::new_with_byte_offset_and_length(&layer_buffer, 0, layer_count as u32);
layer_uint32.copy_to(&mut buffers.readback[layer_base..layer_base + layer_count]);
}
layer_base += layer_count;
}
let active_splats = match old_sort_internal(buffers, total_splats) {
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 sort_splats(
num_splats: u32, readback: Uint16Array, ordering: Uint32Array,
) -> u32 {
let max_splats = readback.length() as usize;
let active_splats = SORT_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 sort_internal(buffers, 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,
dir_x: f32, dir_y: f32, dir_z: f32,
near: f32, far: f32,
num_splats: u32, packed_splats: Uint32Array,
raycast_ellipsoid: bool,
) -> Float32Array {
let mut distances = Vec::<f32>::new();
_ = RAYCAST_BUFFER.with_borrow_mut(|buffer| {
let mut base = 0;
while base < num_splats {
let chunk_size = RAYCAST_BUFFER_COUNT.min(num_splats - base);
let subarray = packed_splats.subarray(4 * base, 4 * (base + chunk_size));
let subbuffer = &mut buffer[0..(4 * chunk_size as usize)];
subarray.copy_to(subbuffer);
if raycast_ellipsoid {
raycast_ellipsoids(subbuffer, &mut distances, [origin_x, origin_y, origin_z], [dir_x, dir_y, dir_z], near, far);
} else {
raycast_spheres(subbuffer, &mut distances, [origin_x, origin_y, origin_z], [dir_x, dir_y, dir_z], near, far);
}
base += chunk_size;
}
});
let output = Float32Array::new_with_length(distances.len() as u32);
output.copy_from(&distances);
output
}
+170
View File
@@ -0,0 +1,170 @@
use half::f16;
use wlg::decode_scale;
const MIN_OPACITY: f32 = 0.1;
pub fn raycast_spheres(
buffer: &[u32], distances: &mut Vec<f32>,
origin: [f32; 3], dir: [f32; 3], near: f32, far: f32,
) {
let quad_a = vec3_dot(dir, dir);
for packed in buffer.chunks(4) {
let opacity = ((packed[0] >> 24) as u8) as f32 / 255.0;
if opacity < MIN_OPACITY {
continue;
}
let origin = vec3_sub(origin, extract_center(packed));
let scale = extract_scale(packed);
// Model the Gsplat as a sphere for faster approximate raycasting
let radius = (scale[0] + scale[1] + scale[2]) / 3.0;
let quad_b = vec3_dot(dir, origin);
let quad_c = vec3_dot(origin, origin) - radius * radius;
let discriminant = quad_b * quad_b - quad_a * quad_c;
if discriminant < 0.0 {
continue;
}
let t = (-quad_b - discriminant.sqrt()) / quad_a;
if t >= near && t <= far {
distances.push(t);
}
}
}
pub fn raycast_ellipsoids(
buffer: &[u32], distances: &mut Vec<f32>,
origin: [f32; 3], dir: [f32; 3], near: f32, far: f32,
) {
for packed in buffer.chunks(4) {
let opacity = ((packed[0] >> 24) as u8) as f32 / 255.0;
if opacity < MIN_OPACITY {
continue;
}
let origin = vec3_sub(origin, extract_center(packed));
let scale = extract_scale(packed);
let quat = extract_quat(packed);
let inv_quat = [-quat[0], -quat[1], -quat[2], quat[3]];
// Model the Gsplat as an ellipsoid for higher quality raycasting
let local_origin = quat_vec(inv_quat, origin);
let local_dir = quat_vec(inv_quat, dir);
let min_scale = scale[0].max(scale[1]).max(scale[2]) * 0.01;
let t = if scale[2] < min_scale {
// Treat it as a flat elliptical disk
if local_dir[2].abs() < 1e-6 {
continue;
}
let t = -local_origin[2] / local_dir[2];
let p_x = local_origin[0] + t * local_dir[0];
let p_y = local_origin[1] + t * local_dir[1];
if sqr(p_x / scale[0]) + sqr(p_y / scale[1]) > 1.0 {
continue;
}
t
} else if scale[1] < min_scale {
// Treat it as a flat elliptical disk
if local_dir[1].abs() < 1e-6 {
continue;
}
let t = -local_origin[1] / local_dir[1];
let p_x = local_origin[0] + t * local_dir[0];
let p_z = local_origin[2] + t * local_dir[2];
if sqr(p_x / scale[0]) + sqr(p_z / scale[2]) > 1.0 {
continue;
}
t
} else if scale[0] < min_scale {
// Treat it as a flat elliptical disk
if local_dir[0].abs() < 1e-6 {
continue;
}
let t = -local_origin[0] / local_dir[0];
let p_y = local_origin[1] + t * local_dir[1];
let p_z = local_origin[2] + t * local_dir[2];
if sqr(p_y / scale[1]) + sqr(p_z / scale[2]) > 1.0 {
continue;
}
t
} else {
let inv_scale = [1.0 / scale[0], 1.0 / scale[1], 1.0 / scale[2]];
let local_origin = vec3_mul(local_origin, inv_scale);
let local_dir = vec3_mul(local_dir, inv_scale);
let a = vec3_dot(local_dir, local_dir);
let b = vec3_dot(local_origin, local_dir);
let c = vec3_dot(local_origin, local_origin) - 1.0;
let discriminant = b * b - a * c;
if discriminant < 0.0 {
continue;
}
(-b - discriminant.sqrt()) / a
};
if t >= near && t <= far {
distances.push(t);
}
}
}
fn extract_center(packed: &[u32]) -> [f32; 3] {
let x = f16::from_bits(packed[1] as u16).to_f32();
let y = f16::from_bits((packed[1] >> 16) as u16).to_f32();
let z = f16::from_bits(packed[2] as u16).to_f32();
[x, y, z]
}
fn extract_scale(packed: &[u32]) -> [f32; 3] {
let scale_x = decode_scale(packed[3] as u8);
let scale_y = decode_scale((packed[3] >> 8) as u8);
let scale_z = decode_scale((packed[3] >> 16) as u8);
[scale_x, scale_y, scale_z]
}
fn extract_quat(packed: &[u32]) -> [f32; 4] {
let quat_x = ((packed[2] >> 16) as i8) as f32 / 127.0;
let quat_y = ((packed[2] >> 24) as i8) as f32 / 127.0;
let quat_z = ((packed[3] >> 24) as i8) as f32 / 127.0;
let quat_w = (1.0 - quat_x * quat_x - quat_y * quat_y - quat_z * quat_z).max(0.0).sqrt();
[quat_x, quat_y, quat_z, quat_w]
}
fn sqr(x: f32) -> f32 {
x * x
}
fn vec3_sub(a: [f32; 3], b: [f32; 3]) -> [f32; 3] {
[a[0] - b[0], a[1] - b[1], a[2] - b[2]]
}
fn vec3_mul(a: [f32; 3], b: [f32; 3]) -> [f32; 3] {
[a[0] * b[0], a[1] * b[1], a[2] * b[2]]
}
fn vec3_dot(a: [f32; 3], b: [f32; 3]) -> f32 {
a[0] * b[0] + a[1] * b[1] + a[2] * b[2]
}
fn vec3_cross(a: [f32; 3], b: [f32; 3]) -> [f32; 3] {
[
a[1] * b[2] - a[2] * b[1],
a[2] * b[0] - a[0] * b[2],
a[0] * b[1] - a[1] * b[0],
]
}
fn quat_vec(q: [f32; 4], v: [f32; 3]) -> [f32; 3] {
let q_vec = [q[0], q[1], q[2]];
let uv = vec3_cross(q_vec, v);
let uuv = vec3_cross(q_vec, uv);
[
v[0] + 2.0 * (q[3] * uv[0] + uuv[0]),
v[1] + 2.0 * (q[3] * uv[1] + uuv[1]),
v[2] + 2.0 * (q[3] * uv[2] + uuv[2]),
]
}
+133
View File
@@ -0,0 +1,133 @@
use anyhow::anyhow;
const DEPTH_INFINITY: u32 = 0x7c00;
const DEPTH_SIZE: usize = DEPTH_INFINITY as usize + 1;
#[derive(Default)]
pub struct OldSortBuffers {
pub readback: Vec<u32>,
pub ordering: Vec<u32>,
pub buckets: Vec<u32>,
}
impl OldSortBuffers {
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.buckets.len() < DEPTH_SIZE {
self.buckets.resize(DEPTH_SIZE, 0);
}
}
}
pub fn old_sort_internal(buffers: &mut OldSortBuffers, total_splats: usize) -> anyhow::Result<u32> {
let OldSortBuffers { readback, ordering, buckets } = buffers;
let readback = &readback[..total_splats];
// Set the bucket counts to zero
buckets.clear();
buckets.resize(DEPTH_SIZE, 0);
// Count the number of splats in each bucket
for &readback_u32 in readback.iter() {
let priority = readback_u32 & 0x7FFF;
if priority < DEPTH_INFINITY {
buckets[priority as usize] += 1;
}
}
// Compute bucket starting offset
let mut active_splats = 0;
for count in buckets.iter_mut() {
let new_total = active_splats + *count;
*count = active_splats;
active_splats = new_total;
}
// Write out splat indices at the right location using bucket offsets
for (index, &readback_u32) in readback.iter().enumerate() {
let priority = readback_u32 & 0x7FFF;
if priority < DEPTH_INFINITY {
ordering[buckets[priority as usize] as usize] = index as u32;
buckets[priority as usize] += 1;
}
}
// Sanity check
if buckets[DEPTH_SIZE - 1] != active_splats {
return Err(anyhow!(
"Expected {} active splats but got {}",
active_splats,
buckets[DEPTH_SIZE - 1]
));
}
Ok(active_splats)
}
#[derive(Default)]
pub struct SortBuffers {
pub readback: Vec<u16>,
pub ordering: Vec<u32>,
pub buckets: Vec<u32>,
}
impl SortBuffers {
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.buckets.len() < DEPTH_SIZE {
self.buckets.resize(DEPTH_SIZE, 0);
}
}
}
pub fn sort_internal(buffers: &mut SortBuffers, num_splats: usize) -> anyhow::Result<u32> {
let SortBuffers { readback, ordering, buckets } = buffers;
let readback = &readback[..num_splats];
// Set the bucket counts to zero
buckets.clear();
buckets.resize(DEPTH_SIZE, 0);
// Count the number of splats in each bucket
for &metric in readback.iter() {
if (metric as u32) < DEPTH_INFINITY {
buckets[metric as usize] += 1;
}
}
// Compute bucket starting offset
let mut active_splats = 0;
for count in buckets.iter_mut().rev().skip(1) {
let new_total = active_splats + *count;
*count = active_splats;
active_splats = new_total;
}
// Write out splat indices at the right location using bucket offsets
for (index, &metric) in readback.iter().enumerate() {
if (metric as u32) < DEPTH_INFINITY {
ordering[buckets[metric as usize] as usize] = index as u32;
buckets[metric as usize] += 1;
}
}
// Sanity check
if buckets[0] != active_splats {
return Err(anyhow!(
"Expected {} active splats but got {}",
active_splats,
buckets[0]
));
}
Ok(active_splats)
}
+15
View File
@@ -0,0 +1,15 @@
[package]
name = "wlg"
version = "0.1.0"
rust-version.workspace = true
edition.workspace = true
license.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
anyhow.workspace = true
bincode.workspace = true
half.workspace = true
serde.workspace = true
ruzstd.workspace = true
+15
View File
@@ -0,0 +1,15 @@
mod ordering;
// WLG0 includes WLGv1 and WLGv2
mod wlg0;
pub use wlg0::{
Wlg0Gaussian, PackedSplats,
Wlg0Settings, Wlg0EncodeSettings,
decode12, decode12_packed,
encode12, encode_scale, decode_scale,
};
// WLG3 will be used for future WLG versions
// mod wlg3;
+201
View File
@@ -0,0 +1,201 @@
// Compute Morton index for 3D coordinates.
pub fn morton_coord_to_index([x, y, z]: [u16; 3]) -> u64 {
fn expand3(x: u16) -> u64 {
let mut x = x as u64;
x = (x | x << 32) & 0x1f00000000ffff;
x = (x | x << 16) & 0x1f0000ff0000ff;
x = (x | x << 8) & 0x100f00f00f00f00f;
x = (x | x << 4) & 0x10c30c30c30c30c3;
x = (x | x << 2) & 0x1249249249249249;
x
}
(expand3(x) << 0) | (expand3(y) << 1) | (expand3(z) << 2)
}
pub fn morton_coord_to_index_24([x, y, z]: [u32; 3]) -> u128 {
fn expand3_24(x: u32) -> u128 {
let mut x = x as u128;
x = (x | x << 64) & 0x3ff0000000000000000ffffffffu128;
x = (x | x << 32) & 0x3ff00000000ffff00000000ffffu128;
x = (x | x << 16) & 0x30000ff0000ff0000ff0000ff0000ffu128;
x = (x | x << 8) & 0x300f00f00f00f00f00f00f00f00f00fu128;
x = (x | x << 4) & 0x30c30c30c30c30c30c30c30c30c30c3u128;
x = (x | x << 2) & 0x9249249249249249249249249249249u128;
x
}
(expand3_24(x) << 0) | (expand3_24(y) << 1) | (expand3_24(z) << 2)
}
// Compute Hilbert index for 3D coordinates.
// Converts a 48-bit Hilbert index to 3D coordinates (x, y, z).
pub fn _hilbert_index_to_coord(mut index: u64) -> (u16, u16, u16) {
const BITS: u32 = 16;
// Helper function to decode a Hilbert quad to coordinate bits.
fn hilbert_decode_step(quad: u16, s: &mut [u16; 3]) {
let (mut x, mut y, mut z) = (0u16, 0u16, 0u16);
match quad {
0 => { x = 0; y = 0; z = 0; },
1 => { x = 0; y = 0; z = 1; },
2 => { x = 0; y = 1; z = 1; },
3 => { x = 0; y = 1; z = 0; },
4 => { x = 1; y = 1; z = 0; },
5 => { x = 1; y = 1; z = 1; },
6 => { x = 1; y = 0; z = 1; },
7 => { x = 1; y = 0; z = 0; },
_ => {},
}
s[0] = x;
s[1] = y;
s[2] = z;
}
let mut mask = 1u16; // Start with the least significant bit
let mut h = [0u16; 3];
let mut s = [0u16; 3];
for _ in 0..BITS {
let quad = (index & 7) as u16; // Extract the last 3 bits
index >>= 3;
hilbert_decode_step(quad, &mut s);
h[0] |= s[0] * mask;
h[1] |= s[1] * mask;
h[2] |= s[2] * mask;
mask <<= 1; // Move to the next bit
}
(h[0], h[1], h[2])
}
// Converts 3D coordinates (x, y, z) to a 48-bit Hilbert index.
pub fn hilbert_coord_to_index([x, y, z]: [u16; 3]) -> u64 {
const BITS: u32 = 16;
// Helper function to encode coordinate bits to a Hilbert quad.
fn hilbert_encode_step(s: &mut [u16; 3]) -> u16 {
let x = s[0];
let y = s[1];
let z = s[2];
let mut quad = 0u16;
if x == 0 && y == 0 && z == 0 { quad = 0; }
else if x == 0 && y == 0 && z == 1 { quad = 1; }
else if x == 0 && y == 1 && z == 1 { quad = 2; }
else if x == 0 && y == 1 && z == 0 { quad = 3; }
else if x == 1 && y == 1 && z == 0 { quad = 4; }
else if x == 1 && y == 1 && z == 1 { quad = 5; }
else if x == 1 && y == 0 && z == 1 { quad = 6; }
else if x == 1 && y == 0 && z == 0 { quad = 7; }
quad
}
let mut index = 0u64;
let mut h = [x, y, z];
let mut s = [0u16; 3];
for _ in 0..BITS {
s[0] = h[0] & 1;
s[1] = h[1] & 1;
s[2] = h[2] & 1;
let quad = hilbert_encode_step(&mut s);
index <<= 3;
index |= quad as u64;
h[0] >>= 1;
h[1] >>= 1;
h[2] >>= 1;
}
index
}
// Converts a 72-bit Hilbert index to 3D coordinates (x, y, z).
pub fn _hilbert_index_coord_24(index: u128) -> (u32, u32, u32) {
const BITS: u32 = 24;
// Helper function to decode a Hilbert quad to coordinates.
fn hilbert_decode_step(quad: u32, s: &mut [u32; 3]) {
match quad {
0 => { s.swap(0, 1); },
1 => {},
2 => {},
3 => { s[2] ^= 1; },
4 => { s[0] ^= 1; s[1] ^= 1; s[2] ^= 1; },
5 => { s[0] ^= 1; s[1] ^= 1; s[2] ^= 1; s.swap(0, 1); },
6 => { s[0] ^= 1; s[1] ^= 1; s[2] ^= 1; s.swap(1, 2); },
7 => { s.swap(0, 2); },
_ => {},
}
}
let mut idx = index;
let mut mask = 1u32 << (BITS - 1);
let mut h = [0u32; 3];
let mut s = [0u32; 3];
for _ in 0..BITS {
let quad = (idx & 7) as u32;
idx >>= 3;
hilbert_decode_step(quad, &mut s);
h[0] |= s[0] & mask;
h[1] |= s[1] & mask;
h[2] |= s[2] & mask;
mask >>= 1;
}
(h[0], h[1], h[2])
}
// Converts 3D coordinates (x, y, z) to a 72-bit Hilbert index.
pub fn hilbert_coord_to_index_24([x, y, z]: [u32; 3]) -> u128 {
const BITS: u32 = 24;
// Helper function to encode coordinates to a Hilbert quad.
fn hilbert_encode_step(s: &mut [u32; 3]) -> u32 {
let mut quad = 0u32;
let mut bits = [s[0], s[1], s[2]];
if bits[0] > bits[1] { bits.swap(0, 1); quad ^= 1; }
if bits[1] > bits[2] { bits.swap(1, 2); quad ^= 2; }
if bits[0] > bits[1] { bits.swap(0, 1); quad ^= 1; }
s[0] = bits[0];
s[1] = bits[1];
s[2] = bits[2];
quad
}
let mut index = 0u128;
let h = [x, y, z];
let mut s = [0u32; 3];
let mut mask = 1u32 << (BITS - 1);
for _ in 0..BITS {
s[0] = (h[0] & mask) >> (BITS - 1);
s[1] = (h[1] & mask) >> (BITS - 1);
s[2] = (h[2] & mask) >> (BITS - 1);
let quad = hilbert_encode_step(&mut s);
index = (index << 3) | quad as u128;
mask >>= 1;
}
index
}
+268
View File
@@ -0,0 +1,268 @@
use anyhow::anyhow;
use std::io::Read;
use half::f16;
use super::{
decompress_data, deobfuscate, encode_scale, first_cumulative,
PackedSplats, Wlg0Gaussian, Wlg0Header, Wlg0Settings,
};
pub(super) trait Wlg0Decoder {
fn init_num_splats(&mut self, num_splats: usize);
fn write_center(&mut self, index: usize, dim: usize, value: f32);
fn write_scale(&mut self, index: usize, dim: usize, value: f32);
fn write_quaternion_f32(&mut self, index: usize, dim: usize, value: f32);
fn write_quaternion_i8(&mut self, index: usize, dim: usize, value: i8);
fn write_quaternion_full_f32(&mut self, index: usize, value: [f32; 4]);
fn write_rgba(&mut self, index: usize, dim: usize, value: u8);
}
impl Wlg0Decoder for Vec<Wlg0Gaussian> {
fn init_num_splats(&mut self, num_splats: usize) {
self.resize(num_splats, Wlg0Gaussian::default());
}
fn write_center(&mut self, index: usize, dim: usize, value: f32) {
self[index].center[dim] = value;
}
fn write_scale(&mut self, index: usize, dim: usize, value: f32) {
self[index].ln_scale[dim] = value.ln();
}
fn write_quaternion_f32(&mut self, index: usize, dim: usize, value: f32) {
self[index].quaternion[dim] = value;
}
fn write_quaternion_i8(&mut self, _index: usize, _dim: usize, _value: i8) {
// No-op sing we use the f32 variant
}
fn write_quaternion_full_f32(&mut self, _index: usize, _value: [f32; 4]) {
// No-op
}
fn write_rgba(&mut self, index: usize, dim: usize, value: u8) {
self[index].color[dim] = value as f32 / 255.0;
}
}
impl Wlg0Decoder for PackedSplats {
fn init_num_splats(&mut self, num_splats: usize) {
self.0.resize(num_splats * 4, 0);
}
fn write_center(&mut self, index: usize, dim: usize, value: f32) {
match dim {
0 => { self.0[index * 4 + 1] |= f16::from_f32(value).to_bits() as u32; },
1 => { self.0[index * 4 + 1] |= (f16::from_f32(value).to_bits() as u32) << 16; },
2 => { self.0[index * 4 + 2] |= f16::from_f32(value).to_bits() as u32; },
_ => unreachable!(),
}
}
fn write_scale(&mut self, index: usize, dim: usize, value: f32) {
let scale8 = encode_scale(value);
match dim {
0 => { self.0[index * 4 + 3] |= scale8 as u32; },
1 => { self.0[index * 4 + 3] |= (scale8 as u32) << 8; },
2 => { self.0[index * 4 + 3] |= (scale8 as u32) << 16; },
_ => unreachable!(),
}
}
fn write_quaternion_f32(&mut self, _index: usize, _dim: usize, _value: f32) {
// No-op since we use the i8 variant
}
fn write_quaternion_i8(&mut self, _index: usize, _dim: usize, _value: i8) {
// Old Qxyz encoding:
// match dim {
// 0 => { self.0[index * 4 + 2] |= ((value as u32) & 0xff) << 16; },
// 1 => { self.0[index * 4 + 2] |= ((value as u32) & 0xff) << 24; },
// 2 => { self.0[index * 4 + 3] |= ((value as u32) & 0xff) << 24; },
// 3 => {
// // Quaternion w is inferred from xyz in PackedSplats
// },
// _ => unreachable!(),
// }
}
fn write_quaternion_full_f32(&mut self, index: usize, value: [f32; 4]) {
// Encode as OctXy88R8
let q = if value[3] < 0.0 { value.map(|v| -v) } else { value };
let theta = 2.0 * q[3].acos();
let xyz_norm = (q[0] * q[0] + q[1] * q[1] + q[2] * q[2]).sqrt();
let axis = if xyz_norm < 1e-6 {
[1.0, 0.0, 0.0]
} else {
[q[0] / xyz_norm, q[1] / xyz_norm, q[2] / xyz_norm]
};
let sum = axis[0].abs() + axis[1].abs() + axis[2].abs();
let p = [axis[0] / sum, axis[1] / sum];
let p = if axis[2] >= 0.0 { p } else {
[
(1.0 - p[1].abs()) * p[0].signum(),
(1.0 - p[0].abs()) * p[1].signum(),
]
};
let uv = [
((p[0] * 0.5 + 0.5) * 255.0).round() as u8,
((p[1] * 0.5 + 0.5) * 255.0).round() as u8,
];
let angle = (theta * (255.0 / 3.14159265359)).round().clamp(0.0, 255.0) as u8;
self.0[index * 4 + 2] |= (uv[0] as u32) << 16;
self.0[index * 4 + 2] |= (uv[1] as u32) << 24;
self.0[index * 4 + 3] |= (angle as u32) << 24;
}
fn write_rgba(&mut self, index: usize, dim: usize, value: u8) {
match dim {
0 => { self.0[index * 4] |= value as u32; },
1 => { self.0[index * 4] |= (value as u32) << 8; },
2 => { self.0[index * 4] |= (value as u32) << 16; },
3 => { self.0[index * 4] |= (value as u32) << 24; },
_ => unreachable!(),
}
}
}
pub(super) fn decode_internal<D: Wlg0Decoder>(
data: &mut [u8],
settings: &Wlg0Settings,
decoder: &mut D,
) -> anyhow::Result<()> {
if settings.enable_obfuscation {
deobfuscate(data);
}
let data = if settings.enable_compression {
decompress_data(data)?
} else {
data.to_vec()
};
let mut reader = std::io::Cursor::new(&data);
let header = Wlg0Header::read(&mut reader)?;
let num_splats = header.num_splats as usize;
decoder.init_num_splats(num_splats);
if !settings.enable_split_dims {
return Err(anyhow!("Unsupported WLG !settings.enable_split_dims"));
}
let mut data_u8: Vec<u8> = vec![0; num_splats];
if settings.enable_split_center_bytes {
let mut data_u32: Vec<u32> = vec![0; num_splats];
for d in 0..3 {
data_u32.iter_mut().for_each(|v| *v = 0);
if settings.enable_24bit_center {
reader.read_exact(&mut data_u8)?;
if settings.enable_first_differences {
first_cumulative(&mut data_u8);
}
for (i, &byte) in data_u8.iter().enumerate() {
data_u32[i] = (byte as u32) << 16;
}
}
reader.read_exact(&mut data_u8)?;
if settings.enable_first_differences {
first_cumulative(&mut data_u8);
}
for (i, &byte) in data_u8.iter().enumerate() {
data_u32[i] |= (byte as u32) << 8;
}
reader.read_exact(&mut data_u8)?;
if settings.enable_first_differences {
first_cumulative(&mut data_u8);
}
for (i, &byte) in data_u8.iter().enumerate() {
data_u32[i] |= byte as u32;
}
let resolution = if settings.enable_24bit_center {
16777215.0
} else {
65535.0
};
let scale = header.center_scale / resolution;
let offset = header.center_offset[d];
for (i, &value) in data_u32.iter().enumerate() {
decoder.write_center(i, d, value as f32 * scale + offset);
}
}
} else {
return Err(anyhow!(
"Unsupported WLG !settings.enable_split_center_bytes"
));
}
for d in 0..3 {
reader.read_exact(&mut data_u8)?;
let scale = header.ln_scale_max - header.ln_scale_min;
let offset = header.ln_scale_min;
for (i, &byte) in data_u8.iter().enumerate() {
let float = ((byte as f32 / 255.0) * scale + offset).exp();
decoder.write_scale(i, d, float);
}
}
{
let mut quaternions: Vec<i8> = vec![0; num_splats * 3];
for d in 0..3 {
reader.read_exact(&mut data_u8)?;
for (i, &byte) in data_u8.iter().enumerate() {
quaternions[i * 3 + d] = byte as i8;
}
}
for i in 0..num_splats {
let quat: [f32; 3] = [
quaternions[i * 3 + 0] as f32 / 127.0,
quaternions[i * 3 + 1] as f32 / 127.0,
quaternions[i * 3 + 2] as f32 / 127.0,
];
let w = (1.0 - quat.iter().map(|&v| v.powi(2)).sum::<f32>())
.max(0.0)
.sqrt();
let w8 = (w * 127.0).round() as i8;
for d in 0..3 {
decoder.write_quaternion_f32(i, d, quat[d]);
decoder.write_quaternion_i8(i, d, quaternions[i * 3 + d]);
}
decoder.write_quaternion_f32(i, 3, w);
decoder.write_quaternion_i8(i, 3, w8);
decoder.write_quaternion_full_f32(i, [quat[0], quat[1], quat[2], w]);
}
}
reader.read_exact(&mut data_u8)?;
if settings.enable_first_differences {
first_cumulative(&mut data_u8);
}
for (i, &byte) in data_u8.iter().enumerate() {
decoder.write_rgba(i, 3, byte);
}
for d in 0..3 {
reader.read_exact(&mut data_u8)?;
if settings.enable_first_differences {
first_cumulative(&mut data_u8);
}
for (i, &byte) in data_u8.iter().enumerate() {
decoder.write_rgba(i, d, byte);
}
}
let position = reader.position() as usize;
if position != data.len() {
return Err(anyhow!("Invalid WLG data size"));
}
Ok(())
}
+283
View File
@@ -0,0 +1,283 @@
use std::io::Write;
use super::{compress_data, obfuscate, Wlg0Gaussian, Wlg0Settings, Wlg0Header};
use crate::ordering::{
hilbert_coord_to_index, hilbert_coord_to_index_24, morton_coord_to_index,
morton_coord_to_index_24,
};
// WLG splat representation with fixed-point encoding.
#[derive(Debug)]
struct WlgSplat {
order: u128,
center: [u32; 3],
scale: [u8; 3],
quaternion: [i8; 3],
opacity: u8,
color: [u8; 3],
}
impl WlgSplat {
fn new(
settings: &Wlg0Settings,
index: usize,
center: [f32; 3],
scale: [f32; 3],
mut quaternion: [f32; 4],
opacity: f32,
color: [f32; 3],
) -> Self {
use std::array::from_fn;
let center = if settings.enable_24bit_center {
from_fn(|d| (center[d] * 16777215.0).clamp(0.0, 16777215.0).round() as u32)
} else {
from_fn(|d| (center[d] * 65535.0).clamp(0.0, 65535.0).round() as u32)
};
if quaternion[3] < 0.0 {
quaternion = from_fn(|d| -quaternion[d]);
}
Self {
order: if settings.enable_hilbert_reordering {
// Convert X/Y/Z into a single Hilbert curve index
if settings.enable_24bit_center {
hilbert_coord_to_index_24(center) as u128
} else {
hilbert_coord_to_index(from_fn(|d| center[d] as u16)) as u128
}
} else if settings.enable_morton_reordering {
// Interleave X/Y/Z bits into a single Morton index
if settings.enable_24bit_center {
morton_coord_to_index_24(center) as u128
} else {
morton_coord_to_index(from_fn(|d| center[d] as u16)) as u128
}
} else {
index as u128
},
center,
scale: from_fn(|d| (scale[d] * 255.0).clamp(0.0, 255.0).round() as u8),
quaternion: from_fn(|d| (quaternion[d] * 127.0).clamp(-127.0, 127.0).round() as i8),
opacity: (opacity * 255.0).clamp(0.0, 255.0).round() as u8,
color: from_fn(|d| (color[d] * 255.0).clamp(0.0, 255.0).round() as u8),
}
}
}
// Encoding functions
fn first_differences(settings: &Wlg0Settings, mut data: Vec<u8>) -> Vec<u8> {
if settings.enable_first_differences {
super::first_differences(&mut data);
}
data
}
pub(super) fn encode_internal(
settings: &Wlg0Settings,
gaussians: &[Wlg0Gaussian],
) -> anyhow::Result<Vec<u8>> {
use std::array::from_fn;
let num_splats = gaussians.len();
// Compute min/max bounds for Gsplat centers
let center_min_max = ([f32::INFINITY; 3], [-f32::INFINITY; 3]);
let (center_min, center_max) = gaussians.iter().fold(center_min_max, |(min, max), g| {
let min = from_fn(|i| min[i].min(g.center[i]));
let max = from_fn(|i| max[i].max(g.center[i]));
(min, max)
});
// println!("Center min: {:?}, max: {:?}", center_min, center_max);
// Compute scale factor and offset for center values
let center_ranges: [f32; 3] = from_fn(|i| center_max[i] - center_min[i]);
let center_scale = center_ranges.into_iter().reduce(|a, b| a.max(b)).unwrap();
let center_offset = center_min;
// Compute min/max bounds for Gsplat ln_scales
let ln_scale_min_max = (f32::INFINITY, -f32::INFINITY);
let (ln_scale_min, ln_scale_max) =
gaussians
.iter()
.fold(ln_scale_min_max, |(mut min, mut max), g| {
g.ln_scale.iter().for_each(|&ln_scale| {
let ln_scale = ln_scale.clamp(-10.0, 10.0);
min = min.min(ln_scale);
max = max.max(ln_scale);
});
(min, max)
});
// println!("ln_scale min: {}, max: {}", ln_scale_min, ln_scale_max);
let mut splats: Vec<_> = gaussians
.iter()
.enumerate()
.map(|(index, g)| {
WlgSplat::new(
settings,
index,
from_fn(|d| (g.center[d] - center_offset[d]) / center_scale),
from_fn(|d| (g.ln_scale[d] - ln_scale_min) / (ln_scale_max - ln_scale_min)),
g.quaternion,
g.opacity,
from_fn(|d| g.color[d]),
)
})
.collect();
// Reorder splats using the provided ordering key, either Morton or original index.
splats.sort_by_key(|splat| splat.order);
// Create buffer to write header and splats to
let mut buffer: Vec<u8> = Vec::new();
let header = Wlg0Header {
center_scale,
center_offset,
ln_scale_min,
ln_scale_max,
num_splats: num_splats as u32,
max_sh_order: 0,
num_sh_splats: [num_splats as u32, 0, 0, 0],
};
header.write(&mut buffer)?;
if settings.enable_split_dims {
// Write dimensions separately, i.e. a column-oriented format.
for d in 0..3 {
if settings.enable_split_center_bytes {
// Split center values into separate bytes (MSB order)
if settings.enable_24bit_center {
let data: Vec<u8> = splats
.iter()
.map(|splat| (splat.center[d] >> 16) as u8)
.collect();
buffer.write_all(&first_differences(settings, data))?;
}
let data: Vec<u8> = splats
.iter()
.map(|splat| (splat.center[d] >> 8) as u8)
.collect();
buffer.write_all(&first_differences(settings, data))?;
let data: Vec<u8> = splats.iter().map(|splat| splat.center[d] as u8).collect();
buffer.write_all(&first_differences(settings, data))?;
} else {
if settings.enable_24bit_center {
// Write 24-bit center values into u8 array in LSB order
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| {
let center = splat.center[d];
[center as u8, (center >> 8) as u8, (center >> 16) as u8]
})
.collect();
buffer.write_all(&data)?;
} else {
let data_u16: Vec<u16> =
splats.iter().map(|splat| splat.center[d] as u16).collect();
let data: Vec<u8> = data_u16
.iter()
.flat_map(|&v| v.to_le_bytes().to_vec())
.collect();
buffer.write_all(&data)?;
}
}
}
for d in 0..3 {
let data: Vec<u8> = splats.iter().map(|splat| splat.scale[d]).collect();
buffer.write_all(&data)?;
}
for d in 0..3 {
let data: Vec<u8> = splats
.iter()
.map(|splat| splat.quaternion[d] as u8)
.collect();
buffer.write_all(&data)?;
}
{
let data: Vec<u8> = splats.iter().map(|splat| splat.opacity).collect();
buffer.write_all(&first_differences(settings, data))?;
}
for d in 0..3 {
let data: Vec<u8> = splats.iter().map(|splat| splat.color[d]).collect();
buffer.write_all(&first_differences(settings, data))?;
}
} else {
// !enable_split_dims, write all dimensions interleaved (array of structs).
if settings.enable_split_center_bytes {
if settings.enable_24bit_center {
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.center.iter().map(|&v| (v >> 16) as u8))
.collect();
buffer.write_all(&first_differences(settings, data))?;
}
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.center.iter().map(|&v| (v >> 8) as u8))
.collect();
buffer.write_all(&first_differences(settings, data))?;
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.center.iter().map(|&v| v as u8))
.collect();
buffer.write_all(&first_differences(settings, data))?;
} else {
if settings.enable_24bit_center {
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| {
splat
.center
.iter()
.flat_map(|&d| [d as u8, (d >> 8) as u8, (d >> 16) as u8])
})
.collect();
buffer.write_all(&data)?;
} else {
let data_u16: Vec<u16> = splats
.iter()
.flat_map(|splat| splat.center.iter().map(|&v| v as u16))
.collect();
let data: Vec<u8> = data_u16
.iter()
.flat_map(|&v| v.to_le_bytes().to_vec())
.collect();
buffer.write_all(&data)?;
}
}
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.scale.iter().copied())
.collect();
buffer.write_all(&data)?;
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.quaternion.iter().map(|&v| v as u8))
.collect();
buffer.write_all(&data)?;
let data: Vec<u8> = splats.iter().map(|splat| splat.opacity).collect();
buffer.write_all(&first_differences(settings, data))?;
let data: Vec<u8> = splats
.iter()
.flat_map(|splat| splat.color.iter().copied())
.collect();
buffer.write_all(&first_differences(settings, data))?;
}
let mut payload = if settings.enable_compression {
compress_data(&buffer, settings.zstd_compression_level)?
} else {
buffer
};
if settings.enable_obfuscation {
obfuscate(&mut payload);
}
Ok(payload)
}
+274
View File
@@ -0,0 +1,274 @@
use serde::{Deserialize, Serialize};
use std::io::{Read, Write};
use anyhow::anyhow;
mod decode;
mod encode;
#[derive(Debug, Default, Clone)]
pub struct Wlg0Gaussian {
pub center: [f32; 3],
pub ln_scale: [f32; 3],
pub quaternion: [f32; 4],
pub opacity: f32,
pub color: [f32; 3],
}
pub const LN_SCALE_MIN: f32 = -9.0;
pub const LN_SCALE_MAX: f32 = 9.0;
pub const LN_RESCALE: f32 = (LN_SCALE_MAX - LN_SCALE_MIN) / 254.0; // 1..=255
pub fn encode_scale(scale: f32) -> u8 {
if scale == 0.0 {
0
} else {
// Allow scales below LN_SCALE_MIN to be encoded as 0, which signifies a 2DGS
((scale.ln() - LN_SCALE_MIN) / LN_RESCALE + 1.0).clamp(0.0, 255.0).round() as u8
}
}
pub fn decode_scale(scale: u8) -> f32 {
if scale == 0 {
0.0
} else {
(LN_SCALE_MIN + (scale - 1) as f32 * LN_RESCALE).exp()
}
}
pub struct PackedSplats(pub Vec<u32>);
#[derive(Debug, Serialize, Deserialize)]
struct Wlg0Signature {
magic: [u8; 4], // "WLG0"
version: u32, // 1 or 2
}
impl Wlg0Signature {
const MAGIC: [u8; 4] = *b"WLG0";
fn new(version: u32) -> Self {
Self {
magic: Self::MAGIC,
version,
}
}
fn check_version(&self) -> Option<u32> {
if self.magic == Self::MAGIC {
Some(self.version)
} else {
None
}
}
fn write<W: Write>(&self, writer: &mut W) -> std::io::Result<()> {
writer.write_all(&bincode::serialize(self).unwrap())
}
fn read<R: Read>(reader: &mut R) -> std::io::Result<Self> {
let mut buf = [0; std::mem::size_of::<Self>()];
reader.read_exact(&mut buf)?;
Ok(bincode::deserialize(&buf).unwrap())
}
}
#[derive(Debug, Clone)]
pub struct Wlg0Settings {
pub version: u32,
pub enable_compression: bool,
pub enable_first_differences: bool,
pub enable_morton_reordering: bool,
pub enable_hilbert_reordering: bool,
pub enable_split_center_bytes: bool,
pub enable_obfuscation: bool,
pub enable_split_dims: bool,
pub enable_24bit_center: bool,
pub zstd_compression_level: i32, // 0 selects default level (3), max is 22
}
const WLG012_SETTINGS: [Wlg0Settings; 3] = [
// 0: WLGv0: Not a valid version, for testing purposes only
Wlg0Settings {
version: 0,
enable_compression: true,
enable_first_differences: true,
enable_morton_reordering: false,
enable_hilbert_reordering: true,
enable_split_center_bytes: true,
enable_obfuscation: true,
enable_split_dims: true,
enable_24bit_center: true,
zstd_compression_level: 0,
},
// 1: WLGv1: Default settings for WLG v1 files
Wlg0Settings {
version: 1,
enable_compression: true,
enable_first_differences: true,
enable_morton_reordering: true,
enable_hilbert_reordering: false,
enable_split_center_bytes: true,
enable_obfuscation: true,
enable_split_dims: true,
enable_24bit_center: false,
zstd_compression_level: 0,
},
// 2: WLGv2: Default settings for WLG v2 files
Wlg0Settings {
version: 2,
enable_compression: true,
enable_first_differences: true,
enable_morton_reordering: true,
enable_hilbert_reordering: false,
enable_split_center_bytes: true,
enable_obfuscation: true,
enable_split_dims: true,
enable_24bit_center: true,
zstd_compression_level: 0,
},
];
#[derive(Debug, Serialize, Deserialize)]
pub struct Wlg0Header {
// center[d] = center_int[d] * center_scale + center_offset
pub center_scale: f32,
pub center_offset: [f32; 3],
// scale[d] = exp(scale_int[d] / 255 * (scale_max - scale_min) + scale_min)
pub ln_scale_min: f32,
pub ln_scale_max: f32,
// Total # splats, regardless of SH order
pub num_splats: u32,
// 0 for "no spherical harmonics"
pub max_sh_order: u32,
// Supports maximum of SH=3
pub num_sh_splats: [u32; 4],
}
impl Wlg0Header {
pub fn write<W: Write>(&self, writer: &mut W) -> std::io::Result<()> {
writer.write_all(&bincode::serialize(self).unwrap())
}
pub fn read<R: Read>(reader: &mut R) -> std::io::Result<Self> {
let mut buf = [0; std::mem::size_of::<Self>()];
reader.read_exact(&mut buf)?;
Ok(bincode::deserialize(&buf).unwrap())
}
}
fn obfuscate(data: &mut [u8]) {
let mut prev: u8 = 0;
let mut state: u32 = 0x1AB51AB5;
for byte in data {
state = state.wrapping_mul(1664525).wrapping_add(1013904223);
let rand_byte = (state >> 8) as u8;
let obf_byte = (*byte ^ prev).rotate_left(3) ^ rand_byte;
prev = obf_byte;
*byte = obf_byte;
}
}
fn deobfuscate(data: &mut [u8]) {
let mut prev: u8 = 0;
let mut state: u32 = 0x1AB51AB5;
for obf_byte in data {
state = state.wrapping_mul(1664525).wrapping_add(1013904223);
let rand_byte = (state >> 8) as u8;
let tmp = (*obf_byte ^ rand_byte).rotate_right(3);
let byte = tmp ^ prev;
prev = *obf_byte;
*obf_byte = byte;
}
}
fn compress_data(data: &[u8], _zstd_compression_level: i32) -> anyhow::Result<Vec<u8>> {
use ruzstd::encoding::{compress_to_vec, CompressionLevel};
let compressed_data = compress_to_vec(data, CompressionLevel::Fastest);
Ok(compressed_data)
}
fn decompress_data(data: &[u8]) -> anyhow::Result<Vec<u8>> {
let mut decoder = ruzstd::decoding::StreamingDecoder::new(data)?;
let mut decompressed_data = Vec::new();
decoder.read_to_end(&mut decompressed_data)?;
Ok(decompressed_data)
}
fn first_differences(data: &mut [u8]) {
// Compute first differences for array to aid compression
for i in (1..data.len()).rev() {
data[i] = data[i].wrapping_sub(data[i - 1]);
}
}
fn first_cumulative(data: &mut [u8]) {
// Invert first differences
for i in 1..data.len() {
data[i] = data[i].wrapping_add(data[i - 1]);
}
}
pub fn decode12(bytes: &mut [u8]) -> anyhow::Result<(Wlg0Settings, Vec<Wlg0Gaussian>)> {
let (offset, version) = {
let mut reader = std::io::Cursor::new(&bytes);
let signature = Wlg0Signature::read(&mut reader)?;
let Some(version) = signature.check_version() else {
return Err(anyhow!("Invalid WLG signature"));
};
if (version != 1) && (version != 2) {
return Err(anyhow!("Unsupported WLG version"));
}
let offset = reader.position() as usize;
(offset, version)
};
let settings = WLG012_SETTINGS[version as usize].clone();
let mut gaussians = Vec::new();
decode::decode_internal(&mut bytes[offset..], &settings, &mut gaussians)?;
Ok((settings, gaussians))
}
pub fn decode12_packed(bytes: &mut [u8]) -> anyhow::Result<(Wlg0Settings, PackedSplats)> {
let (offset, version) = {
let mut reader = std::io::Cursor::new(&bytes);
let signature = Wlg0Signature::read(&mut reader)?;
let Some(version) = signature.check_version() else {
return Err(anyhow!("Invalid WLG signature"));
};
if (version != 1) && (version != 2) {
return Err(anyhow!("Unsupported WLG version"));
}
let offset = reader.position() as usize;
(offset, version)
};
let settings = WLG012_SETTINGS[version as usize].clone();
let mut packed_splats = PackedSplats(Vec::new());
decode::decode_internal(&mut bytes[offset..], &settings, &mut packed_splats)?;
Ok((settings, packed_splats))
}
pub enum Wlg0EncodeSettings {
Wlg1,
Wlg2,
Custom(Wlg0Settings),
}
pub fn encode12(
settings: &Wlg0EncodeSettings,
gaussians: &[Wlg0Gaussian],
) -> anyhow::Result<(Wlg0Settings, Vec<u8>)> {
let settings = match settings {
Wlg0EncodeSettings::Wlg1 => &WLG012_SETTINGS[1],
Wlg0EncodeSettings::Wlg2 => &WLG012_SETTINGS[2],
Wlg0EncodeSettings::Custom(custom) => custom,
}
.clone();
let payload = encode::encode_internal(&settings, gaussians)?;
let mut buffer = Vec::new();
Wlg0Signature::new(settings.version).write(&mut buffer)?;
buffer.write_all(&payload)?;
Ok((settings, buffer))
}