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,