Implementation PC SOGS file format loading via meta.json root file. Added webp dependency for accurate decoding. Updated viewer and editor to accept new file type. (#73)

This commit is contained in:
Andreas Sundquist
2025-06-30 08:57:51 -07:00
committed by GitHub
parent eb1b5029ce
commit 36c3c61d0b
9 changed files with 389 additions and 41 deletions
+14 -5
View File
@@ -68,7 +68,7 @@
import * as THREE from "three";
import { OrbitControls } from "three/addons/controls/OrbitControls.js";
import { GUI } from "lil-gui";
import { constructGrid, SparkControls, SparkRenderer, SplatMesh, textSplats, dyno, transcodeSpz } from "@sparkjsdev/spark";
import { constructGrid, SparkControls, SparkRenderer, SplatMesh, textSplats, dyno, transcodeSpz, isPcSogs } from "@sparkjsdev/spark";
import { getAssetFileURL } from "/examples/js/get-asset-url.js";
const scene = new THREE.Scene();
@@ -154,7 +154,7 @@
loadFiles(urls);
guiOptions.loadFromText = ""; // Clear after loading
} else {
alert("No valid URLs found in text. URLs must start with http:// or https:// and end with .ply, .spz, .splat, or .ksplat");
alert("No valid URLs found in text. URLs must start with http:// or https:// and end with .ply, .spz, .splat, .ksplat, or .json");
}
}
},
@@ -357,6 +357,7 @@
try {
let fileBytes;
let fileName;
let url = null;
// Check if splatFile is a URL string or a File object
if (typeof splatFile === "string") {
@@ -365,6 +366,10 @@
fileBytes = new Uint8Array(await fetchWithProgress(splatFile));
// Extract filename from URL
fileName = splatFile.split("/").pop().split("?")[0] || "downloaded-file";
if (isPcSogs(fileBytes)) {
url = splatFile;
}
} else {
// It's a File object
fileBytes = new Uint8Array(await splatFile.arrayBuffer());
@@ -375,14 +380,18 @@
writeOptions.filename = fileName.split(".")[0];
}
const splatMesh = new SplatMesh({ fileBytes: fileBytes.slice(), fileName });
const init = url ? { url } : { fileBytes: fileBytes.slice(), fileName };
const splatMesh = new SplatMesh(init);
const translate = guiOptions.loadOffset * index
splatMesh.position.set(translate, 0.5 * translate, 0.1 * translate);
splatMesh.enableWorldToView = true;
splatMesh.worldModifier = makeWorldModifier(splatMesh);
await splatMesh.initialized;
inputs.push({ fileBytes, pathOrUrl: fileName, object: splatMesh });
if (!url) {
// PC SOGS transcode not supported yet
inputs.push({ fileBytes, pathOrUrl: fileName, object: splatMesh });
}
frame.add(splatMesh);
console.log(`Loaded ${fileName} with ${splatMesh.numSplats} splats`);
@@ -659,7 +668,7 @@
// Parse URLs from text (handles single URLs, multiple lines, mixed content)
function parseURLsFromText(text) {
const supportedExtensions = [".ply", ".spz", ".splat", ".ksplat"];
const supportedExtensions = [".ply", ".spz", ".splat", ".ksplat", ".json"];
const urls = [];
// Split by lines, commas, and semicolons
+6 -9
View File
@@ -263,15 +263,17 @@
loadSplatFile(event.target.files[0]);
};
var loadedSplat;
async function loadSplatFile(splatFile) {
const reader = new FileReader();
fileBytes = new Uint8Array(await splatFile.arrayBuffer());
fileName = splatFile.name;
setSplatFile({ fileBytes: fileBytes.slice(), fileName });
}
var loadedSplat;
function setSplatFile(init) {
if (loadedSplat) { scene.remove(loadedSplat); }
loadedSplat = new SplatMesh({ fileBytes: fileBytes.slice(), fileName });
loadedSplat = new SplatMesh(init);
loadedSplat.quaternion.set(1, 0, 0, 0);
scene.add(loadedSplat);
@@ -289,12 +291,7 @@
fileName = splatURL.split("/").pop().split("?")[0];
document.querySelector('.container').classList.add('hidden');
document.querySelector('.canvas-container').classList.remove('invisible');
fetch(splatURL)
.then(res => res.blob())
.then(blob => {
loadSplatFile(blob);
})
.catch(err => console.error('Download failed:', err));
setSplatFile({ url: splatURL });
}
const urlFormEl = document.querySelector('.url-form');
+1
View File
@@ -58,6 +58,7 @@
"spark-internal-rs": "file:rust/spark-internal-rs/pkg"
},
"dependencies": {
"@jsquash/webp": "^1.5.0",
"fflate": "^0.8.2"
},
"keywords": ["3d", "three.js", "gsplats", "3dgs", "gaussian", "splats"]
+7 -14
View File
@@ -1,7 +1,7 @@
import * as THREE from "three";
import type { GsplatGenerator } from "./SplatGenerator";
import { type SplatFileType, unpackSplats } from "./SplatLoader";
import { type SplatFileType, SplatLoader, unpackSplats } from "./SplatLoader";
import { SPLAT_TEX_HEIGHT, SPLAT_TEX_WIDTH } from "./defines";
import {
DynoProgram,
@@ -120,20 +120,12 @@ export class PackedSplats {
}
async asyncInitialize(options: PackedSplatsOptions) {
let { url, fileBytes, construct } = options;
const { url, fileBytes, construct } = options;
if (url) {
fileBytes = await fetch(url).then(async (response) => {
if (!response.ok) {
throw new Error(
`${response.status} "${response.statusText}" fetching URL: ${url}`,
);
}
const arrayBuffer = await response.arrayBuffer();
return arrayBuffer;
});
}
if (fileBytes) {
const loader = new SplatLoader();
loader.packedSplats = this;
await loader.loadAsync(url);
} else if (fileBytes) {
const unpacked = await unpackSplats({
input: fileBytes,
fileType: options.fileType,
@@ -141,6 +133,7 @@ export class PackedSplats {
});
this.initialize(unpacked);
}
if (construct) {
const maybePromise = construct(this);
// If construct returns a promise, wait for it to complete
+167 -7
View File
@@ -14,6 +14,7 @@ import { decompressPartialGzip, getTextureSize } from "./utils";
export class SplatLoader extends Loader {
fileLoader: FileLoader;
fileType?: SplatFileType;
packedSplats?: PackedSplats;
constructor(manager?: LoadingManager) {
super(manager);
@@ -37,12 +38,61 @@ export class SplatLoader extends Loader {
async (response) => {
if (onLoad) {
const input = response as ArrayBuffer;
const decoded = await unpackSplats({
input,
fileType: this.fileType,
pathOrUrl: url,
});
onLoad(new PackedSplats(decoded));
const extraFiles: Record<string, ArrayBuffer> = {};
const promises = [];
let fileType = this.fileType;
try {
const pcSogsJson = tryPcSogs(input);
if (this.fileType === SplatFileType.PCSOGS) {
if (pcSogsJson === undefined) {
throw new Error("Invalid PC SOGS file");
}
}
if (pcSogsJson !== undefined) {
fileType = SplatFileType.PCSOGS;
for (const key of ["means", "scales", "quats", "sh0", "shN"]) {
const prop = pcSogsJson[key as keyof PcSogsJson];
if (prop) {
const files = prop.files;
for (const file of files) {
const fileUrl = new URL(file, url).toString();
this.manager.itemStart(fileUrl);
const promise = this.loadExtra(fileUrl)
.then((data) => {
extraFiles[file] = data;
})
.catch((error) => {
this.manager.itemError(fileUrl);
throw error;
})
.finally(() => {
this.manager.itemEnd(fileUrl);
});
promises.push(promise);
}
}
}
}
await Promise.all(promises);
const decoded = await unpackSplats({
input,
extraFiles,
fileType,
pathOrUrl: url,
});
if (this.packedSplats) {
this.packedSplats.initialize(decoded);
onLoad(this.packedSplats);
} else {
onLoad(new PackedSplats(decoded));
}
} catch (error) {
onError?.(error);
}
}
},
onProgress,
@@ -66,6 +116,17 @@ export class SplatLoader extends Loader {
});
}
async loadExtra(url: string): Promise<ArrayBuffer> {
return new Promise((resolve, reject) => {
this.fileLoader.load(
url,
(response) => resolve(response as ArrayBuffer),
undefined,
(error) => reject(error),
);
});
}
parse(packedSplats: PackedSplats): SplatMesh {
return new SplatMesh({ packedSplats });
}
@@ -76,6 +137,7 @@ export enum SplatFileType {
SPZ = "spz",
SPLAT = "splat",
KSPLAT = "ksplat",
PCSOGS = "pcsogs",
}
export function getSplatFileType(
@@ -133,12 +195,96 @@ export function getSplatFileTypeFromPath(
return undefined;
}
export type PcSogsJson = {
means: {
shape: number[];
dtype: string;
mins: number[];
maxs: number[];
files: string[];
};
scales: {
shape: number[];
dtype: string;
mins: number[];
maxs: number[];
files: string[];
};
quats: { shape: number[]; dtype: string; encoding?: string; files: string[] };
sh0: {
shape: number[];
dtype: string;
mins: number[];
maxs: number[];
files: string[];
};
shN: {
shape: number[];
dtype: string;
mins: number;
maxs: number;
quantization: number;
files: string[];
};
};
export function isPcSogs(input: ArrayBuffer | Uint8Array | string): boolean {
// Returns true if the input seems to be a valid PC SOGS file
return tryPcSogs(input) !== undefined;
}
export function tryPcSogs(
input: ArrayBuffer | Uint8Array | string,
): PcSogsJson | undefined {
// Try to parse input as SOGS JSON and see if it's valid
try {
let text: string;
if (typeof input === "string") {
text = input;
} else {
const fileBytes =
input instanceof ArrayBuffer ? new Uint8Array(input) : input;
if (fileBytes.length > 65536) {
// Should be only a few KB, definitely not a SOGS JSON file
return undefined;
}
text = new TextDecoder().decode(fileBytes);
}
const json = JSON.parse(text);
if (!json || typeof json !== "object" || Array.isArray(json)) {
return undefined;
}
for (const key of ["means", "scales", "quats", "sh0"]) {
if (
!json[key] ||
typeof json[key] !== "object" ||
Array.isArray(json[key])
) {
return undefined;
}
if (!json[key].shape || !json[key].files) {
return undefined;
}
if (key !== "quats" && (!json[key].mins || !json[key].maxs)) {
return undefined;
}
}
// This is probably a PC SOGS file
return json as PcSogsJson;
} catch {
return undefined;
}
}
export async function unpackSplats({
input,
extraFiles,
fileType,
pathOrUrl,
}: {
input: Uint8Array | ArrayBuffer;
extraFiles?: Record<string, ArrayBuffer>;
fileType?: SplatFileType;
pathOrUrl?: string;
}): Promise<{
@@ -201,7 +347,7 @@ export async function unpackSplats({
return { packedArray, numSplats };
});
}
case SplatFileType.KSPLAT:
case SplatFileType.KSPLAT: {
return await withWorker(async (worker) => {
const { packedArray, numSplats, extra } = (await worker.call(
"decodeKsplat",
@@ -213,6 +359,20 @@ export async function unpackSplats({
};
return { packedArray, numSplats, extra };
});
}
case SplatFileType.PCSOGS: {
return await withWorker(async (worker) => {
const { packedArray, numSplats, extra } = (await worker.call(
"decodePcSogs",
{ fileBytes, extraFiles },
)) as {
packedArray: Uint32Array;
numSplats: number;
extra: Record<string, unknown>;
};
return { packedArray, numSplats, extra };
});
}
default: {
throw new Error(`Unknown splat file type: ${splatFileType}`);
}
+1
View File
@@ -8,6 +8,7 @@ export {
unpackSplats,
SplatFileType,
getSplatFileType,
isPcSogs,
} from "./SplatLoader";
export { PlyReader } from "./ply";
export { SpzReader, SpzWriter, transcodeSpz } from "./spz";
+160
View File
@@ -0,0 +1,160 @@
import { decode as decodeWebp } from "@jsquash/webp";
import type { PcSogsJson } from "./SplatLoader";
import {
computeMaxSplats,
encodeSh1Rgb,
encodeSh2Rgb,
encodeSh3Rgb,
setPackedSplatCenter,
setPackedSplatQuat,
setPackedSplatRgba,
setPackedSplatScales,
} from "./utils";
export async function unpackPcSogs(
fileBytes: Uint8Array,
extraFiles: Record<string, ArrayBuffer>,
): Promise<{
packedArray: Uint32Array;
numSplats: number;
extra: Record<string, unknown>;
}> {
const json = JSON.parse(new TextDecoder().decode(fileBytes)) as PcSogsJson;
if (json.quats.encoding !== "quaternion_packed") {
throw new Error("Unsupported quaternion encoding");
}
const numSplats = json.means.shape[0];
const maxSplats = computeMaxSplats(numSplats);
const packedArray = new Uint32Array(maxSplats * 4);
const extra: Record<string, unknown> = {};
const means = await Promise.all([
decodeImageRgba(extraFiles[json.means.files[0]]),
decodeImageRgba(extraFiles[json.means.files[1]]),
]);
for (let i = 0; i < numSplats; ++i) {
const i4 = i * 4;
const fx = (means[0][i4 + 0] + (means[1][i4 + 0] << 8)) / 65535;
const fy = (means[0][i4 + 1] + (means[1][i4 + 1] << 8)) / 65535;
const fz = (means[0][i4 + 2] + (means[1][i4 + 2] << 8)) / 65535;
let x = json.means.mins[0] + (json.means.maxs[0] - json.means.mins[0]) * fx;
let y = json.means.mins[1] + (json.means.maxs[1] - json.means.mins[1]) * fy;
let z = json.means.mins[2] + (json.means.maxs[2] - json.means.mins[2]) * fz;
x = Math.sign(x) * (Math.exp(Math.abs(x)) - 1);
y = Math.sign(y) * (Math.exp(Math.abs(y)) - 1);
z = Math.sign(z) * (Math.exp(Math.abs(z)) - 1);
setPackedSplatCenter(packedArray, i, x, y, z);
}
const scales = await decodeImageRgba(extraFiles[json.scales.files[0]]);
for (let i = 0; i < numSplats; ++i) {
const i4 = i * 4;
const fx = scales[i4 + 0] / 255;
const fy = scales[i4 + 1] / 255;
const fz = scales[i4 + 2] / 255;
const x =
json.scales.mins[0] + (json.scales.maxs[0] - json.scales.mins[0]) * fx;
const y =
json.scales.mins[1] + (json.scales.maxs[1] - json.scales.mins[1]) * fy;
const z =
json.scales.mins[2] + (json.scales.maxs[2] - json.scales.mins[2]) * fz;
setPackedSplatScales(packedArray, i, Math.exp(x), Math.exp(y), Math.exp(z));
}
const quats = await decodeImageRgba(extraFiles[json.quats.files[0]]);
const SQRT2 = Math.sqrt(2);
for (let i = 0; i < numSplats; ++i) {
const i4 = i * 4;
const r0 = (quats[i4 + 0] / 255 - 0.5) * SQRT2;
const r1 = (quats[i4 + 1] / 255 - 0.5) * SQRT2;
const r2 = (quats[i4 + 2] / 255 - 0.5) * SQRT2;
const rr = Math.sqrt(Math.max(0, 1.0 - r0 * r0 - r1 * r1 - r2 * r2));
const rOrder = quats[i4 + 3] - 252;
const quatX = rOrder === 0 ? r0 : rOrder === 1 ? rr : r1;
const quatY = rOrder <= 1 ? r1 : rOrder === 2 ? rr : r2;
const quatZ = rOrder <= 2 ? r2 : rr;
const quatW = rOrder === 0 ? rr : r0;
setPackedSplatQuat(packedArray, i, quatX, quatY, quatZ, quatW);
}
const sh0 = await decodeImageRgba(extraFiles[json.sh0.files[0]]);
const SH_C0 = 0.28209479177387814;
for (let i = 0; i < numSplats; ++i) {
const i4 = i * 4;
const f0 = sh0[i4 + 0] / 255;
const f1 = sh0[i4 + 1] / 255;
const f2 = sh0[i4 + 2] / 255;
const f3 = sh0[i4 + 3] / 255;
const dc0 = json.sh0.mins[0] + (json.sh0.maxs[0] - json.sh0.mins[0]) * f0;
const dc1 = json.sh0.mins[1] + (json.sh0.maxs[1] - json.sh0.mins[1]) * f1;
const dc2 = json.sh0.mins[2] + (json.sh0.maxs[2] - json.sh0.mins[2]) * f2;
const opa = json.sh0.mins[3] + (json.sh0.maxs[3] - json.sh0.mins[3]) * f3;
const r = SH_C0 * dc0 + 0.5;
const g = SH_C0 * dc1 + 0.5;
const b = SH_C0 * dc2 + 0.5;
const a = 1.0 / (1.0 + Math.exp(-opa));
setPackedSplatRgba(packedArray, i, r, g, b, a);
}
if (json.shN) {
extra.sh1 = new Uint32Array(numSplats * 2);
extra.sh2 = new Uint32Array(numSplats * 4);
extra.sh3 = new Uint32Array(numSplats * 4);
const sh1 = new Float32Array(9);
const sh2 = new Float32Array(15);
const sh3 = new Float32Array(21);
const [centroids, labels] = await Promise.all([
decodeImage(extraFiles[json.shN.files[0]]),
decodeImage(extraFiles[json.shN.files[1]]),
]);
for (let i = 0; i < numSplats; ++i) {
const i4 = i * 4;
const label = labels.rgba[i4 + 0] + (labels.rgba[i4 + 1] << 8);
const col = (label & 63) * 15;
const row = label >>> 6;
const offset = row * centroids.width + col;
for (let d = 0; d < 3; ++d) {
for (let k = 0; k < 3; ++k) {
sh1[k * 3 + d] =
json.shN.mins +
((json.shN.maxs - json.shN.mins) *
centroids.rgba[(offset + k) * 4 + d]) /
255;
}
for (let k = 0; k < 5; ++k) {
sh2[k * 3 + d] =
json.shN.mins +
((json.shN.maxs - json.shN.mins) *
centroids.rgba[(offset + 3 + k) * 4 + d]) /
255;
}
for (let k = 0; k < 7; ++k) {
sh3[k * 3 + d] =
json.shN.mins +
((json.shN.maxs - json.shN.mins) *
centroids.rgba[(offset + 8 + k) * 4 + d]) /
255;
}
}
encodeSh1Rgb(extra.sh1 as Uint32Array, i, sh1);
encodeSh2Rgb(extra.sh2 as Uint32Array, i, sh2);
encodeSh3Rgb(extra.sh3 as Uint32Array, i, sh3);
}
}
return { packedArray, numSplats, extra };
}
async function decodeImage(fileBytes: ArrayBuffer) {
const { data: rgba, width, height } = await decodeWebp(fileBytes);
return { rgba, width, height };
}
async function decodeImageRgba(fileBytes: ArrayBuffer) {
const { rgba } = await decodeImage(fileBytes);
return rgba;
}
+17
View File
@@ -504,6 +504,23 @@ export function setPackedSplatQuat(
packedSplats[i4 + 3] = (packedSplats[i4 + 3] & 0x00ffffff) | (uQuatZ << 24);
}
// Encode the RGBA color in the packedSplats Uint32Array, leaving other fields alone.
export function setPackedSplatRgba(
packedSplats: Uint32Array,
index: number,
r: number,
g: number,
b: number,
a: number,
) {
const uR = floatToUint8(r);
const uG = floatToUint8(g);
const uB = floatToUint8(b);
const uA = floatToUint8(a);
const i4 = index * 4;
packedSplats[i4] = uR | (uG << 8) | (uB << 16) | (uA << 24);
}
// Encode the RGB color in the packedSplats Uint32Array, leaving other fields alone.
export function setPackedSplatRgb(
packedSplats: Uint32Array,
+16 -6
View File
@@ -1,13 +1,9 @@
import init_wasm, { sort_splats } from "spark-internal-rs";
import type { TranscodeSpzInput } from "./SplatLoader";
import { unpackAntiSplat } from "./antisplat";
import {
SCALE_MIN,
SPLAT_TEX_HEIGHT,
SPLAT_TEX_WIDTH,
WASM_SPLAT_SORT,
} from "./defines";
import { SCALE_MIN, WASM_SPLAT_SORT } from "./defines";
import { unpackKsplat } from "./ksplat";
import { unpackPcSogs } from "./pcsogs";
import { PlyReader } from "./ply";
import { SpzReader, transcodeSpz } from "./spz";
import {
@@ -85,6 +81,20 @@ async function onMessage(event: MessageEvent) {
};
break;
}
case "decodePcSogs": {
const { fileBytes, extraFiles } = args as {
fileBytes: Uint8Array;
extraFiles: Record<string, ArrayBuffer>;
};
const decoded = await unpackPcSogs(fileBytes, extraFiles);
result = {
id,
numSplats: decoded.numSplats,
packedArray: decoded.packedArray,
extra: decoded.extra,
};
break;
}
case "sortSplats": {
// Sort maxSplats splats using readback data, which encodes one uint32 per
// Gsplats, with the low bytes encoding a float16 distance sort metric.