From cb658935e8c372917c813fad248676bdd4d08c2c Mon Sep 17 00:00:00 2001 From: Noeri Huisman <8823461+mrxz@users.noreply.github.com> Date: Fri, 3 Jul 2026 16:26:27 +0200 Subject: [PATCH 1/2] Refactor get_file to get_image_data in sogs.rs --- rust/spark-lib/src/sogs.rs | 64 +++++++++++++++----------------------- 1 file changed, 25 insertions(+), 39 deletions(-) diff --git a/rust/spark-lib/src/sogs.rs b/rust/spark-lib/src/sogs.rs index 6495d708..d8f843af 100644 --- a/rust/spark-lib/src/sogs.rs +++ b/rust/spark-lib/src/sogs.rs @@ -169,20 +169,21 @@ fn decode_sogs(bytes: &[u8], splats: &mut T, _pathname: Option let mut file_cache: HashMap> = HashMap::new(); preload_all(&meta, &prefix, &mut zip, &mut file_cache)?; - let mut get_file = |name: &str| -> anyhow::Result> { - file_cache.get(name).cloned().ok_or_else(|| anyhow!("Missing file {name} in cache")) + let mut get_image_data = |name: &str| -> anyhow::Result { + let file = file_cache.get(name).cloned().ok_or_else(|| anyhow!("Missing file {name} in cache")); + decode_image(&file?) }; match meta { - PcSogsRoot::V2(v2) => decode_v2(v2, splats, &mut get_file), - PcSogsRoot::V1(v1) => decode_v1(v1, splats, &mut get_file), + PcSogsRoot::V2(v2) => decode_v2(v2, splats, &mut get_image_data), + PcSogsRoot::V1(v1) => decode_v1(v1, splats, &mut get_image_data), } } fn decode_v2( meta: PcSogsV2, splats: &mut T, - get_file: &mut dyn FnMut(&str) -> anyhow::Result>, + get_image_data: &mut dyn FnMut(&str) -> anyhow::Result, ) -> anyhow::Result<()> { let _ = meta.version; let num_splats = meta.count; @@ -191,16 +192,11 @@ fn decode_v2( }).unwrap_or(0); splats.init_splats(&SplatInit { num_splats, max_sh_degree, lod_tree: false })?; - let means0 = decode_rgba(&get_file(&meta.means.files[0])?) - .context("decode means[0]")?; - let means1 = decode_rgba(&get_file(&meta.means.files[1])?) - .context("decode means[1]")?; - let scales_img = decode_rgba(&get_file(&meta.scales.files[0])?) - .context("decode scales")?; - let quats_img = decode_rgba(&get_file(&meta.quats.files[0])?) - .context("decode quats")?; - let sh0_img = decode_rgba(&get_file(&meta.sh0.files[0])?) - .context("decode sh0")?; + let means0 = get_image_data(&meta.means.files[0])?; + let means1 = get_image_data(&meta.means.files[1])?; + let scales_img = get_image_data(&meta.scales.files[0])?; + let quats_img = get_image_data(&meta.quats.files[0])?; + let sh0_img = get_image_data(&meta.sh0.files[0])?; let mut center = vec![0.0f32; num_splats * 3]; let mut scale = vec![0.0f32; num_splats * 3]; @@ -222,7 +218,7 @@ fn decode_v2( if let Some(shn) = meta.shn { decode_shn_v2( shn, - get_file, + get_image_data, num_splats, &mut sh1, &mut sh2, @@ -248,7 +244,7 @@ fn decode_v2( fn decode_v1( meta: PcSogsV1, splats: &mut T, - get_file: &mut dyn FnMut(&str) -> anyhow::Result>, + get_image_data: &mut dyn FnMut(&str) -> anyhow::Result, ) -> anyhow::Result<()> { let num_splats = meta.means.shape[0]; if meta.quats.encoding.as_deref() != Some("quaternion_packed") { @@ -270,16 +266,11 @@ fn decode_v1( splats.init_splats(&SplatInit { num_splats, max_sh_degree, lod_tree: false })?; - let means0 = decode_rgba(&get_file(&meta.means.files[0])?) - .context("decode means[0]")?; - let means1 = decode_rgba(&get_file(&meta.means.files[1])?) - .context("decode means[1]")?; - let scales_img = decode_rgba(&get_file(&meta.scales.files[0])?) - .context("decode scales")?; - let quats_img = decode_rgba(&get_file(&meta.quats.files[0])?) - .context("decode quats")?; - let sh0_img = decode_rgba(&get_file(&meta.sh0.files[0])?) - .context("decode sh0")?; + let means0 = get_image_data(&meta.means.files[0])?; + let means1 = get_image_data(&meta.means.files[1])?; + let scales_img = get_image_data(&meta.scales.files[0])?; + let quats_img = get_image_data(&meta.quats.files[0])?; + let sh0_img = get_image_data(&meta.sh0.files[0])?; let mut center = vec![0.0f32; num_splats * 3]; let mut scale = vec![0.0f32; num_splats * 3]; @@ -298,7 +289,7 @@ fn decode_v1( if let Some(shn) = meta.shn { decode_shn_v1( shn, - get_file, + get_image_data, num_splats, max_sh_degree, &mut sh1, @@ -447,14 +438,14 @@ fn decode_sh0_v1(mins: &[f32; 4], maxs: &[f32; 4], img: &ImageData, out_rgb: &mu fn decode_shn_v2( shn: ShNV2, - get_file: &mut dyn FnMut(&str) -> anyhow::Result>, + get_image_data: &mut dyn FnMut(&str) -> anyhow::Result, num_splats: usize, sh1: &mut [f32], sh2: &mut [f32], sh3: &mut [f32], ) -> anyhow::Result<()> { - let centroids = decode_image(&get_file(&shn.files[0])?)?; - let labels = decode_image(&get_file(&shn.files[1])?)?; + let centroids = get_image_data(&shn.files[0])?; + let labels = get_image_data(&shn.files[1])?; let lookup = shn.codebook; let use_sh1 = shn.bands >= 1; let use_sh2 = shn.bands >= 2; @@ -490,15 +481,15 @@ fn decode_shn_v2( fn decode_shn_v1( shn: ShNV1, - get_file: &mut dyn FnMut(&str) -> anyhow::Result>, + get_image_data: &mut dyn FnMut(&str) -> anyhow::Result, num_splats: usize, max_sh_degree: usize, sh1: &mut [f32], sh2: &mut [f32], sh3: &mut [f32], ) -> anyhow::Result<()> { - let centroids = decode_image(&get_file(&shn.files[0])?)?; - let labels = decode_image(&get_file(&shn.files[1])?)?; + let centroids = get_image_data(&shn.files[0])?; + let labels = get_image_data(&shn.files[1])?; let lookup: Vec = (0..256) .map(|i| shn.mins + (shn.maxs - shn.mins) * (i as f32 / 255.0)) .collect(); @@ -629,11 +620,6 @@ struct ImageData { height: usize, } -fn decode_rgba(bytes: &[u8]) -> anyhow::Result { - let img = decode_image(bytes)?; - Ok(img) -} - fn decode_image(bytes: &[u8]) -> anyhow::Result { let img = ImageReader::new(Cursor::new(bytes)) .with_guessed_format()? From 773f6508cf6854b5aa7fced800271a007276bf87 Mon Sep 17 00:00:00 2001 From: Noeri Huisman <8823461+mrxz@users.noreply.github.com> Date: Fri, 3 Jul 2026 18:09:39 +0200 Subject: [PATCH 2/2] Decode images on the JS side and add support for unbundled SOG files --- rust/spark-lib/Cargo.toml | 1 + rust/spark-lib/src/decoder.rs | 8 +- rust/spark-lib/src/sogs.rs | 59 ++++++++++- rust/spark-rs/Cargo.toml | 2 +- src/sogs.ts | 181 ++++++++++++++++++++++++++++++++++ src/worker.ts | 35 ++++++- 6 files changed, 278 insertions(+), 8 deletions(-) create mode 100644 src/sogs.ts diff --git a/rust/spark-lib/Cargo.toml b/rust/spark-lib/Cargo.toml index 75dee09b..4de46847 100644 --- a/rust/spark-lib/Cargo.toml +++ b/rust/spark-lib/Cargo.toml @@ -14,6 +14,7 @@ ksplat = [] ply = [] rad = ["dep:miniz_oxide"] sogs = ["dep:zip", "dep:image"] +sogs_web_decode = ["sogs"] spz = ["dep:miniz_oxide"] csplat = [] gsplat = [] diff --git a/rust/spark-lib/src/decoder.rs b/rust/spark-lib/src/decoder.rs index 2aa59a7a..2693499c 100644 --- a/rust/spark-lib/src/decoder.rs +++ b/rust/spark-lib/src/decoder.rs @@ -13,7 +13,7 @@ use crate::ply::{PLY_MAGIC, PlyDecoder}; #[cfg(feature = "rad")] use crate::rad::{RAD_CHUNK_MAGIC, RAD_MAGIC, RadDecoder}; #[cfg(feature = "sogs")] -use crate::sogs::{PK_MAGIC, SogsDecoder}; +use crate::sogs::{PK_MAGIC, CUSTOM_SOGS_MAGIC, SogsDecoder}; #[cfg(feature = "spz")] use crate::spz::{SPZ_MAGIC, SpzDecoder}; @@ -377,7 +377,7 @@ impl SplatFileType { #[cfg(feature = "ksplat")] "ksplat" => Ok(Self::KSPLAT), #[cfg(feature = "sogs")] - "pcsogszip" => Ok(Self::SOGS), + "pcsogs" | "pcsogszip" => Ok(Self::SOGS), #[cfg(feature = "rad")] "rad" => Ok(Self::RAD), _ => Err(anyhow::anyhow!("Invalid file type: {}", enum_str)), @@ -538,6 +538,10 @@ impl ChunkReceiver for MultiDecoder { } } } + #[cfg(feature = "sogs")] + (CUSTOM_SOGS_MAGIC, _) => { + return self.init_file_type(SplatFileType::SOGS); + } #[cfg(feature = "rad")] (RAD_MAGIC, _) | (RAD_CHUNK_MAGIC, _) => { return self.init_file_type(SplatFileType::RAD); diff --git a/rust/spark-lib/src/sogs.rs b/rust/spark-lib/src/sogs.rs index d8f843af..7707b343 100644 --- a/rust/spark-lib/src/sogs.rs +++ b/rust/spark-lib/src/sogs.rs @@ -1,14 +1,19 @@ -use std::{collections::HashMap, io::Cursor}; +use std::collections::HashMap; +#[cfg(not(feature = "sogs_web_decode"))] +use std::io::Cursor; use anyhow::{anyhow, Context}; +#[cfg(not(feature = "sogs_web_decode"))] use image::{DynamicImage, GenericImageView, ImageReader}; use serde_json; use serde::Deserialize; +#[cfg(not(feature = "sogs_web_decode"))] use zip::ZipArchive; use crate::decoder::{ChunkReceiver, SplatInit, SplatProps, SplatReceiver}; pub const PK_MAGIC: u32 = 0x04034b50; +pub const CUSTOM_SOGS_MAGIC: u32 = 0x53474F53; const SH_C0: f32 = 0.28209479177387814; const MAX_SPLAT_CHUNK: usize = 65536; @@ -133,7 +138,7 @@ impl ChunkReceiver for SogsDecoder { return Err(anyhow!("SOGS file too small")); } let magic = u32::from_le_bytes([self.buffer[0], self.buffer[1], self.buffer[2], self.buffer[3]]); - if magic != PK_MAGIC { + if magic != PK_MAGIC && magic != CUSTOM_SOGS_MAGIC { return Err(anyhow!("Not a ZIP/SOGS file")); } decode_sogs(&self.buffer, &mut self.splats, None)?; @@ -141,6 +146,53 @@ impl ChunkReceiver for SogsDecoder { } } +#[cfg(feature = "sogs_web_decode")] +fn decode_sogs(bytes: &[u8], splats: &mut T, _pathname: Option<&str>) -> anyhow::Result<()> { + let mut file_cache: HashMap> = HashMap::new(); + let mut offset: usize = 4; // Skip magic number + + while offset < bytes.len() { + let name_len = u16::from_le_bytes(bytes[offset..offset + 2].try_into().unwrap()) as usize; + offset += 2; + let name_bytes = &bytes[offset..offset + name_len]; + offset += name_len; + let name = std::str::from_utf8(name_bytes).unwrap(); + + let data_size = u32::from_le_bytes(bytes[offset..offset + 4].try_into().unwrap()) as usize; + offset += 4; + let data = &bytes[offset..offset + data_size]; + offset += data_size; + + file_cache.insert(name.to_string(), data.to_vec()); + } + + let mut get_image_data = |name: &str| -> anyhow::Result { + let mut file = file_cache.get(name).cloned().ok_or_else(|| anyhow!("Missing file {name} in cache"))?; + + let width = u32::from_le_bytes( + file[0..4].try_into().unwrap() + ) as usize; + + let height = u32::from_le_bytes( + file[4..8].try_into().unwrap() + ) as usize; + + let rgba = file.split_off(8); + + Ok(ImageData { width, height, rgba }) + }; + + let meta_bytes = file_cache.get("meta.json").cloned().ok_or_else(|| anyhow!("Missing meta.json in cache"))?; + let meta: PcSogsRoot = serde_json::from_slice(&meta_bytes) + .context("Failed to parse meta.json for SOGS")?; + + match meta { + PcSogsRoot::V2(v2) => decode_v2(v2, splats, &mut get_image_data), + PcSogsRoot::V1(v1) => decode_v1(v1, splats, &mut get_image_data), + } +} + +#[cfg(not(feature = "sogs_web_decode"))] fn decode_sogs(bytes: &[u8], splats: &mut T, _pathname: Option<&str>) -> anyhow::Result<()> { let cursor = Cursor::new(bytes); let mut zip = ZipArchive::new(cursor)?; @@ -560,6 +612,7 @@ fn emit_to_receiver( splats.finish() } +#[cfg(not(feature = "sogs_web_decode"))] fn preload_all( meta: &PcSogsRoot, prefix: &str, @@ -589,6 +642,7 @@ fn preload_all( Ok(()) } +#[cfg(not(feature = "sogs_web_decode"))] fn preload_file( zip: &mut ZipArchive>, prefix: &str, @@ -620,6 +674,7 @@ struct ImageData { height: usize, } +#[cfg(not(feature = "sogs_web_decode"))] fn decode_image(bytes: &[u8]) -> anyhow::Result { let img = ImageReader::new(Cursor::new(bytes)) .with_guessed_format()? diff --git a/rust/spark-rs/Cargo.toml b/rust/spark-rs/Cargo.toml index 32870ac8..779c8ac3 100644 --- a/rust/spark-rs/Cargo.toml +++ b/rust/spark-rs/Cargo.toml @@ -19,7 +19,7 @@ antisplat = ["spark-lib/antisplat"] ksplat = ["spark-lib/ksplat"] ply = ["spark-lib/ply"] rad = ["spark-lib/rad"] -sogs = ["spark-lib/sogs"] +sogs = ["spark-lib/sogs", "spark-lib/sogs_web_decode"] spz = ["spark-lib/spz"] csplat = ["spark-lib/csplat"] gsplat = ["spark-lib/gsplat"] diff --git a/src/sogs.ts b/src/sogs.ts new file mode 100644 index 00000000..726471fd --- /dev/null +++ b/src/sogs.ts @@ -0,0 +1,181 @@ +import { unzipSync } from "fflate"; + +// Custom magic number for (unzipped) and decoded SOG files +const HEADER = new Uint8Array([0x53, 0x4f, 0x47, 0x53]); + +const temp = new Uint8Array(4); +const dataView = new DataView(temp.buffer); +const textDecoder = new TextDecoder(); +const textEncoder = new TextEncoder(); + +type ImageData = { width: number; height: number; rgba: Uint8Array }; + +export function unzipAndDecodeImages(zipSize: number) { + const data = new Uint8Array(zipSize); + let processed = 0; + + return new TransformStream({ + start(controller) { + controller.enqueue(HEADER); + }, + async transform(chunk, controller) { + const chunkData = await chunk; + data.set(chunkData, processed); + processed += chunkData.length; + }, + async flush(controller) { + const unzipped = unzipSync(data); + + const promises: Array> = []; + for (const fileName in unzipped) { + if (fileName.endsWith(".webp")) { + promises.push( + decodeImage(unzipped[fileName].buffer as ArrayBuffer).then( + (imageData) => enqueueImage(controller, fileName, imageData), + ), + ); + } else { + enqueueFile(controller, fileName, unzipped[fileName]); + } + } + + await Promise.allSettled(promises); + }, + }); +} + +export function fetchAndDecodeImages(url: string) { + // Strip the file name + const baseUrl = url.substring(0, url.lastIndexOf("/")); + + return new ReadableStream({ + async start(controller) { + // Fetch the meta.json file + const arrayBuffer = await (await fetch(url)).arrayBuffer(); + const json = JSON.parse(textDecoder.decode(arrayBuffer)); + const refFiles = [ + ...json.means.files, + ...json.scales.files, + ...json.quats.files, + ...json.sh0.files, + ...(json.shN?.files ?? []), + ]; + + // Start outputting + controller.enqueue(HEADER); + enqueueFile(controller, "meta.json", new Uint8Array(arrayBuffer)); + + const promises = refFiles.map(async (imageFile) => { + const response = await fetch(`${baseUrl}/${imageFile}`); + const arrayBuffer = await response.arrayBuffer(); + const imageData = await decodeImage(arrayBuffer); + + enqueueImage(controller, imageFile, imageData); + }); + + await Promise.allSettled(promises); + controller.close(); + }, + }); +} + +function enqueueFileName( + controller: + | ReadableStreamDefaultController + | TransformStreamDefaultController, + fileName: string, +) { + const encodedFileName = textEncoder.encode(fileName); + dataView.setUint16(0, encodedFileName.byteLength, true); + controller.enqueue(temp.slice(0, 2)); + controller.enqueue(encodedFileName); +} + +function enqueueFile( + controller: + | ReadableStreamDefaultController + | TransformStreamDefaultController, + fileName: string, + data: Uint8Array, +) { + enqueueFileName(controller, fileName); + + dataView.setUint32(0, data.byteLength, true); + controller.enqueue(temp.slice()); + controller.enqueue(data); +} + +function enqueueImage( + controller: + | ReadableStreamDefaultController + | TransformStreamDefaultController, + fileName: string, + imageData: ImageData, +) { + enqueueFileName(controller, fileName); + + // byte size + dataView.setUint32(0, imageData.rgba.byteLength + 8, true); + controller.enqueue(temp.slice()); + + // width + dataView.setUint32(0, imageData.width, true); + controller.enqueue(temp.slice()); + // height + dataView.setUint32(0, imageData.height, true); + controller.enqueue(temp.slice()); + // rgba + controller.enqueue(imageData.rgba); +} + +// WebGL context for reading raw pixel data of WebP images +let offscreenGlContext: WebGL2RenderingContext | null = null; + +export async function decodeImage(fileBytes: ArrayBuffer) { + if (!offscreenGlContext) { + const canvas = new OffscreenCanvas(1, 1); + offscreenGlContext = canvas.getContext("webgl2"); + if (!offscreenGlContext) { + throw new Error("Failed to create WebGL2 context"); + } + } + + const imageBlob = new Blob([fileBytes]); + const bitmap = await createImageBitmap(imageBlob, { + premultiplyAlpha: "none", + }); + + const gl = offscreenGlContext; + const texture = gl.createTexture(); + gl.bindTexture(gl.TEXTURE_2D, texture); + gl.pixelStorei(gl.UNPACK_FLIP_Y_WEBGL, true); + gl.texImage2D(gl.TEXTURE_2D, 0, gl.RGBA, gl.RGBA, gl.UNSIGNED_BYTE, bitmap); + gl.texParameteri(gl.TEXTURE_2D, gl.TEXTURE_MAG_FILTER, gl.NEAREST); + gl.texParameteri(gl.TEXTURE_2D, gl.TEXTURE_MIN_FILTER, gl.NEAREST); + + const framebuffer = gl.createFramebuffer(); + gl.bindFramebuffer(gl.FRAMEBUFFER, framebuffer); + gl.framebufferTexture2D( + gl.FRAMEBUFFER, + gl.COLOR_ATTACHMENT0, + gl.TEXTURE_2D, + texture, + 0, + ); + + const data = new Uint8Array(bitmap.width * bitmap.height * 4); + gl.readPixels( + 0, + 0, + bitmap.width, + bitmap.height, + gl.RGBA, + gl.UNSIGNED_BYTE, + data, + ); + + gl.deleteTexture(texture); + gl.deleteFramebuffer(framebuffer); + + return { rgba: data, width: bitmap.width, height: bitmap.height }; +} diff --git a/src/worker.ts b/src/worker.ts index 3687a25c..02cccda8 100644 --- a/src/worker.ts +++ b/src/worker.ts @@ -19,9 +19,15 @@ import init_wasm, { bhatt_lod_extsplats, get_lod_tree_level, } from "spark-rs"; -import type { ExtResult, PackedResult, SplatEncoding } from "./defines"; - -const rpcHandlers = { +import { + type ExtResult, + type PackedResult, + type SplatEncoding, + SplatFileType, +} from "./defines"; +import { fetchAndDecodeImages, unzipAndDecodeImages } from "./sogs"; + +export const rpcHandlers = { sortSplats16, sortSplats32, loadPackedSplats, @@ -99,8 +105,10 @@ function sortSplats32({ async function decodeBytesUrl({ decoder, + fileType, fileBytes, url, + pathName, requestHeader, withCredentials, chunked, @@ -108,8 +116,10 @@ async function decodeBytesUrl({ sendStatus, }: { decoder: ChunkDecoder; + fileType?: string; fileBytes?: Uint8Array; url?: string; + pathName?: string; requestHeader?: Record; withCredentials?: boolean; chunked?: boolean; @@ -127,6 +137,9 @@ async function decodeBytesUrl({ }, }); streamLength = fileBytes.length; + } else if (url && fileType === SplatFileType.PCSOGS) { + // Unbundled SOG files require fetching and decoding + readStream = fetchAndDecodeImages(url); } else if (url) { const request = new Request(url, { headers: requestHeader ? new Headers(requestHeader) : undefined, @@ -174,6 +187,14 @@ async function decodeBytesUrl({ throw new Error("No url or fileBytes provided"); } + // Handle SOG files + if ( + fileType === SplatFileType.PCSOGSZIP || + (fileType === undefined && /\.(sogs?|zip)$/.test(pathName ?? url ?? "")) + ) { + readStream = readStream.pipeThrough(unzipAndDecodeImages(streamLength)); + } + const reader = readStream.getReader(); let loaded = 0; while (true) { @@ -276,8 +297,10 @@ async function loadPackedSplats( ); const decoded = await decodeBytesUrl({ decoder, + fileType, fileBytes, url, + pathName, requestHeader, withCredentials, chunked, @@ -294,8 +317,10 @@ async function loadPackedSplats( const decoder = decode_to_csplatarray(fileType, pathName ?? url, encoding); const decoded = await decodeBytesUrl({ decoder, + fileType, fileBytes, url, + pathName, requestHeader, withCredentials, chunked, @@ -434,8 +459,10 @@ async function loadExtSplats( ); const decoded = await decodeBytesUrl({ decoder, + fileType, fileBytes, url, + pathName, requestHeader, withCredentials, chunked, @@ -452,8 +479,10 @@ async function loadExtSplats( const decoder = decode_to_gsplatarray(fileType, pathName ?? url); const decoded = await decodeBytesUrl({ decoder, + fileType, fileBytes, url, + pathName, requestHeader, withCredentials, chunked,