Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions rust/spark-lib/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@ license.workspace = true
authors.workspace = true
repository.workspace = true

[features]
sogs_web_decode = []

[dependencies]
ahash.workspace = true
anyhow.workspace = true
Expand Down
7 changes: 5 additions & 2 deletions rust/spark-lib/src/decoder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use crate::{
ksplat::KsplatDecoder,
ply::{PLY_MAGIC, PlyDecoder},
rad::{RAD_CHUNK_MAGIC, RAD_MAGIC, RadDecoder},
sogs::SogsDecoder,
sogs::{SogsDecoder, PK_MAGIC, CUSTOM_SOGS_MAGIC},
spz::{SPZ_MAGIC, SpzDecoder}
};

Expand Down Expand Up @@ -355,6 +355,7 @@ impl SplatFileType {
"spz" => Ok(Self::SPZ),
"splat" => Ok(Self::ANTISPLAT),
"ksplat" => Ok(Self::KSPLAT),
"pcsogs" => Ok(Self::SOGS),
"pcsogszip" => Ok(Self::SOGS),
"rad" => Ok(Self::RAD),
_ => Err(anyhow::anyhow!("Invalid file type: {}", enum_str)),
Expand Down Expand Up @@ -488,13 +489,15 @@ impl<T: SplatReceiver> ChunkReceiver for MultiDecoder<T> {
}
}
}
} else if magic == 0x04034b50 {
} else if magic == PK_MAGIC {
detection_complete = true;
if let Some(pathname) = &self.pathname {
if let Some(SplatFileType::SOGS) = SplatFileType::from_pathname(pathname) {
return self.init_file_type(SplatFileType::SOGS);
}
}
} else if magic == CUSTOM_SOGS_MAGIC {
return self.init_file_type(SplatFileType::SOGS);
} else if magic == RAD_MAGIC || magic == RAD_CHUNK_MAGIC {
return self.init_file_type(SplatFileType::RAD);
} else {
Expand Down
125 changes: 83 additions & 42 deletions rust/spark-lib/src/sogs.rs
Original file line number Diff line number Diff line change
@@ -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};

const PK_MAGIC: u32 = 0x04034b50;
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;

Expand Down Expand Up @@ -133,14 +138,61 @@ impl<T: SplatReceiver> ChunkReceiver for SogsDecoder<T> {
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)?;
Ok(())
}
}

#[cfg(feature = "sogs_web_decode")]
fn decode_sogs<T: SplatReceiver>(bytes: &[u8], splats: &mut T, _pathname: Option<&str>) -> anyhow::Result<()> {
let mut file_cache: HashMap<String, Vec<u8>> = 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<ImageData> {
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<T: SplatReceiver>(bytes: &[u8], splats: &mut T, _pathname: Option<&str>) -> anyhow::Result<()> {
let cursor = Cursor::new(bytes);
let mut zip = ZipArchive::new(cursor)?;
Expand Down Expand Up @@ -169,20 +221,21 @@ fn decode_sogs<T: SplatReceiver>(bytes: &[u8], splats: &mut T, _pathname: Option
let mut file_cache: HashMap<String, Vec<u8>> = HashMap::new();
preload_all(&meta, &prefix, &mut zip, &mut file_cache)?;

let mut get_file = |name: &str| -> anyhow::Result<Vec<u8>> {
file_cache.get(name).cloned().ok_or_else(|| anyhow!("Missing file {name} in cache"))
let mut get_image_data = |name: &str| -> anyhow::Result<ImageData> {
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<T: SplatReceiver>(
meta: PcSogsV2,
splats: &mut T,
get_file: &mut dyn FnMut(&str) -> anyhow::Result<Vec<u8>>,
get_image_data: &mut dyn FnMut(&str) -> anyhow::Result<ImageData>,
) -> anyhow::Result<()> {
let _ = meta.version;
let num_splats = meta.count;
Expand All @@ -191,16 +244,11 @@ fn decode_v2<T: SplatReceiver>(
}).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];
Expand All @@ -222,7 +270,7 @@ fn decode_v2<T: SplatReceiver>(
if let Some(shn) = meta.shn {
decode_shn_v2(
shn,
get_file,
get_image_data,
num_splats,
&mut sh1,
&mut sh2,
Expand All @@ -248,7 +296,7 @@ fn decode_v2<T: SplatReceiver>(
fn decode_v1<T: SplatReceiver>(
meta: PcSogsV1,
splats: &mut T,
get_file: &mut dyn FnMut(&str) -> anyhow::Result<Vec<u8>>,
get_image_data: &mut dyn FnMut(&str) -> anyhow::Result<ImageData>,
) -> anyhow::Result<()> {
let num_splats = meta.means.shape[0];
if meta.quats.encoding.as_deref() != Some("quaternion_packed") {
Expand All @@ -270,16 +318,11 @@ fn decode_v1<T: SplatReceiver>(

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];
Expand All @@ -298,7 +341,7 @@ fn decode_v1<T: SplatReceiver>(
if let Some(shn) = meta.shn {
decode_shn_v1(
shn,
get_file,
get_image_data,
num_splats,
max_sh_degree,
&mut sh1,
Expand Down Expand Up @@ -447,14 +490,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<Vec<u8>>,
get_image_data: &mut dyn FnMut(&str) -> anyhow::Result<ImageData>,
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;
Expand Down Expand Up @@ -490,15 +533,15 @@ fn decode_shn_v2(

fn decode_shn_v1(
shn: ShNV1,
get_file: &mut dyn FnMut(&str) -> anyhow::Result<Vec<u8>>,
get_image_data: &mut dyn FnMut(&str) -> anyhow::Result<ImageData>,
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<f32> = (0..256)
.map(|i| shn.mins + (shn.maxs - shn.mins) * (i as f32 / 255.0))
.collect();
Expand Down Expand Up @@ -569,6 +612,7 @@ fn emit_to_receiver<T: SplatReceiver>(
splats.finish()
}

#[cfg(not(feature = "sogs_web_decode"))]
fn preload_all(
meta: &PcSogsRoot,
prefix: &str,
Expand Down Expand Up @@ -598,6 +642,7 @@ fn preload_all(
Ok(())
}

#[cfg(not(feature = "sogs_web_decode"))]
fn preload_file(
zip: &mut ZipArchive<Cursor<&[u8]>>,
prefix: &str,
Expand Down Expand Up @@ -629,11 +674,7 @@ struct ImageData {
height: usize,
}

fn decode_rgba(bytes: &[u8]) -> anyhow::Result<ImageData> {
let img = decode_image(bytes)?;
Ok(img)
}

#[cfg(not(feature = "sogs_web_decode"))]
fn decode_image(bytes: &[u8]) -> anyhow::Result<ImageData> {
let img = ImageReader::new(Cursor::new(bytes))
.with_guessed_format()?
Expand Down
2 changes: 1 addition & 1 deletion rust/spark-rs/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ ordered-float.workspace = true
smallvec.workspace = true
wasm-bindgen.workspace = true
web-sys = { workspace = true, features = ["Window", "Performance"] }
spark-lib = { path = "../spark-lib" }
spark-lib = { path = "../spark-lib", features = ["sogs_web_decode"] }
serde-wasm-bindgen.workspace = true
serde_json.workspace = true
itertools.workspace = true
Expand Down
Loading
Loading