From 156bbf400ef4c628430f7778ea64a619d44a3231 Mon Sep 17 00:00:00 2001 From: Andreas Sundquist Date: Mon, 26 May 2025 13:46:16 -0700 Subject: [PATCH] SPZ writing, exporting, & viewer. Creates a new `npm run to-spz` which runs `scripts/to-spz.js` to convert any Gsplat input file to compress .spz. Also includes a configurable viewer examples/viewer that can load splats via drag-n-drop, from url query param, the URL text list, can edit sets of SplatMesh, and export composite .SPZ with some filters. Refactored some loading code from worker.ts to loader files so they can be called from Nodejs. Created api transcodeSpz with options to combine multiple source Gsplat files into one without losing coding quality along the way. --- examples/viewer/index.html | 673 ++++++++++++++++++++++++++++++++++++- package.json | 1 + scripts/to-spz.js | 270 +++++++++++++++ src/PackedSplats.ts | 4 +- src/SplatGenerator.ts | 19 +- src/SplatLoader.ts | 331 ++++++++++++++---- src/SplatMesh.ts | 11 +- src/antisplat.ts | 120 +++++++ src/controls.ts | 15 +- src/dyno/transform.ts | 34 +- src/generators/snow.ts | 2 +- src/generators/static.ts | 2 +- src/index.ts | 2 +- src/ksplat.ts | 576 +++++++++++++++++++++++++++++++ src/splatConstructors.ts | 7 +- src/spz.ts | 627 +++++++++++++++++++++++++++++++++- src/utils.ts | 13 + src/worker.ts | 398 +--------------------- vite.config.ts | 8 + 19 files changed, 2605 insertions(+), 508 deletions(-) create mode 100644 scripts/to-spz.js create mode 100644 src/antisplat.ts create mode 100644 src/ksplat.ts diff --git a/examples/viewer/index.html b/examples/viewer/index.html index 38cbbd3..cb4f796 100644 --- a/examples/viewer/index.html +++ b/examples/viewer/index.html @@ -2,43 +2,692 @@ - Forge • Hello World + Forge • Viewer +
+
+
+ +
+
+ diff --git a/package.json b/package.json index 3e7ea3a..3c7c958 100644 --- a/package.json +++ b/package.json @@ -11,6 +11,7 @@ "build:wasm": "node rust/build_wasm.js", "build:watch": "onchange 'src/**/*.{ts,glsl}' -- npm run build", "clean": "rm -rf dist/ && rm -rf node_modules/ && rm -rf rust/target/ && rm -rf rust/forge-internal-rs/pkg/ && npm run assets:clean", + "to-spz": "node scripts/to-spz.js", "dev": "npm run build && (vite --host & npm run build:watch)", "deploy": "node scripts/deploy.js", "docs": "mkdocs serve", diff --git a/scripts/to-spz.js b/scripts/to-spz.js new file mode 100644 index 0000000..26d03cb --- /dev/null +++ b/scripts/to-spz.js @@ -0,0 +1,270 @@ +#!/usr/bin/env node + +import { execSync } from "node:child_process"; +import fs from "node:fs/promises"; +import path from "node:path"; +import { URL } from "node:url"; + +// Import directly from source to avoid worker system +import { transcodeSpz } from "../dist/forge.module.js"; + +async function main() { + const args = process.argv.slice(2); + if (args.length === 0 || args.includes("-h") || args.includes("--help")) { + console.log(`Usage: to-spz.js [options] ... + +Convert Gaussian Splat files to optimized .spz format. + +Supported input formats: + .ply PLY Gaussian Splat files + .wlg WorldLabs Gaussian format + .spz SPZ format (for reprocessing/filtering) + .splat AntiSplat format + .ksplat KSplat format + http(s) Direct URLs to splat files + +Options: + --filter-zero Filter out splats with zero opacity + --filter-opacity N Filter out splats with opacity <= N (0.0-1.0) + --min x,y,z Minimum AABB (e.g. --min 0,0,0) + --max x,y,z Maximum AABB (e.g. --max 1,1,1) + --min-x N Minimum X coordinate + --max-x N Maximum X coordinate + --min-y N Minimum Y coordinate + --max-y N Maximum Y coordinate + --min-z N Minimum Z coordinate + --max-z N Maximum Z coordinate + --x-range min,max X coordinate range (e.g. --x-range -1,1) + --y-range min,max Y coordinate range (e.g. --y-range -1,1) + --z-range min,max Z coordinate range (e.g. --z-range -1,1) + --max-sh N Maximum SH degree to output (0-3, default: auto-detect from input) + --fractional-bits N Fractional bits for coordinate precision (default: 12, range: 6-24) +`); + process.exit(0); + } + + let opacityThreshold = null; // null = no filtering + let minAABB = [ + Number.NEGATIVE_INFINITY, + Number.NEGATIVE_INFINITY, + Number.NEGATIVE_INFINITY, + ]; + let maxAABB = [ + Number.POSITIVE_INFINITY, + Number.POSITIVE_INFINITY, + Number.POSITIVE_INFINITY, + ]; + let maxShDegree = null; // null = auto-detect + let fractionalBits = 12; // default value + const inputs = []; + + for (let i = 0; i < args.length; i++) { + const arg = args[i]; + if (arg === "--filter-zero") { + opacityThreshold = 0; + } else if (arg === "--filter-opacity") { + const val = Number(args[++i]); + if (Number.isNaN(val) || val < 0 || val > 1) { + throw new Error("Invalid --filter-opacity value, expected 0.0-1.0"); + } + opacityThreshold = val; + } else if (arg === "--min") { + const val = args[++i]; + minAABB = val.split(",").map(Number); + if (minAABB.length !== 3 || minAABB.some(Number.isNaN)) { + throw new Error("Invalid --min value, expected comma-separated x,y,z"); + } + } else if (arg === "--max") { + const val = args[++i]; + maxAABB = val.split(",").map(Number); + if (maxAABB.length !== 3 || maxAABB.some(Number.isNaN)) { + throw new Error("Invalid --max value, expected comma-separated x,y,z"); + } + } else if (arg === "--max-sh") { + const val = Number(args[++i]); + if (Number.isNaN(val) || val < 0 || val > 3) { + throw new Error("Invalid --max-sh value, expected 0-3"); + } + maxShDegree = val; + } else if (arg === "--fractional-bits") { + const val = Number(args[++i]); + if (Number.isNaN(val) || val < 6 || val > 24) { + throw new Error("Invalid --fractional-bits value, expected 6-24"); + } + fractionalBits = val; + } else if (arg === "--min-x") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --min-x value, expected a number"); + } + minAABB[0] = val; + } else if (arg === "--max-x") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --max-x value, expected a number"); + } + maxAABB[0] = val; + } else if (arg === "--min-y") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --min-y value, expected a number"); + } + minAABB[1] = val; + } else if (arg === "--max-y") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --max-y value, expected a number"); + } + maxAABB[1] = val; + } else if (arg === "--min-z") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --min-z value, expected a number"); + } + minAABB[2] = val; + } else if (arg === "--max-z") { + const val = Number(args[++i]); + if (Number.isNaN(val)) { + throw new Error("Invalid --max-z value, expected a number"); + } + maxAABB[2] = val; + } else if (arg === "--x-range") { + const val = args[++i]; + const range = val.split(",").map(Number); + if (range.length !== 2 || range.some(Number.isNaN)) { + throw new Error("Invalid --x-range value, expected min,max"); + } + minAABB[0] = range[0]; + maxAABB[0] = range[1]; + } else if (arg === "--y-range") { + const val = args[++i]; + const range = val.split(",").map(Number); + if (range.length !== 2 || range.some(Number.isNaN)) { + throw new Error("Invalid --y-range value, expected min,max"); + } + minAABB[1] = range[0]; + maxAABB[1] = range[1]; + } else if (arg === "--z-range") { + const val = args[++i]; + const range = val.split(",").map(Number); + if (range.length !== 2 || range.some(Number.isNaN)) { + throw new Error("Invalid --z-range value, expected min,max"); + } + minAABB[2] = range[0]; + maxAABB[2] = range[1]; + } else { + inputs.push(arg); + } + } + + if (inputs.length === 0) { + console.error("No input files or URLs specified."); + process.exit(1); + } + + // Collect all file inputs first + const transcodeInputs = []; + + for (const input of inputs) { + console.log(`Loading ${input}...`); + let fileBytes; + let pathOrUrl; + + if (input.startsWith("http://") || input.startsWith("https://")) { + const res = await fetch(input); + if (!res.ok) { + console.error( + `Failed to fetch ${input}: ${res.status} ${res.statusText}`, + ); + continue; + } + const buffer = await res.arrayBuffer(); + fileBytes = new Uint8Array(buffer); + pathOrUrl = input; + } else { + const data = await fs.readFile(input); + fileBytes = new Uint8Array(data); + pathOrUrl = input; + } + + transcodeInputs.push({ + fileBytes: fileBytes.slice(), + pathOrUrl, + transform: { + translate: [0, 0, 0], + quaternion: [0, 0, 0, 1], + scale: 1, + }, + }); + } + + if (transcodeInputs.length === 0) { + console.error("No valid inputs to process."); + process.exit(1); + } + + // Setup clipping bounds if specified + let clipXyz = undefined; + if ( + minAABB[0] !== Number.NEGATIVE_INFINITY || + minAABB[1] !== Number.NEGATIVE_INFINITY || + minAABB[2] !== Number.NEGATIVE_INFINITY || + maxAABB[0] !== Number.POSITIVE_INFINITY || + maxAABB[1] !== Number.POSITIVE_INFINITY || + maxAABB[2] !== Number.POSITIVE_INFINITY + ) { + clipXyz = { + min: minAABB, + max: maxAABB, + }; + } + + // Setup opacity threshold filtering + if (opacityThreshold !== null) { + console.log( + `Applying opacity filtering: removing splats with opacity <= ${opacityThreshold}`, + ); + } + + console.log( + `Processing ${transcodeInputs.length} input file(s) with transcodeSpz...`, + ); + const transcode = await transcodeSpz({ + inputs: transcodeInputs, + maxSh: maxShDegree, + clipXyz, + fractionalBits, + opacityThreshold, + }); + + // Write output file + const firstInput = transcodeInputs[0]; + const baseName = path.basename(firstInput.pathOrUrl); + const nameNoExt = baseName.replace(/\.[^.]+$/, ""); + const outDir = + firstInput.pathOrUrl.startsWith("http://") || + firstInput.pathOrUrl.startsWith("https://") + ? process.cwd() + : path.dirname(firstInput.pathOrUrl); + + const outputName = + transcodeInputs.length === 1 + ? `${nameNoExt}.spz` + : `combined_${transcodeInputs.length}_files.spz`; + + const outPath = path.join(outDir, outputName); + await fs.writeFile(outPath, transcode.fileBytes); + console.log(`Wrote ${outPath} (${transcode.fileBytes.length} bytes)`); + + if (transcode.clippedCount && transcode.clippedCount > 0) { + console.log(`Clipped ${transcode.clippedCount} splats.`); + console.log( + `Consider decreasing fractional-bits from ${fractionalBits} to reduce clipping.`, + ); + } +} + +main().catch((err) => { + console.error(err); + process.exit(1); +}); diff --git a/src/PackedSplats.ts b/src/PackedSplats.ts index 7f2a43f..35b829a 100644 --- a/src/PackedSplats.ts +++ b/src/PackedSplats.ts @@ -29,6 +29,8 @@ export type PackedSplatsOptions = { // auto-detected (.splat, .ksplat). (default: undefined auto-detects other // formats from file contents) fileType?: SplatFileType; + // File name to use for type detection. (default: undefined) + fileName?: string; // Reserve space for at least this many splats when constructing the collection // initially. The array will automatically resize past maxSplats so setting it is // an optional optimization. (default: 0) @@ -135,7 +137,7 @@ export class PackedSplats { const unpacked = await unpackSplats({ input: fileBytes, fileType: options.fileType, - pathOrUrl: url, + pathOrUrl: options.fileName ?? url, }); this.initialize(unpacked); } diff --git a/src/SplatGenerator.ts b/src/SplatGenerator.ts index e1407b9..9c5ebdc 100644 --- a/src/SplatGenerator.ts +++ b/src/SplatGenerator.ts @@ -8,7 +8,9 @@ import { DynoVec4, Gsplat, dynoBlock, + transformDir, transformGsplat, + transformPos, } from "./dyno"; // A GsplatGenerator is a dyno program that maps an index to a Gsplat's properties @@ -81,8 +83,23 @@ export class SplatTransformer { }); } + // Apply the transform to a Vec3 position in a dyno program. + apply(position: DynoVal<"vec3">): DynoVal<"vec3"> { + return transformPos(position, { + scale: this.scale, + rotate: this.rotate, + translate: this.translate, + }); + } + + applyDir(dir: DynoVal<"vec3">): DynoVal<"vec3"> { + return transformDir(dir, { + rotate: this.rotate, + }); + } + // Apply the transform to a Gsplat in a dyno program. - modify(gsplat: DynoVal): DynoVal { + applyGsplat(gsplat: DynoVal): DynoVal { return transformGsplat(gsplat, { scale: this.scale, rotate: this.rotate, diff --git a/src/SplatLoader.ts b/src/SplatLoader.ts index 108c644..948c059 100644 --- a/src/SplatLoader.ts +++ b/src/SplatLoader.ts @@ -118,6 +118,28 @@ export function getFileExtension(pathOrUrl: string): string { return filename.slice(lastDot + 1).toLowerCase(); } +export function getSplatFileTypeFromPath( + pathOrUrl: string, +): SplatFileType | undefined { + const extension = getFileExtension(pathOrUrl); + if (extension === "ply") { + return SplatFileType.PLY; + } + if (extension === "wlg") { + return SplatFileType.WLG0; + } + if (extension === "spz") { + return SplatFileType.SPZ; + } + if (extension === "splat") { + return SplatFileType.SPLAT; + } + if (extension === "ksplat") { + return SplatFileType.KSPLAT; + } + return undefined; +} + export async function unpackSplats({ input, fileType, @@ -137,82 +159,245 @@ export async function unpackSplats({ if (!fileType) { splatFileType = getSplatFileType(fileBytes); if (!splatFileType && pathOrUrl) { - const extension = getFileExtension(pathOrUrl); - if (extension === "ply") { - splatFileType = SplatFileType.PLY; - } else if (extension === "wlg") { - splatFileType = SplatFileType.WLG0; - } else if (extension === "spz") { - splatFileType = SplatFileType.SPZ; - } else if (extension === "splat") { - splatFileType = SplatFileType.SPLAT; - } else if (extension === "ksplat") { - splatFileType = SplatFileType.KSPLAT; - } + splatFileType = getSplatFileTypeFromPath(pathOrUrl); } } - if (splatFileType === SplatFileType.WLG0) { - return await withWorker(async (worker) => { - const { packedArray, numSplats } = (await worker.call("decodeWlg", { - fileBytes, - })) as { packedArray: Uint32Array; numSplats: number }; - return { packedArray, numSplats }; - }); - } - if (splatFileType === SplatFileType.PLY) { - const ply = new PlyReader({ fileBytes }); - await ply.parseHeader(); - const numSplats = ply.numSplats; - const maxSplats = getTextureSize(numSplats).maxSplats; - const args = { fileBytes, packedArray: new Uint32Array(maxSplats * 4) }; - return await withWorker(async (worker) => { - const { packedArray, numSplats, extra } = (await worker.call( - "unpackPly", - args, - )) as { - packedArray: Uint32Array; - numSplats: number; - extra: Record; - }; - return { packedArray, numSplats, extra }; - }); - } - if (splatFileType === SplatFileType.SPZ) { - return await withWorker(async (worker) => { - const { packedArray, numSplats, extra } = (await worker.call( - "decodeSpz", - { + switch (splatFileType) { + case SplatFileType.WLG0: + return await withWorker(async (worker) => { + const { packedArray, numSplats } = (await worker.call("decodeWlg", { fileBytes, - }, - )) as { - packedArray: Uint32Array; - numSplats: number; - extra: Record; - }; - return { packedArray, numSplats, extra }; - }); + })) as { packedArray: Uint32Array; numSplats: number }; + return { packedArray, numSplats }; + }); + case SplatFileType.PLY: { + const ply = new PlyReader({ fileBytes }); + await ply.parseHeader(); + const numSplats = ply.numSplats; + const maxSplats = getTextureSize(numSplats).maxSplats; + const args = { fileBytes, packedArray: new Uint32Array(maxSplats * 4) }; + return await withWorker(async (worker) => { + const { packedArray, numSplats, extra } = (await worker.call( + "unpackPly", + args, + )) as { + packedArray: Uint32Array; + numSplats: number; + extra: Record; + }; + return { packedArray, numSplats, extra }; + }); + } + case SplatFileType.SPZ: { + return await withWorker(async (worker) => { + const { packedArray, numSplats, extra } = (await worker.call( + "decodeSpz", + { + fileBytes, + }, + )) as { + packedArray: Uint32Array; + numSplats: number; + extra: Record; + }; + return { packedArray, numSplats, extra }; + }); + } + case SplatFileType.SPLAT: { + return await withWorker(async (worker) => { + const { packedArray, numSplats } = (await worker.call( + "decodeAntiSplat", + { + fileBytes, + }, + )) as { packedArray: Uint32Array; numSplats: number }; + return { packedArray, numSplats }; + }); + } + case SplatFileType.KSPLAT: + return await withWorker(async (worker) => { + const { packedArray, numSplats, extra } = (await worker.call( + "decodeKsplat", + { fileBytes }, + )) as { + packedArray: Uint32Array; + numSplats: number; + extra: Record; + }; + return { packedArray, numSplats, extra }; + }); + default: { + throw new Error(`Unknown splat file type: ${splatFileType}`); + } } - if (splatFileType === SplatFileType.SPLAT) { - return await withWorker(async (worker) => { - const { packedArray, numSplats } = (await worker.call("decodeAntiSplat", { - fileBytes, - })) as { packedArray: Uint32Array; numSplats: number }; - return { packedArray, numSplats }; - }); - } - if (splatFileType === SplatFileType.KSPLAT) { - return await withWorker(async (worker) => { - const { packedArray, numSplats, extra } = (await worker.call( - "decodeKsplat", - { fileBytes }, - )) as { - packedArray: Uint32Array; - numSplats: number; - extra: Record; - }; - return { packedArray, numSplats, extra }; - }); - } - throw new Error("Unknown splat file type"); } + +export class SplatData { + numSplats: number; + maxSplats: number; + centers: Float32Array; + scales: Float32Array; + quaternions: Float32Array; + opacities: Float32Array; + colors: Float32Array; + sh1?: Float32Array; + sh2?: Float32Array; + sh3?: Float32Array; + + constructor({ maxSplats = 1 }: { maxSplats?: number } = {}) { + this.numSplats = 0; + this.maxSplats = getTextureSize(maxSplats).maxSplats; + this.centers = new Float32Array(this.maxSplats * 3); + this.scales = new Float32Array(this.maxSplats * 3); + this.quaternions = new Float32Array(this.maxSplats * 4); + this.opacities = new Float32Array(this.maxSplats); + this.colors = new Float32Array(this.maxSplats * 3); + } + + pushSplat(): number { + const index = this.numSplats; + this.ensureIndex(index); + this.numSplats += 1; + return index; + } + + unpushSplat(index: number) { + if (index === this.numSplats - 1) { + this.numSplats -= 1; + } else { + throw new Error("Cannot unpush splat from non-last position"); + } + } + + ensureCapacity(numSplats: number) { + if (numSplats > this.maxSplats) { + const targetSplats = Math.max(numSplats, this.maxSplats * 2); + const newCenters = new Float32Array(targetSplats * 3); + const newScales = new Float32Array(targetSplats * 3); + const newQuaternions = new Float32Array(targetSplats * 4); + const newOpacities = new Float32Array(targetSplats); + const newColors = new Float32Array(targetSplats * 3); + newCenters.set(this.centers); + newScales.set(this.scales); + newQuaternions.set(this.quaternions); + newOpacities.set(this.opacities); + newColors.set(this.colors); + this.centers = newCenters; + this.scales = newScales; + this.quaternions = newQuaternions; + this.opacities = newOpacities; + this.colors = newColors; + + if (this.sh1) { + const newSh1 = new Float32Array(targetSplats * 9); + newSh1.set(this.sh1); + this.sh1 = newSh1; + } + if (this.sh2) { + const newSh2 = new Float32Array(targetSplats * 15); + newSh2.set(this.sh2); + this.sh2 = newSh2; + } + if (this.sh3) { + const newSh3 = new Float32Array(targetSplats * 21); + newSh3.set(this.sh3); + this.sh3 = newSh3; + } + + this.maxSplats = targetSplats; + // console.log("Reallocated capacity", this.maxSplats); + } + } + + ensureIndex(index: number) { + this.ensureCapacity(index + 1); + } + + setCenter(index: number, x: number, y: number, z: number) { + this.centers[index * 3] = x; + this.centers[index * 3 + 1] = y; + this.centers[index * 3 + 2] = z; + } + + setScale(index: number, scaleX: number, scaleY: number, scaleZ: number) { + this.scales[index * 3] = scaleX; + this.scales[index * 3 + 1] = scaleY; + this.scales[index * 3 + 2] = scaleZ; + } + + setQuaternion(index: number, x: number, y: number, z: number, w: number) { + this.quaternions[index * 4] = x; + this.quaternions[index * 4 + 1] = y; + this.quaternions[index * 4 + 2] = z; + this.quaternions[index * 4 + 3] = w; + } + + setOpacity(index: number, opacity: number) { + this.opacities[index] = opacity; + } + + setColor(index: number, r: number, g: number, b: number) { + this.colors[index * 3] = r; + this.colors[index * 3 + 1] = g; + this.colors[index * 3 + 2] = b; + } + + setSh1(index: number, sh1: Float32Array) { + if (!this.sh1) { + this.sh1 = new Float32Array(this.maxSplats * 9); + // console.log("setSh1 creating sh1", this.sh1.length); + } + for (let j = 0; j < 9; ++j) { + this.sh1[index * 9 + j] = sh1[j]; + } + } + + setSh2(index: number, sh2: Float32Array) { + if (!this.sh2) { + this.sh2 = new Float32Array(this.maxSplats * 15); + // console.log("setSh2 creating sh2", this.sh2.length); + } + for (let j = 0; j < 15; ++j) { + this.sh2[index * 15 + j] = sh2[j]; + } + } + + setSh3(index: number, sh3: Float32Array) { + if (!this.sh3) { + this.sh3 = new Float32Array(this.maxSplats * 21); + // console.log("setSh3 creating sh3", this.sh3.length); + } + for (let j = 0; j < 21; ++j) { + this.sh3[index * 21 + j] = sh3[j]; + } + } +} + +export async function transcodeSpz( + input: TranscodeSpzInput, +): Promise<{ input: TranscodeSpzInput; fileBytes: Uint8Array }> { + return await withWorker(async (worker) => { + // console.log("withWorker calling transcodeSpz", input); + const result = (await worker.call("transcodeSpz", input)) as { + input: TranscodeSpzInput; + fileBytes: Uint8Array; + }; + return result; + }); +} + +export type FileInput = { + fileBytes: Uint8Array; + fileType?: SplatFileType; + pathOrUrl?: string; + transform?: { translate?: number[]; quaternion?: number[]; scale?: number }; +}; + +export type TranscodeSpzInput = { + inputs: FileInput[]; + maxSh?: number; + clipXyz?: { min: number[]; max: number[] }; + fractionalBits?: number; + opacityThreshold?: number; +}; diff --git a/src/SplatMesh.ts b/src/SplatMesh.ts index d2554dc..6a1241b 100644 --- a/src/SplatMesh.ts +++ b/src/SplatMesh.ts @@ -45,6 +45,8 @@ export type SplatMeshOptions = { // auto-detected (.splat, .ksplat). (default: undefined auto-detects other // formats from file contents) fileType?: SplatFileType; + // File name to use for type detection. (default: undefined) + fileName?: string; // Use an existing PackedSplats object as the source instead of loading from // a file. Can be used to share a collection of Gsplats among multiple SplatMeshes // (default: undefined creates a new empty PackedSplats or decoded from a @@ -224,12 +226,14 @@ export class SplatMesh extends SplatGenerator { } async asyncInitialize(options: SplatMeshOptions) { - const { url, fileBytes, fileType, maxSplats, constructSplats } = options; + const { url, fileBytes, fileType, fileName, maxSplats, constructSplats } = + options; if (url || fileBytes || constructSplats) { const packedSplatsOptions = { url, fileBytes, fileType, + fileName, maxSplats, construct: constructSplats, }; @@ -295,7 +299,8 @@ export class SplatMesh extends SplatGenerator { this.packedSplats.dispose(); } - constructGenerator({ transform, viewToObject, recolor }: SplatMeshContext) { + constructGenerator(context: SplatMeshContext) { + const { transform, viewToObject, recolor } = context; const generator = dynoBlock( { index: "int" }, { gsplat: Gsplat }, @@ -354,7 +359,7 @@ export class SplatMesh extends SplatGenerator { } // Transform from object to world-space - gsplat = transform.modify(gsplat); + gsplat = transform.applyGsplat(gsplat); // Apply any global recoloring and opacity const recolorRgba = mul(recolor, splitGsplat(gsplat).outputs.rgba); diff --git a/src/antisplat.ts b/src/antisplat.ts new file mode 100644 index 0000000..d24017d --- /dev/null +++ b/src/antisplat.ts @@ -0,0 +1,120 @@ +import { computeMaxSplats, setPackedSplat } from "./utils"; + +export function decodeAntiSplat( + fileBytes: Uint8Array, + initNumSplats: (numSplats: number) => void, + splatCallback: ( + index: number, + x: number, + y: number, + z: number, + scaleX: number, + scaleY: number, + scaleZ: number, + quatX: number, + quatY: number, + quatZ: number, + quatW: number, + opacity: number, + r: number, + g: number, + b: number, + ) => void, +) { + const numSplats = Math.floor(fileBytes.length / 32); // 32 bytes per splat + if (numSplats * 32 !== fileBytes.length) { + throw new Error("Invalid .splat file size"); + } + initNumSplats(numSplats); + + const f32 = new Float32Array(fileBytes.buffer); + for (let i = 0; i < numSplats; ++i) { + const i32 = i * 32; + const i8 = i * 8; + const x = f32[i8 + 0]; + const y = f32[i8 + 1]; + const z = f32[i8 + 2]; + const scaleX = f32[i8 + 3]; + const scaleY = f32[i8 + 4]; + const scaleZ = f32[i8 + 5]; + const r = fileBytes[i32 + 24] / 255; + const g = fileBytes[i32 + 25] / 255; + const b = fileBytes[i32 + 26] / 255; + const opacity = fileBytes[i32 + 27] / 255; + const quatW = (fileBytes[i32 + 28] - 128) / 128; + const quatX = (fileBytes[i32 + 29] - 128) / 128; + const quatY = (fileBytes[i32 + 30] - 128) / 128; + const quatZ = (fileBytes[i32 + 31] - 128) / 128; + splatCallback( + i, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ); + } +} + +export function unpackAntiSplat(fileBytes: Uint8Array): { + packedArray: Uint32Array; + numSplats: number; +} { + let numSplats = 0; + let maxSplats = 0; + let packedArray = new Uint32Array(0); + decodeAntiSplat( + fileBytes, + (cbNumSplats) => { + numSplats = cbNumSplats; + maxSplats = computeMaxSplats(numSplats); + packedArray = new Uint32Array(maxSplats * 4); + }, + ( + index, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ) => { + setPackedSplat( + packedArray, + index, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ); + }, + ); + return { packedArray, numSplats }; +} diff --git a/src/controls.ts b/src/controls.ts index e7bcfd2..bbb2e32 100644 --- a/src/controls.ts +++ b/src/controls.ts @@ -341,6 +341,7 @@ export class PointerControls { rotateSpeed: number; slideSpeed: number; scrollSpeed: number; + swapRotateSlide: boolean; reverseRotate: boolean; reverseSlide: boolean; reverseSwipe: boolean; @@ -385,6 +386,8 @@ export class PointerControls { slideSpeed, // Speed of movement when using mouse scroll wheel (default DEFAULT_SCROLL_SPEED) scrollSpeed, + // Swap the direction of rotation and sliding (default: false) + swapRotateSlide, // Reverse the direction of rotation (default: false) reverseRotate, // Reverse the direction of sliding (default: false) @@ -404,6 +407,7 @@ export class PointerControls { rotateSpeed?: number; slideSpeed?: number; scrollSpeed?: number; + swapRotateSlide?: boolean; reverseRotate?: boolean; reverseSlide?: boolean; reverseSwipe?: boolean; @@ -419,6 +423,7 @@ export class PointerControls { this.rotateSpeed = rotateSpeed ?? DEFAULT_ROTATE_SPEED; this.slideSpeed = slideSpeed ?? DEFAULT_SLIDE_SPEED; this.scrollSpeed = scrollSpeed ?? DEFAULT_SCROLL_SPEED; + this.swapRotateSlide = swapRotateSlide ?? false; this.reverseRotate = reverseRotate ?? false; this.reverseSlide = reverseSlide ?? false; this.reverseSwipe = reverseSwipe ?? false; @@ -446,7 +451,15 @@ export class PointerControls { // Determine if we're starting a rotation pointer action const isRotate = - !this.rotating && (event.pointerType !== "mouse" || event.button === 0); + (!this.swapRotateSlide && + !this.rotating && + (event.pointerType !== "mouse" || event.button === 0)) || + (this.swapRotateSlide && + this.sliding && + !this.rotating && + (event.pointerType !== "mouse" || event.button === 1)); + // const isRotate = + // !this.rotating && (event.pointerType !== "mouse" || event.button === 0); const { pointerId, timeStamp } = event; if (isRotate) { diff --git a/src/dyno/transform.ts b/src/dyno/transform.ts index 97fcdd2..15146df 100644 --- a/src/dyno/transform.ts +++ b/src/dyno/transform.ts @@ -18,8 +18,8 @@ export const transformPos = ( return new TransformPosition({ position, scale, scales, rotate, translate }) .outputs.position; }; -export const transformRay = ( - ray: DynoVal<"vec3">, +export const transformDir = ( + dir: DynoVal<"vec3">, { scale, scales, @@ -30,7 +30,7 @@ export const transformRay = ( rotate?: DynoVal<"vec4">; }, ): DynoVal<"vec3"> => { - return new TransformRay({ ray, scale, scales, rotate }).outputs.ray; + return new TransformDir({ dir, scale, scales, rotate }).outputs.dir; }; export const transformQuat = ( quaternion: DynoVal<"vec4">, @@ -90,36 +90,36 @@ export class TransformPosition extends Dyno< } } -export class TransformRay extends Dyno< - { ray: "vec3"; scale: "float"; scales: "vec3"; rotate: "vec4" }, - { ray: "vec3" } +export class TransformDir extends Dyno< + { dir: "vec3"; scale: "float"; scales: "vec3"; rotate: "vec4" }, + { dir: "vec3" } > { constructor({ - ray, + dir, scale, scales, rotate, }: { - ray?: DynoVal<"vec3">; + dir?: DynoVal<"vec3">; scale?: DynoVal<"float">; scales?: DynoVal<"vec3">; rotate?: DynoVal<"vec4">; }) { super({ - inTypes: { ray: "vec3", scale: "float", scales: "vec3", rotate: "vec4" }, - outTypes: { ray: "vec3" }, - inputs: { ray, scale, scales, rotate }, + inTypes: { dir: "vec3", scale: "float", scales: "vec3", rotate: "vec4" }, + outTypes: { dir: "vec3" }, + inputs: { dir, scale, scales, rotate }, statements: ({ inputs, outputs }) => { - const { ray } = outputs; - if (!ray) { + const { dir } = outputs; + if (!dir) { return []; } const { scale, scales, rotate } = inputs; return [ - `${ray} = ${inputs.ray ?? "vec3(0.0, 0.0, 0.0)"};`, - !scale ? null : `${ray} *= ${scale};`, - !scales ? null : `${ray} *= ${scales};`, - !rotate ? null : `${ray} = quatVec(${rotate}, ${ray});`, + `${dir} = ${inputs.dir ?? "vec3(0.0, 0.0, 0.0)"};`, + !scale ? null : `${dir} *= ${scale};`, + !scales ? null : `${dir} *= ${scales};`, + !rotate ? null : `${dir} = quatVec(${rotate}, ${dir});`, ].filter(Boolean) as string[]; }, }); diff --git a/src/generators/snow.ts b/src/generators/snow.ts index 971b430..9174588 100644 --- a/src/generators/snow.ts +++ b/src/generators/snow.ts @@ -231,7 +231,7 @@ export function snowBox({ rgb, opacity: dynoOpacity, }); - gsplat = transformer.modify(gsplat); + gsplat = transformer.applyGsplat(gsplat); return { gsplat }; }, { diff --git a/src/generators/static.ts b/src/generators/static.ts index e909be2..cba3c62 100644 --- a/src/generators/static.ts +++ b/src/generators/static.ts @@ -99,7 +99,7 @@ export function staticBox({ quaternion, rgba, }); - gsplat = transformer.modify(gsplat); + gsplat = transformer.applyGsplat(gsplat); return { gsplat }; }, { diff --git a/src/index.ts b/src/index.ts index 594d274..2cd4649 100644 --- a/src/index.ts +++ b/src/index.ts @@ -10,7 +10,7 @@ export { getSplatFileType, } from "./SplatLoader"; export { PlyReader } from "./ply"; -export { SpzReader } from "./spz"; +export { SpzReader, SpzWriter, transcodeSpz } from "./spz"; export { PackedSplats, type PackedSplatsOptions } from "./PackedSplats"; export { diff --git a/src/ksplat.ts b/src/ksplat.ts new file mode 100644 index 0000000..d034a3e --- /dev/null +++ b/src/ksplat.ts @@ -0,0 +1,576 @@ +import { + computeMaxSplats, + encodeSh1Rgb, + encodeSh2Rgb, + encodeSh3Rgb, + fromHalf, + setPackedSplat, +} from "./utils"; + +type KsplatCompression = { + bytesPerCenter: number; + bytesPerScale: number; + bytesPerRotation: number; + bytesPerColor: number; + bytesPerSphericalHarmonicsComponent: number; + scaleOffsetBytes: number; + rotationOffsetBytes: number; + colorOffsetBytes: number; + sphericalHarmonicsOffsetBytes: number; + scaleRange: number; +}; + +const KSPLAT_COMPRESSION: Record = { + 0: { + bytesPerCenter: 12, + bytesPerScale: 12, + bytesPerRotation: 16, + bytesPerColor: 4, + bytesPerSphericalHarmonicsComponent: 4, + scaleOffsetBytes: 12, + rotationOffsetBytes: 24, + colorOffsetBytes: 40, + sphericalHarmonicsOffsetBytes: 44, + scaleRange: 1, + }, + 1: { + bytesPerCenter: 6, + bytesPerScale: 6, + bytesPerRotation: 8, + bytesPerColor: 4, + bytesPerSphericalHarmonicsComponent: 2, + scaleOffsetBytes: 6, + rotationOffsetBytes: 12, + colorOffsetBytes: 20, + sphericalHarmonicsOffsetBytes: 24, + scaleRange: 32767, + }, + 2: { + bytesPerCenter: 6, + bytesPerScale: 6, + bytesPerRotation: 8, + bytesPerColor: 4, + bytesPerSphericalHarmonicsComponent: 1, + scaleOffsetBytes: 6, + rotationOffsetBytes: 12, + colorOffsetBytes: 20, + sphericalHarmonicsOffsetBytes: 24, + scaleRange: 32767, + }, +}; + +const KSPLAT_SH_DEGREE_TO_COMPONENTS: Record = { + 0: 0, + 1: 9, + 2: 24, + 3: 45, +}; + +export function decodeKsplat( + fileBytes: Uint8Array, + initNumSplats: (numSplats: number) => void, + splatCallback: ( + index: number, + x: number, + y: number, + z: number, + scaleX: number, + scaleY: number, + scaleZ: number, + quatX: number, + quatY: number, + quatZ: number, + quatW: number, + opacity: number, + r: number, + g: number, + b: number, + ) => void, + shCallback?: ( + index: number, + sh1: Float32Array, + sh2?: Float32Array, + sh3?: Float32Array, + ) => void, +) { + const HEADER_BYTES = 4096; + const SECTION_BYTES = 1024; + + let headerOffset = 0; + const header = new DataView(fileBytes.buffer, headerOffset, HEADER_BYTES); + headerOffset += HEADER_BYTES; + + const versionMajor = header.getUint8(0); + const versionMinor = header.getUint8(1); + if (versionMajor !== 0 || versionMinor < 1) { + throw new Error( + `Unsupported .ksplat version: ${versionMajor}.${versionMinor}`, + ); + } + const maxSectionCount = header.getUint32(4, true); + // const sectionCount = header.getUint32(8, true); + // const maxSplatCount = header.getUint32(12, true); + const splatCount = header.getUint32(16, true); + const compressionLevel = header.getUint16(20, true); + if (compressionLevel < 0 || compressionLevel > 2) { + throw new Error(`Invalid .ksplat compression level: ${compressionLevel}`); + } + // const sceneCenterX = header.getFloat32(24, true); + // const sceneCenterY = header.getFloat32(28, true); + // const sceneCenterZ = header.getFloat32(32, true); + const minSphericalHarmonicsCoeff = header.getFloat32(36, true) || -1.5; + const maxSphericalHarmonicsCoeff = header.getFloat32(40, true) || 1.5; + + const numSplats = splatCount; + initNumSplats(numSplats); + const maxSplats = computeMaxSplats(numSplats); + const packedArray = new Uint32Array(maxSplats * 4); + const extra: Record = {}; + + let sectionBase = HEADER_BYTES + maxSectionCount * SECTION_BYTES; + + for (let section = 0; section < maxSectionCount; ++section) { + const section = new DataView(fileBytes.buffer, headerOffset, SECTION_BYTES); + headerOffset += SECTION_BYTES; + + const sectionSplatCount = section.getUint32(0, true); + const sectionMaxSplatCount = section.getUint32(4, true); + const bucketSize = section.getUint32(8, true); + const bucketCount = section.getUint32(12, true); + const bucketBlockSize = section.getFloat32(16, true); + const bucketStorageSizeBytes = section.getUint16(20, true); + const compressionScaleRange = + (section.getUint32(24, true) || + KSPLAT_COMPRESSION[compressionLevel]?.scaleRange) ?? + 1; + // const fullBucketCount = section.getUint32(32, true); + const partiallyFilledBucketCount = section.getUint32(36, true); + const bucketsMetaDataSizeBytes = partiallyFilledBucketCount * 4; + const bucketsStorageSizeBytes = + bucketStorageSizeBytes * bucketCount + bucketsMetaDataSizeBytes; + const sphericalHarmonicsDegree = section.getUint16(40, true); + const shComponents = + KSPLAT_SH_DEGREE_TO_COMPONENTS[sphericalHarmonicsDegree]; + + const { + bytesPerCenter, + bytesPerScale, + bytesPerRotation, + bytesPerColor, + bytesPerSphericalHarmonicsComponent, + scaleOffsetBytes, + rotationOffsetBytes, + colorOffsetBytes, + sphericalHarmonicsOffsetBytes, + } = KSPLAT_COMPRESSION[compressionLevel]; + const bytesPerSplat = + bytesPerCenter + + bytesPerScale + + bytesPerRotation + + bytesPerColor + + shComponents * bytesPerSphericalHarmonicsComponent; + const splatDataStorageSizeBytes = bytesPerSplat * sectionMaxSplatCount; + const storageSizeBytes = + splatDataStorageSizeBytes + bucketsStorageSizeBytes; + + const sh1Index = [0, 3, 6, 1, 4, 7, 2, 5, 8]; + const sh2Index = [ + 9, 14, 19, 10, 15, 20, 11, 16, 21, 12, 17, 22, 13, 18, 23, + ]; + const sh3Index = [ + 24, 31, 38, 25, 32, 39, 26, 33, 40, 27, 34, 41, 28, 35, 42, 29, 36, 43, + 30, 37, 44, + ]; + const sh1 = + sphericalHarmonicsDegree >= 1 ? new Float32Array(3 * 3) : undefined; + const sh2 = + sphericalHarmonicsDegree >= 2 ? new Float32Array(5 * 3) : undefined; + const sh3 = + sphericalHarmonicsDegree >= 3 ? new Float32Array(7 * 3) : undefined; + + const compressionScaleFactor = bucketBlockSize / 2 / compressionScaleRange; + const bucketsBase = sectionBase + bucketsMetaDataSizeBytes; + const dataBase = sectionBase + bucketsStorageSizeBytes; + const data = new DataView( + fileBytes.buffer, + dataBase, + splatDataStorageSizeBytes, + ); + const bucketArray = new Float32Array( + fileBytes.buffer, + bucketsBase, + bucketCount * 3, + ); + + function getSh(splatOffset: number, component: number) { + if (compressionLevel === 0) { + return data.getFloat32( + splatOffset + sphericalHarmonicsOffsetBytes + component * 4, + true, + ); + } + if (compressionLevel === 1) { + return fromHalf( + data.getUint16( + splatOffset + sphericalHarmonicsOffsetBytes + component * 2, + true, + ), + ); + } + const t = + data.getUint8(splatOffset + sphericalHarmonicsOffsetBytes + component) / + 255; + return ( + minSphericalHarmonicsCoeff + + t * (maxSphericalHarmonicsCoeff - minSphericalHarmonicsCoeff) + ); + } + + for (let i = 0; i < sectionSplatCount; ++i) { + const splatOffset = i * bytesPerSplat; + const bucketIndex = Math.floor(i / bucketSize); + + const x = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 0, true) + : (data.getUint16(splatOffset + 0, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 0]; + const y = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 4, true) + : (data.getUint16(splatOffset + 2, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 1]; + const z = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 8, true) + : (data.getUint16(splatOffset + 4, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 2]; + + const scaleX = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 0, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 0, true)); + const scaleY = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 4, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 2, true)); + const scaleZ = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 8, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 4, true)); + + const quatW = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 0, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 0, true), + ); + const quatX = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 4, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 2, true), + ); + const quatY = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 8, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 4, true), + ); + const quatZ = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 12, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 6, true), + ); + + const r = data.getUint8(splatOffset + colorOffsetBytes + 0) / 255; + const g = data.getUint8(splatOffset + colorOffsetBytes + 1) / 255; + const b = data.getUint8(splatOffset + colorOffsetBytes + 2) / 255; + const opacity = data.getUint8(splatOffset + colorOffsetBytes + 3) / 255; + + splatCallback( + i, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ); + + if (sphericalHarmonicsDegree >= 1 && sh1) { + shCallback?.(i, sh1, sh2, sh3); + } + } + sectionBase += storageSizeBytes; + } +} + +export function unpackKsplat(fileBytes: Uint8Array): { + packedArray: Uint32Array; + numSplats: number; + extra: Record; +} { + const HEADER_BYTES = 4096; + const SECTION_BYTES = 1024; + + let headerOffset = 0; + const header = new DataView(fileBytes.buffer, headerOffset, HEADER_BYTES); + headerOffset += HEADER_BYTES; + + const versionMajor = header.getUint8(0); + const versionMinor = header.getUint8(1); + if (versionMajor !== 0 || versionMinor < 1) { + throw new Error( + `Unsupported .ksplat version: ${versionMajor}.${versionMinor}`, + ); + } + const maxSectionCount = header.getUint32(4, true); + // const sectionCount = header.getUint32(8, true); + // const maxSplatCount = header.getUint32(12, true); + const splatCount = header.getUint32(16, true); + const compressionLevel = header.getUint16(20, true); + if (compressionLevel < 0 || compressionLevel > 2) { + throw new Error(`Invalid .ksplat compression level: ${compressionLevel}`); + } + // const sceneCenterX = header.getFloat32(24, true); + // const sceneCenterY = header.getFloat32(28, true); + // const sceneCenterZ = header.getFloat32(32, true); + const minSphericalHarmonicsCoeff = header.getFloat32(36, true) || -1.5; + const maxSphericalHarmonicsCoeff = header.getFloat32(40, true) || 1.5; + + const numSplats = splatCount; + const maxSplats = computeMaxSplats(numSplats); + const packedArray = new Uint32Array(maxSplats * 4); + const extra: Record = {}; + + let sectionBase = HEADER_BYTES + maxSectionCount * SECTION_BYTES; + + for (let section = 0; section < maxSectionCount; ++section) { + const section = new DataView(fileBytes.buffer, headerOffset, SECTION_BYTES); + headerOffset += SECTION_BYTES; + + const sectionSplatCount = section.getUint32(0, true); + const sectionMaxSplatCount = section.getUint32(4, true); + const bucketSize = section.getUint32(8, true); + const bucketCount = section.getUint32(12, true); + const bucketBlockSize = section.getFloat32(16, true); + const bucketStorageSizeBytes = section.getUint16(20, true); + const compressionScaleRange = + (section.getUint32(24, true) || + KSPLAT_COMPRESSION[compressionLevel]?.scaleRange) ?? + 1; + // const fullBucketCount = section.getUint32(32, true); + const partiallyFilledBucketCount = section.getUint32(36, true); + const bucketsMetaDataSizeBytes = partiallyFilledBucketCount * 4; + const bucketsStorageSizeBytes = + bucketStorageSizeBytes * bucketCount + bucketsMetaDataSizeBytes; + const sphericalHarmonicsDegree = section.getUint16(40, true); + const shComponents = + KSPLAT_SH_DEGREE_TO_COMPONENTS[sphericalHarmonicsDegree]; + + const { + bytesPerCenter, + bytesPerScale, + bytesPerRotation, + bytesPerColor, + bytesPerSphericalHarmonicsComponent, + scaleOffsetBytes, + rotationOffsetBytes, + colorOffsetBytes, + sphericalHarmonicsOffsetBytes, + } = KSPLAT_COMPRESSION[compressionLevel]; + const bytesPerSplat = + bytesPerCenter + + bytesPerScale + + bytesPerRotation + + bytesPerColor + + shComponents * bytesPerSphericalHarmonicsComponent; + const splatDataStorageSizeBytes = bytesPerSplat * sectionMaxSplatCount; + const storageSizeBytes = + splatDataStorageSizeBytes + bucketsStorageSizeBytes; + + const sh1Index = [0, 3, 6, 1, 4, 7, 2, 5, 8]; + const sh2Index = [ + 9, 14, 19, 10, 15, 20, 11, 16, 21, 12, 17, 22, 13, 18, 23, + ]; + const sh3Index = [ + 24, 31, 38, 25, 32, 39, 26, 33, 40, 27, 34, 41, 28, 35, 42, 29, 36, 43, + 30, 37, 44, + ]; + const sh1 = + sphericalHarmonicsDegree >= 1 ? new Float32Array(3 * 3) : undefined; + const sh2 = + sphericalHarmonicsDegree >= 2 ? new Float32Array(5 * 3) : undefined; + const sh3 = + sphericalHarmonicsDegree >= 3 ? new Float32Array(7 * 3) : undefined; + + const compressionScaleFactor = bucketBlockSize / 2 / compressionScaleRange; + const bucketsBase = sectionBase + bucketsMetaDataSizeBytes; + const dataBase = sectionBase + bucketsStorageSizeBytes; + const data = new DataView( + fileBytes.buffer, + dataBase, + splatDataStorageSizeBytes, + ); + const bucketArray = new Float32Array( + fileBytes.buffer, + bucketsBase, + bucketCount * 3, + ); + + function getSh(splatOffset: number, component: number) { + if (compressionLevel === 0) { + return data.getFloat32( + splatOffset + sphericalHarmonicsOffsetBytes + component * 4, + true, + ); + } + if (compressionLevel === 1) { + return fromHalf( + data.getUint16( + splatOffset + sphericalHarmonicsOffsetBytes + component * 2, + true, + ), + ); + } + const t = + data.getUint8(splatOffset + sphericalHarmonicsOffsetBytes + component) / + 255; + return ( + minSphericalHarmonicsCoeff + + t * (maxSphericalHarmonicsCoeff - minSphericalHarmonicsCoeff) + ); + } + + for (let i = 0; i < sectionSplatCount; ++i) { + const splatOffset = i * bytesPerSplat; + const bucketIndex = Math.floor(i / bucketSize); + + const x = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 0, true) + : (data.getUint16(splatOffset + 0, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 0]; + const y = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 4, true) + : (data.getUint16(splatOffset + 2, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 1]; + const z = + compressionLevel === 0 + ? data.getFloat32(splatOffset + 8, true) + : (data.getUint16(splatOffset + 4, true) - compressionScaleRange) * + compressionScaleFactor + + bucketArray[3 * bucketIndex + 2]; + + const scaleX = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 0, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 0, true)); + const scaleY = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 4, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 2, true)); + const scaleZ = + compressionLevel === 0 + ? data.getFloat32(splatOffset + scaleOffsetBytes + 8, true) + : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 4, true)); + + const quatW = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 0, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 0, true), + ); + const quatX = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 4, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 2, true), + ); + const quatY = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 8, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 4, true), + ); + const quatZ = + compressionLevel === 0 + ? data.getFloat32(splatOffset + rotationOffsetBytes + 12, true) + : fromHalf( + data.getUint16(splatOffset + rotationOffsetBytes + 6, true), + ); + + const r = data.getUint8(splatOffset + colorOffsetBytes + 0) / 255; + const g = data.getUint8(splatOffset + colorOffsetBytes + 1) / 255; + const b = data.getUint8(splatOffset + colorOffsetBytes + 2) / 255; + const opacity = data.getUint8(splatOffset + colorOffsetBytes + 3) / 255; + + setPackedSplat( + packedArray, + i, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ); + + if (sphericalHarmonicsDegree >= 1) { + if (sh1) { + if (!extra.sh1) { + extra.sh1 = new Uint32Array(numSplats * 2); + } + for (const [i, key] of sh1Index.entries()) { + sh1[i] = getSh(splatOffset, key); + } + encodeSh1Rgb(extra.sh1 as Uint32Array, i, sh1); + } + if (sh2) { + if (!extra.sh2) { + extra.sh2 = new Uint32Array(numSplats * 4); + } + for (const [i, key] of sh2Index.entries()) { + sh2[i] = getSh(splatOffset, key); + } + encodeSh2Rgb(extra.sh2 as Uint32Array, i, sh2); + } + if (sh3) { + if (!extra.sh3) { + extra.sh3 = new Uint32Array(numSplats * 4); + } + for (const [i, key] of sh3Index.entries()) { + sh3[i] = getSh(splatOffset, key); + } + encodeSh3Rgb(extra.sh3 as Uint32Array, i, sh3); + } + } + } + sectionBase += storageSizeBytes; + } + return { packedArray, numSplats, extra }; +} diff --git a/src/splatConstructors.ts b/src/splatConstructors.ts index 8122bce..b9bac1a 100644 --- a/src/splatConstructors.ts +++ b/src/splatConstructors.ts @@ -214,6 +214,8 @@ export function textSplats({ textAlign, // line spacing multiplier, lines delimited by "\n" (default: 1.0) lineHeight, + // Coordinate scale in object-space (default: 1.0) + objectScale, }: { text: string; font?: string; @@ -223,6 +225,7 @@ export function textSplats({ dotRadius?: number; textAlign?: "left" | "center" | "right" | "start" | "end"; lineHeight?: number; + objectScale?: number; }) { font = font ?? "Arial"; fontSize = fontSize ?? 32; @@ -230,6 +233,7 @@ export function textSplats({ dotRadius = dotRadius ?? 0.8; textAlign = textAlign ?? "start"; lineHeight = lineHeight ?? 1; + objectScale = objectScale ?? 1; const lines = text.split("\n"); const canvas = document.createElement("canvas"); @@ -276,7 +280,7 @@ export function textSplats({ const rgba = new Uint8Array(imageData.data.buffer); const splats = new PackedSplats(); const center = new THREE.Vector3(); - const scales = new THREE.Vector3().setScalar(dotRadius); + const scales = new THREE.Vector3().setScalar(dotRadius * objectScale); const quaternion = new THREE.Quaternion(0, 0, 0, 1); rgb = rgb ?? new THREE.Color(1, 1, 1); @@ -287,6 +291,7 @@ export function textSplats({ if (a > 0) { const opacity = a / 255; center.set(x - 0.5 * (width - 1), 0.5 * (height - 1) - y, 0); + center.multiplyScalar(objectScale); splats.pushSplat(center, scales, quaternion, opacity, rgb); } offset += 4; diff --git a/src/spz.ts b/src/spz.ts index 0648d19..2de5b4a 100644 --- a/src/spz.ts +++ b/src/spz.ts @@ -1,4 +1,18 @@ -import { GunzipReader, fromHalf } from "./utils"; +import * as THREE from "three"; +import { + SplatData, + SplatFileType, + type TranscodeSpzInput, + getSplatFileType, + getSplatFileTypeFromPath, +} from "./SplatLoader"; +import { GunzipReader, fromHalf, unpackSplat } from "./utils"; + +import init_wasm, { decode_wlg } from "forge-internal-rs"; +import { decodeAntiSplat } from "./antisplat"; +import { SPLAT_TEX_HEIGHT, SPLAT_TEX_WIDTH } from "./defines"; +import { decodeKsplat } from "./ksplat"; +import { PlyReader } from "./ply"; // SPZ file format reader @@ -39,16 +53,16 @@ export class SpzReader { } parseSplats( - centerCallback: (index: number, x: number, y: number, z: number) => void, - alphaCallback: (index: number, alpha: number) => void, - rgbCallback: (index: number, r: number, g: number, b: number) => void, - scalesCallback: ( + centerCallback?: (index: number, x: number, y: number, z: number) => void, + alphaCallback?: (index: number, alpha: number) => void, + rgbCallback?: (index: number, r: number, g: number, b: number) => void, + scalesCallback?: ( index: number, scaleX: number, scaleY: number, scaleZ: number, ) => void, - quatCallback: ( + quatCallback?: ( index: number, quatX: number, quatY: number, @@ -76,7 +90,7 @@ export class SpzReader { const x = fromHalf(centerUint16[i3]); const y = fromHalf(centerUint16[i3 + 1]); const z = fromHalf(centerUint16[i3 + 2]); - centerCallback(i, x, y, z); + centerCallback?.(i, x, y, z); } } else if (this.version === 2) { // 24-bit fixed-point centers @@ -102,7 +116,7 @@ export class SpzReader { (centerBytes[i9 + 6] << 8)) >> 8) / fixed; - centerCallback(i, x, y, z); + centerCallback?.(i, x, y, z); } } else { throw new Error("Unreachable"); @@ -111,7 +125,7 @@ export class SpzReader { { const bytes = this.reader.read(this.numSplats); for (let i = 0; i < this.numSplats; i++) { - alphaCallback(i, bytes[i] / 255); + alphaCallback?.(i, bytes[i] / 255); } } { @@ -122,7 +136,7 @@ export class SpzReader { const r = (rgbBytes[i3] / 255 - 0.5) * scale + 0.5; const g = (rgbBytes[i3 + 1] / 255 - 0.5) * scale + 0.5; const b = (rgbBytes[i3 + 2] / 255 - 0.5) * scale + 0.5; - rgbCallback(i, r, g, b); + rgbCallback?.(i, r, g, b); } } { @@ -132,7 +146,7 @@ export class SpzReader { const scaleX = Math.exp(scalesBytes[i3] / 16 - 10); const scaleY = Math.exp(scalesBytes[i3 + 1] / 16 - 10); const scaleZ = Math.exp(scalesBytes[i3 + 2] / 16 - 10); - scalesCallback(i, scaleX, scaleY, scaleZ); + scalesCallback?.(i, scaleX, scaleY, scaleZ); } } { @@ -145,7 +159,7 @@ export class SpzReader { const quatW = Math.sqrt( Math.max(0, 1 - quatX * quatX - quatY * quatY - quatZ * quatZ), ); - quatCallback(i, quatX, quatY, quatZ, quatW); + quatCallback?.(i, quatX, quatY, quatZ, quatW); } } @@ -175,7 +189,7 @@ export class SpzReader { } offset += 21; } - shCallback(i, sh1, sh2, sh3); + shCallback?.(i, sh1, sh2, sh3); } } } @@ -183,3 +197,590 @@ export class SpzReader { const SH_DEGREE_TO_VECS: Record = { 1: 3, 2: 8, 3: 15 }; const SH_C0 = 0.28209479177387814; + +export const SPZ_MAGIC = 0x5053474e; // NGSP = Niantic gaussian splat +export const SPZ_VERSION = 2; +export const FLAG_ANTIALIASED = 0x1; + +export class SpzWriter { + buffer: ArrayBuffer; + view: DataView; + numSplats: number; + shDegree: number; + fractionalBits: number; + fraction: number; + flagAntiAlias: boolean; + clippedCount = 0; + + constructor({ + numSplats, + shDegree, + fractionalBits = 12, + flagAntiAlias = true, + }: { + numSplats: number; + shDegree: number; + fractionalBits?: number; + flagAntiAlias?: boolean; + }) { + const splatSize = + 9 + + 1 + + 3 + + 3 + + 3 + + (shDegree >= 1 ? 9 : 0) + + (shDegree >= 2 ? 15 : 0) + + (shDegree >= 3 ? 21 : 0); + const bufferSize = 16 + numSplats * splatSize; + this.buffer = new ArrayBuffer(bufferSize); + this.view = new DataView(this.buffer); + + this.view.setUint32(0, SPZ_MAGIC, true); // NGSP + this.view.setUint32(4, SPZ_VERSION, true); + this.view.setUint32(8, numSplats, true); + this.view.setUint8(12, shDegree); + this.view.setUint8(13, fractionalBits); + this.view.setUint8(14, flagAntiAlias ? FLAG_ANTIALIASED : 0); + this.view.setUint8(15, 0); // Reserved + + this.numSplats = numSplats; + this.shDegree = shDegree; + this.fractionalBits = fractionalBits; + this.fraction = 1 << fractionalBits; + this.flagAntiAlias = flagAntiAlias; + } + + setCenter(index: number, x: number, y: number, z: number) { + // Divide by this.fraction and round to nearest integer, + // then write as 3-bytes per x then y then z. + const xRounded = Math.round(x * this.fraction); + const xInt = Math.max(-0x7fffff, Math.min(0x7fffff, xRounded)); + const yRounded = Math.round(y * this.fraction); + const yInt = Math.max(-0x7fffff, Math.min(0x7fffff, yRounded)); + const zRounded = Math.round(z * this.fraction); + const zInt = Math.max(-0x7fffff, Math.min(0x7fffff, zRounded)); + const clipped = xRounded !== xInt || yRounded !== yInt || zRounded !== zInt; + if (clipped) { + this.clippedCount += 1; + // if (this.clippedCount < 10) { + // // Write x y z also in hex + // console.log(`Clipped ${index}: ${x}, ${y}, ${z} (0x${x.toString(16)}, 0x${y.toString(16)}, 0x${z.toString(16)}) -> ${xRounded}, ${yRounded}, ${zRounded} (0x${xRounded.toString(16)}, 0x${yRounded.toString(16)}, 0x${zRounded.toString(16)}) -> ${xInt}, ${yInt}, ${zInt} (0x${xInt.toString(16)}, 0x${yInt.toString(16)}, 0x${zInt.toString(16)})`); + // } + } + const i9 = index * 9; + const base = 16 + i9; + this.view.setUint8(base, xInt & 0xff); + this.view.setUint8(base + 1, (xInt >> 8) & 0xff); + this.view.setUint8(base + 2, (xInt >> 16) & 0xff); + this.view.setUint8(base + 3, yInt & 0xff); + this.view.setUint8(base + 4, (yInt >> 8) & 0xff); + this.view.setUint8(base + 5, (yInt >> 16) & 0xff); + this.view.setUint8(base + 6, zInt & 0xff); + this.view.setUint8(base + 7, (zInt >> 8) & 0xff); + this.view.setUint8(base + 8, (zInt >> 16) & 0xff); + } + + setAlpha(index: number, alpha: number) { + const base = 16 + this.numSplats * 9 + index; + this.view.setUint8( + base, + Math.max(0, Math.min(255, Math.round(alpha * 255))), + ); + } + + static scaleRgb(r: number) { + const v = ((r - 0.5) / (SH_C0 / 0.15) + 0.5) * 255; + return Math.max(0, Math.min(255, Math.round(v))); + } + + setRgb(index: number, r: number, g: number, b: number) { + const base = 16 + this.numSplats * 10 + index * 3; + this.view.setUint8(base, SpzWriter.scaleRgb(r)); + this.view.setUint8(base + 1, SpzWriter.scaleRgb(g)); + this.view.setUint8(base + 2, SpzWriter.scaleRgb(b)); + } + + setScale(index: number, scaleX: number, scaleY: number, scaleZ: number) { + const base = 16 + this.numSplats * 13 + index * 3; + this.view.setUint8( + base, + Math.max(0, Math.min(255, Math.round((Math.log(scaleX) + 10) * 16))), + ); + this.view.setUint8( + base + 1, + Math.max(0, Math.min(255, Math.round((Math.log(scaleY) + 10) * 16))), + ); + this.view.setUint8( + base + 2, + Math.max(0, Math.min(255, Math.round((Math.log(scaleZ) + 10) * 16))), + ); + } + + setQuat( + index: number, + quatX: number, + quatY: number, + quatZ: number, + quatW: number, + ) { + const base = 16 + this.numSplats * 16 + index * 3; + const quatNeg = quatW < 0; + this.view.setUint8( + base, + Math.max( + 0, + Math.min(255, Math.round(((quatNeg ? -quatX : quatX) + 1) * 127.5)), + ), + ); + this.view.setUint8( + base + 1, + Math.max( + 0, + Math.min(255, Math.round(((quatNeg ? -quatY : quatY) + 1) * 127.5)), + ), + ); + this.view.setUint8( + base + 2, + Math.max( + 0, + Math.min(255, Math.round(((quatNeg ? -quatZ : quatZ) + 1) * 127.5)), + ), + ); + } + + static quantizeSh(sh: number, bits: number) { + const value = Math.round(sh * 128) + 128; + const bucketSize = 1 << (8 - bits); + const quantized = + Math.floor((value + bucketSize / 2) / bucketSize) * bucketSize; + return Math.max(0, Math.min(255, quantized)); + } + + setSh( + index: number, + sh1: Float32Array, + sh2?: Float32Array, + sh3?: Float32Array, + ) { + const shVecs = SH_DEGREE_TO_VECS[this.shDegree] || 0; + const base1 = 16 + this.numSplats * 19 + index * shVecs * 3; + for (let j = 0; j < 9; ++j) { + this.view.setUint8(base1 + j, SpzWriter.quantizeSh(sh1[j], 5)); + } + if (sh2) { + const base2 = base1 + 9; + for (let j = 0; j < 15; ++j) { + this.view.setUint8(base2 + j, SpzWriter.quantizeSh(sh2[j], 4)); + } + if (sh3) { + const base3 = base2 + 15; + for (let j = 0; j < 21; ++j) { + this.view.setUint8(base3 + j, SpzWriter.quantizeSh(sh3[j], 4)); + } + } + } + } + + async finalize(): Promise { + const input = new Uint8Array(this.buffer); + const stream = new ReadableStream({ + async start(controller) { + controller.enqueue(input); + controller.close(); + }, + }); + const compressed = stream.pipeThrough(new CompressionStream("gzip")); + const response = new Response(compressed); + const buffer = await response.arrayBuffer(); + console.log( + "Compressed", + input.length, + "bytes to", + buffer.byteLength, + "bytes", + ); + return new Uint8Array(buffer); + } +} + +export async function transcodeSpz(input: TranscodeSpzInput) { + const splats = new SplatData(); + const { + inputs, + clipXyz, + maxSh, + fractionalBits = 12, + opacityThreshold, + } = input; + for (const input of inputs) { + const scale = input.transform?.scale ?? 1; + const quaternion = new THREE.Quaternion().fromArray( + input.transform?.quaternion ?? [0, 0, 0, 1], + ); + const translate = new THREE.Vector3().fromArray( + input.transform?.translate ?? [0, 0, 0], + ); + const clip = clipXyz + ? new THREE.Box3( + new THREE.Vector3().fromArray(clipXyz.min), + new THREE.Vector3().fromArray(clipXyz.max), + ) + : undefined; + + function transformPos(pos: THREE.Vector3) { + pos.multiplyScalar(scale); + pos.applyQuaternion(quaternion); + pos.add(translate); + return pos; + } + + function transformScales(scales: THREE.Vector3) { + scales.multiplyScalar(scale); + return scales; + } + + function transformQuaternion(quat: THREE.Quaternion) { + quat.premultiply(quaternion); + return quat; + } + + function withinClip(p: THREE.Vector3) { + return !clip || clip.containsPoint(p); + } + + function withinOpacity(opacity: number) { + return opacityThreshold !== undefined + ? opacity >= opacityThreshold + : true; + } + + let fileType = input.fileType; + if (!fileType) { + fileType = getSplatFileType(input.fileBytes); + if (!fileType && input.pathOrUrl) { + fileType = getSplatFileTypeFromPath(input.pathOrUrl); + } + } + switch (fileType) { + case SplatFileType.WLG0: { + await init_wasm(); + const decoded = decode_wlg( + input.fileBytes, + SPLAT_TEX_WIDTH, + SPLAT_TEX_HEIGHT, + ) as { numSplats: number; packedSplats: Uint32Array }; + for (let i = 0; i < decoded.numSplats; ++i) { + let { center, scales, quaternion, opacity, color } = unpackSplat( + decoded.packedSplats, + i, + ); + center = transformPos(center); + if (withinClip(center) && withinOpacity(opacity)) { + const index = splats.pushSplat(); + splats.setCenter(index, center.x, center.y, center.z); + scales = transformScales(scales); + splats.setScale(index, scales.x, scales.y, scales.z); + quaternion = transformQuaternion(quaternion); + splats.setQuaternion( + index, + quaternion.x, + quaternion.y, + quaternion.z, + quaternion.w, + ); + splats.setOpacity(index, opacity); + splats.setColor(index, color.r, color.g, color.b); + } + } + break; + } + case SplatFileType.PLY: { + const ply = new PlyReader({ fileBytes: input.fileBytes }); + await ply.parseHeader(); + let lastIndex: number | null = null; + ply.parseSplats( + ( + index, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ) => { + const center = transformPos(new THREE.Vector3(x, y, z)); + if (withinClip(center) && withinOpacity(opacity)) { + lastIndex = splats.pushSplat(); + splats.setCenter(lastIndex, center.x, center.y, center.z); + const scales = transformScales( + new THREE.Vector3(scaleX, scaleY, scaleZ), + ); + splats.setScale(lastIndex, scales.x, scales.y, scales.z); + const quaternion = transformQuaternion( + new THREE.Quaternion(quatX, quatY, quatZ, quatW), + ); + splats.setQuaternion( + lastIndex, + quaternion.x, + quaternion.y, + quaternion.z, + quaternion.w, + ); + splats.setOpacity(lastIndex, opacity); + splats.setColor(lastIndex, r, g, b); + } else { + lastIndex = null; + } + }, + (index, sh1, sh2, sh3) => { + if (sh1 && lastIndex !== null) { + splats.setSh1(lastIndex, sh1); + } + if (sh2 && lastIndex !== null) { + splats.setSh2(lastIndex, sh2); + } + if (sh3 && lastIndex !== null) { + splats.setSh3(lastIndex, sh3); + } + }, + ); + break; + } + case SplatFileType.SPZ: { + const spz = new SpzReader({ fileBytes: input.fileBytes }); + const mapping = new Int32Array(spz.numSplats); + mapping.fill(-1); + const centers = new Float32Array(spz.numSplats * 3); + const center = new THREE.Vector3(); + spz.parseSplats( + (index, x, y, z) => { + const center = transformPos(new THREE.Vector3(x, y, z)); + centers[index * 3] = center.x; + centers[index * 3 + 1] = center.y; + centers[index * 3 + 2] = center.z; + }, + (index, alpha) => { + center.fromArray(centers, index * 3); + if (withinClip(center) && withinOpacity(alpha)) { + mapping[index] = splats.pushSplat(); + splats.setCenter(mapping[index], center.x, center.y, center.z); + splats.setOpacity(mapping[index], alpha); + } + }, + (index, r, g, b) => { + if (mapping[index] >= 0) { + splats.setColor(mapping[index], r, g, b); + } + }, + (index, scaleX, scaleY, scaleZ) => { + if (mapping[index] >= 0) { + const scales = transformScales( + new THREE.Vector3(scaleX, scaleY, scaleZ), + ); + splats.setScale(mapping[index], scales.x, scales.y, scales.z); + } + }, + (index, quatX, quatY, quatZ, quatW) => { + if (mapping[index] >= 0) { + const quaternion = transformQuaternion( + new THREE.Quaternion(quatX, quatY, quatZ, quatW), + ); + splats.setQuaternion( + mapping[index], + quaternion.x, + quaternion.y, + quaternion.z, + quaternion.w, + ); + } + }, + (index, sh1, sh2, sh3) => { + if (mapping[index] >= 0) { + splats.setSh1(mapping[index], sh1); + if (sh2) { + splats.setSh2(mapping[index], sh2); + } + if (sh3) { + splats.setSh3(mapping[index], sh3); + } + } + }, + ); + break; + } + case SplatFileType.SPLAT: + decodeAntiSplat( + input.fileBytes, + (numSplats) => {}, + ( + index, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ) => { + const center = transformPos(new THREE.Vector3(x, y, z)); + if (withinClip(center) && withinOpacity(opacity)) { + const index = splats.pushSplat(); + splats.setCenter(index, center.x, center.y, center.z); + const scales = transformScales( + new THREE.Vector3(scaleX, scaleY, scaleZ), + ); + splats.setScale(index, scales.x, scales.y, scales.z); + const quaternion = transformQuaternion( + new THREE.Quaternion(quatX, quatY, quatZ, quatW), + ); + splats.setQuaternion( + index, + quaternion.x, + quaternion.y, + quaternion.z, + quaternion.w, + ); + splats.setOpacity(index, opacity); + splats.setColor(index, r, g, b); + } + }, + ); + break; + case SplatFileType.KSPLAT: { + let lastIndex: number | null = null; + decodeKsplat( + input.fileBytes, + (numSplats) => {}, + ( + index, + x, + y, + z, + scaleX, + scaleY, + scaleZ, + quatX, + quatY, + quatZ, + quatW, + opacity, + r, + g, + b, + ) => { + const center = transformPos(new THREE.Vector3(x, y, z)); + if (withinClip(center) && withinOpacity(opacity)) { + lastIndex = splats.pushSplat(); + splats.setCenter(lastIndex, center.x, center.y, center.z); + const scales = transformScales( + new THREE.Vector3(scaleX, scaleY, scaleZ), + ); + splats.setScale(lastIndex, scales.x, scales.y, scales.z); + const quaternion = transformQuaternion( + new THREE.Quaternion(quatX, quatY, quatZ, quatW), + ); + splats.setQuaternion( + lastIndex, + quaternion.x, + quaternion.y, + quaternion.z, + quaternion.w, + ); + splats.setOpacity(lastIndex, opacity); + splats.setColor(lastIndex, r, g, b); + } else { + lastIndex = null; + } + }, + (index, sh1, sh2, sh3) => { + if (lastIndex !== null) { + splats.setSh1(lastIndex, sh1); + if (sh2) { + splats.setSh2(lastIndex, sh2); + } + if (sh3) { + splats.setSh3(lastIndex, sh3); + } + } + }, + ); + break; + } + default: + throw new Error(`transcodeSpz not implemented for ${fileType}`); + } + } + + const shDegree = Math.min( + maxSh ?? 3, + splats.sh3 ? 3 : splats.sh2 ? 2 : splats.sh1 ? 1 : 0, + ); + const spz = new SpzWriter({ + numSplats: splats.numSplats, + shDegree, + fractionalBits, + flagAntiAlias: true, + }); + + for (let i = 0; i < splats.numSplats; ++i) { + const i3 = i * 3; + const i4 = i * 4; + spz.setCenter( + i, + splats.centers[i3], + splats.centers[i3 + 1], + splats.centers[i3 + 2], + ); + spz.setScale( + i, + splats.scales[i3], + splats.scales[i3 + 1], + splats.scales[i3 + 2], + ); + spz.setQuat( + i, + splats.quaternions[i4], + splats.quaternions[i4 + 1], + splats.quaternions[i4 + 2], + splats.quaternions[i4 + 3], + ); + spz.setAlpha(i, splats.opacities[i]); + spz.setRgb( + i, + splats.colors[i3], + splats.colors[i3 + 1], + splats.colors[i3 + 2], + ); + if (splats.sh1 && shDegree >= 1) { + spz.setSh( + i, + splats.sh1.slice(i * 9, (i + 1) * 9), + shDegree >= 2 && splats.sh2 + ? splats.sh2.slice(i * 15, (i + 1) * 15) + : undefined, + shDegree >= 3 && splats.sh3 + ? splats.sh3.slice(i * 21, (i + 1) * 21) + : undefined, + ); + } + } + + const spzBytes = await spz.finalize(); + return { fileBytes: spzBytes, clippedCount: spz.clippedCount }; +} diff --git a/src/utils.ts b/src/utils.ts index a01141c..9cf2a9a 100644 --- a/src/utils.ts +++ b/src/utils.ts @@ -624,6 +624,19 @@ export function getTextureSize(numSplats: number): { return { width, height, depth, maxSplats }; } +export function computeMaxSplats(numSplats: number): number { + // Compute the size of a Gsplat array texture (2048x2048xD) that can fit + // numSplats splats, and return the total number of splats that can be stored + // in such a texture. + const width = SPLAT_TEX_WIDTH; + const height = Math.max( + SPLAT_TEX_MIN_HEIGHT, + Math.min(SPLAT_TEX_HEIGHT, Math.ceil(numSplats / width)), + ); + const depth = Math.ceil(numSplats / (width * height)); + return width * height * depth; +} + // Heuristic function to determine if we are running on a mobile device. export function isMobile(): boolean { if (navigator.maxTouchPoints > 0) { diff --git a/src/worker.ts b/src/worker.ts index 3dda89e..c5e62ec 100644 --- a/src/worker.ts +++ b/src/worker.ts @@ -3,20 +3,22 @@ import init_wasm, { old_sort_splats, sort_splats, } from "forge-internal-rs"; +import type { TranscodeSpzInput } from "./SplatLoader"; +import { unpackAntiSplat } from "./antisplat"; import { SCALE_MIN, SPLAT_TEX_HEIGHT, - SPLAT_TEX_MIN_HEIGHT, SPLAT_TEX_WIDTH, WASM_SPLAT_SORT, } from "./defines"; +import { unpackKsplat } from "./ksplat"; import { PlyReader } from "./ply"; -import { SpzReader } from "./spz"; +import { SpzReader, transcodeSpz } from "./spz"; import { + computeMaxSplats, encodeSh1Rgb, encodeSh2Rgb, encodeSh3Rgb, - fromHalf, getArrayBuffers, setPackedSplat, setPackedSplatCenter, @@ -161,6 +163,16 @@ async function onMessage(event: MessageEvent) { } break; } + case "transcodeSpz": { + const input = args as TranscodeSpzInput; + const spzBytes = await transcodeSpz(input); + result = { + id, + fileBytes: spzBytes, + input, + }; + break; + } default: { throw new Error(`Unknown name: ${name}`); } @@ -254,19 +266,6 @@ async function unpackPly({ return { packedArray, numSplats, extra }; } -function computeMaxSplats(numSplats: number): number { - // Compute the size of a Gsplat array texture (2048x2048xD) that can fit - // numSplats splats, and return the total number of splats that can be stored - // in such a texture. - const width = SPLAT_TEX_WIDTH; - const height = Math.max( - SPLAT_TEX_MIN_HEIGHT, - Math.min(SPLAT_TEX_HEIGHT, Math.ceil(numSplats / width)), - ); - const depth = Math.ceil(numSplats / (width * height)); - return width * height * depth; -} - function unpackSpz(fileBytes: Uint8Array): { packedArray: Uint32Array; numSplats: number; @@ -318,373 +317,6 @@ function unpackSpz(fileBytes: Uint8Array): { return { packedArray, numSplats, extra }; } -function unpackAntiSplat(fileBytes: Uint8Array): { - packedArray: Uint32Array; - numSplats: number; -} { - const numSplats = Math.floor(fileBytes.length / 32); // 32 bytes per splat - if (numSplats * 32 !== fileBytes.length) { - throw new Error("Invalid .splat file size"); - } - const maxSplats = computeMaxSplats(numSplats); - const packedArray = new Uint32Array(maxSplats * 4); - const f32 = new Float32Array(fileBytes.buffer); - - for (let i = 0; i < numSplats; ++i) { - const i32 = i * 32; - const i8 = i * 8; - const x = f32[i8 + 0]; - const y = f32[i8 + 1]; - const z = f32[i8 + 2]; - const scaleX = f32[i8 + 3]; - const scaleY = f32[i8 + 4]; - const scaleZ = f32[i8 + 5]; - const r = fileBytes[i32 + 24] / 255; - const g = fileBytes[i32 + 25] / 255; - const b = fileBytes[i32 + 26] / 255; - const opacity = fileBytes[i32 + 27] / 255; - const quatW = (fileBytes[i32 + 28] - 128) / 128; - const quatX = (fileBytes[i32 + 29] - 128) / 128; - const quatY = (fileBytes[i32 + 30] - 128) / 128; - const quatZ = (fileBytes[i32 + 31] - 128) / 128; - setPackedSplat( - packedArray, - i, - x, - y, - z, - scaleX, - scaleY, - scaleZ, - quatX, - quatY, - quatZ, - quatW, - opacity, - r, - g, - b, - ); - } - return { packedArray, numSplats }; -} - -type KsplatCompression = { - bytesPerCenter: number; - bytesPerScale: number; - bytesPerRotation: number; - bytesPerColor: number; - bytesPerSphericalHarmonicsComponent: number; - scaleOffsetBytes: number; - rotationOffsetBytes: number; - colorOffsetBytes: number; - sphericalHarmonicsOffsetBytes: number; - scaleRange: number; -}; - -const KSPLAT_COMPRESSION: Record = { - 0: { - bytesPerCenter: 12, - bytesPerScale: 12, - bytesPerRotation: 16, - bytesPerColor: 4, - bytesPerSphericalHarmonicsComponent: 4, - scaleOffsetBytes: 12, - rotationOffsetBytes: 24, - colorOffsetBytes: 40, - sphericalHarmonicsOffsetBytes: 44, - scaleRange: 1, - }, - 1: { - bytesPerCenter: 6, - bytesPerScale: 6, - bytesPerRotation: 8, - bytesPerColor: 4, - bytesPerSphericalHarmonicsComponent: 2, - scaleOffsetBytes: 6, - rotationOffsetBytes: 12, - colorOffsetBytes: 20, - sphericalHarmonicsOffsetBytes: 24, - scaleRange: 32767, - }, - 2: { - bytesPerCenter: 6, - bytesPerScale: 6, - bytesPerRotation: 8, - bytesPerColor: 4, - bytesPerSphericalHarmonicsComponent: 1, - scaleOffsetBytes: 6, - rotationOffsetBytes: 12, - colorOffsetBytes: 20, - sphericalHarmonicsOffsetBytes: 24, - scaleRange: 32767, - }, -}; - -const KSPLAT_SH_DEGREE_TO_COMPONENTS: Record = { - 0: 0, - 1: 9, - 2: 24, - 3: 45, -}; - -function unpackKsplat(fileBytes: Uint8Array): { - packedArray: Uint32Array; - numSplats: number; - extra: Record; -} { - const HEADER_BYTES = 4096; - const SECTION_BYTES = 1024; - - let headerOffset = 0; - const header = new DataView(fileBytes.buffer, headerOffset, HEADER_BYTES); - headerOffset += HEADER_BYTES; - - const versionMajor = header.getUint8(0); - const versionMinor = header.getUint8(1); - if (versionMajor !== 0 || versionMinor < 1) { - throw new Error( - `Unsupported .ksplat version: ${versionMajor}.${versionMinor}`, - ); - } - const maxSectionCount = header.getUint32(4, true); - // const sectionCount = header.getUint32(8, true); - // const maxSplatCount = header.getUint32(12, true); - const splatCount = header.getUint32(16, true); - const compressionLevel = header.getUint16(20, true); - if (compressionLevel < 0 || compressionLevel > 2) { - throw new Error(`Invalid .ksplat compression level: ${compressionLevel}`); - } - // const sceneCenterX = header.getFloat32(24, true); - // const sceneCenterY = header.getFloat32(28, true); - // const sceneCenterZ = header.getFloat32(32, true); - const minSphericalHarmonicsCoeff = header.getFloat32(36, true) || -1.5; - const maxSphericalHarmonicsCoeff = header.getFloat32(40, true) || 1.5; - - const numSplats = splatCount; - const maxSplats = computeMaxSplats(numSplats); - const packedArray = new Uint32Array(maxSplats * 4); - const extra: Record = {}; - - let sectionBase = HEADER_BYTES + maxSectionCount * SECTION_BYTES; - - for (let section = 0; section < maxSectionCount; ++section) { - const section = new DataView(fileBytes.buffer, headerOffset, SECTION_BYTES); - headerOffset += SECTION_BYTES; - - const sectionSplatCount = section.getUint32(0, true); - const sectionMaxSplatCount = section.getUint32(4, true); - const bucketSize = section.getUint32(8, true); - const bucketCount = section.getUint32(12, true); - const bucketBlockSize = section.getFloat32(16, true); - const bucketStorageSizeBytes = section.getUint16(20, true); - const compressionScaleRange = - (section.getUint32(24, true) || - KSPLAT_COMPRESSION[compressionLevel]?.scaleRange) ?? - 1; - // const fullBucketCount = section.getUint32(32, true); - const partiallyFilledBucketCount = section.getUint32(36, true); - const bucketsMetaDataSizeBytes = partiallyFilledBucketCount * 4; - const bucketsStorageSizeBytes = - bucketStorageSizeBytes * bucketCount + bucketsMetaDataSizeBytes; - const sphericalHarmonicsDegree = section.getUint16(40, true); - const shComponents = - KSPLAT_SH_DEGREE_TO_COMPONENTS[sphericalHarmonicsDegree]; - - const { - bytesPerCenter, - bytesPerScale, - bytesPerRotation, - bytesPerColor, - bytesPerSphericalHarmonicsComponent, - scaleOffsetBytes, - rotationOffsetBytes, - colorOffsetBytes, - sphericalHarmonicsOffsetBytes, - } = KSPLAT_COMPRESSION[compressionLevel]; - const bytesPerSplat = - bytesPerCenter + - bytesPerScale + - bytesPerRotation + - bytesPerColor + - shComponents * bytesPerSphericalHarmonicsComponent; - const splatDataStorageSizeBytes = bytesPerSplat * sectionMaxSplatCount; - const storageSizeBytes = - splatDataStorageSizeBytes + bucketsStorageSizeBytes; - - const sh1Index = [0, 3, 6, 1, 4, 7, 2, 5, 8]; - const sh2Index = [ - 9, 14, 19, 10, 15, 20, 11, 16, 21, 12, 17, 22, 13, 18, 23, - ]; - const sh3Index = [ - 24, 31, 38, 25, 32, 39, 26, 33, 40, 27, 34, 41, 28, 35, 42, 29, 36, 43, - 30, 37, 44, - ]; - const sh1 = - sphericalHarmonicsDegree >= 1 ? new Float32Array(3 * 3) : undefined; - const sh2 = - sphericalHarmonicsDegree >= 2 ? new Float32Array(5 * 3) : undefined; - const sh3 = - sphericalHarmonicsDegree >= 3 ? new Float32Array(7 * 3) : undefined; - - const compressionScaleFactor = bucketBlockSize / 2 / compressionScaleRange; - const bucketsBase = sectionBase + bucketsMetaDataSizeBytes; - const dataBase = sectionBase + bucketsStorageSizeBytes; - const data = new DataView( - fileBytes.buffer, - dataBase, - splatDataStorageSizeBytes, - ); - const bucketArray = new Float32Array( - fileBytes.buffer, - bucketsBase, - bucketCount * 3, - ); - - function getSh(splatOffset: number, component: number) { - if (compressionLevel === 0) { - return data.getFloat32( - splatOffset + sphericalHarmonicsOffsetBytes + component * 4, - true, - ); - } - if (compressionLevel === 1) { - return fromHalf( - data.getUint16( - splatOffset + sphericalHarmonicsOffsetBytes + component * 2, - true, - ), - ); - } - const t = - data.getUint8(splatOffset + sphericalHarmonicsOffsetBytes + component) / - 255; - return ( - minSphericalHarmonicsCoeff + - t * (maxSphericalHarmonicsCoeff - minSphericalHarmonicsCoeff) - ); - } - - for (let i = 0; i < sectionSplatCount; ++i) { - const splatOffset = i * bytesPerSplat; - const bucketIndex = Math.floor(i / bucketSize); - - const x = - compressionLevel === 0 - ? data.getFloat32(splatOffset + 0, true) - : (data.getUint16(splatOffset + 0, true) - compressionScaleRange) * - compressionScaleFactor + - bucketArray[3 * bucketIndex + 0]; - const y = - compressionLevel === 0 - ? data.getFloat32(splatOffset + 4, true) - : (data.getUint16(splatOffset + 2, true) - compressionScaleRange) * - compressionScaleFactor + - bucketArray[3 * bucketIndex + 1]; - const z = - compressionLevel === 0 - ? data.getFloat32(splatOffset + 8, true) - : (data.getUint16(splatOffset + 4, true) - compressionScaleRange) * - compressionScaleFactor + - bucketArray[3 * bucketIndex + 2]; - - const scaleX = - compressionLevel === 0 - ? data.getFloat32(splatOffset + scaleOffsetBytes + 0, true) - : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 0, true)); - const scaleY = - compressionLevel === 0 - ? data.getFloat32(splatOffset + scaleOffsetBytes + 4, true) - : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 2, true)); - const scaleZ = - compressionLevel === 0 - ? data.getFloat32(splatOffset + scaleOffsetBytes + 8, true) - : fromHalf(data.getUint16(splatOffset + scaleOffsetBytes + 4, true)); - - const quatW = - compressionLevel === 0 - ? data.getFloat32(splatOffset + rotationOffsetBytes + 0, true) - : fromHalf( - data.getUint16(splatOffset + rotationOffsetBytes + 0, true), - ); - const quatX = - compressionLevel === 0 - ? data.getFloat32(splatOffset + rotationOffsetBytes + 4, true) - : fromHalf( - data.getUint16(splatOffset + rotationOffsetBytes + 2, true), - ); - const quatY = - compressionLevel === 0 - ? data.getFloat32(splatOffset + rotationOffsetBytes + 8, true) - : fromHalf( - data.getUint16(splatOffset + rotationOffsetBytes + 4, true), - ); - const quatZ = - compressionLevel === 0 - ? data.getFloat32(splatOffset + rotationOffsetBytes + 12, true) - : fromHalf( - data.getUint16(splatOffset + rotationOffsetBytes + 6, true), - ); - - const r = data.getUint8(splatOffset + colorOffsetBytes + 0) / 255; - const g = data.getUint8(splatOffset + colorOffsetBytes + 1) / 255; - const b = data.getUint8(splatOffset + colorOffsetBytes + 2) / 255; - const opacity = data.getUint8(splatOffset + colorOffsetBytes + 3) / 255; - - setPackedSplat( - packedArray, - i, - x, - y, - z, - scaleX, - scaleY, - scaleZ, - quatX, - quatY, - quatZ, - quatW, - opacity, - r, - g, - b, - ); - - if (sphericalHarmonicsDegree >= 1) { - if (sh1) { - if (!extra.sh1) { - extra.sh1 = new Uint32Array(numSplats * 2); - } - for (const [i, key] of sh1Index.entries()) { - sh1[i] = getSh(splatOffset, key); - } - encodeSh1Rgb(extra.sh1 as Uint32Array, i, sh1); - } - if (sh2) { - if (!extra.sh2) { - extra.sh2 = new Uint32Array(numSplats * 4); - } - for (const [i, key] of sh2Index.entries()) { - sh2[i] = getSh(splatOffset, key); - } - encodeSh2Rgb(extra.sh2 as Uint32Array, i, sh2); - } - if (sh3) { - if (!extra.sh3) { - extra.sh3 = new Uint32Array(numSplats * 4); - } - for (const [i, key] of sh3Index.entries()) { - sh3[i] = getSh(splatOffset, key); - } - encodeSh3Rgb(extra.sh3 as Uint32Array, i, sh3); - } - } - } - sectionBase += storageSizeBytes; - } - return { packedArray, numSplats, extra }; -} - // Array of buckets for sorting float16 distances with range [0, DEPTH_INFINITY]. const DEPTH_INFINITY = 0x7c00; const DEPTH_SIZE = DEPTH_INFINITY + 1; diff --git a/vite.config.ts b/vite.config.ts index f257ba2..77296c1 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -73,6 +73,14 @@ export default defineConfig(({ mode }) => { emptyOutDir: isFirstPass, }, + worker: { + plugins: () => [ + glsl({ + include: ["**/*.glsl"], + }), + ], + }, + server: { watch: { usePolling: true,