File size: 5,253 Bytes
0b9c87f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | #!/usr/bin/env python3
# filepath: decimate_splat.py
"""
Decimate a Gaussian Splat PLY file and output as PLY or SPLAT format.
"""
import argparse
import numpy as np
from io import BytesIO
from pathlib import Path
from plyfile import PlyData, PlyElement
def load_gaussian_ply(ply_file_path: str) -> PlyData:
"""Load a Gaussian splat PLY file."""
return PlyData.read(ply_file_path)
def compute_importance_scores(vert) -> np.ndarray:
"""
Compute importance scores for each Gaussian.
Higher scores = more important (larger and more opaque).
"""
scales = np.exp(vert["scale_0"] + vert["scale_1"] + vert["scale_2"])
opacities = 1 / (1 + np.exp(-vert["opacity"]))
return scales * opacities
def decimate_ply(plydata: PlyData, keep_ratio: float) -> PlyData:
"""
Decimate the PLY data by keeping only a fraction of the Gaussians.
Keeps the most important Gaussians based on scale and opacity.
"""
vert = plydata["vertex"]
total_points = len(vert.data)
keep_count = max(1, int(total_points * keep_ratio))
# Compute importance and get indices of top Gaussians
importance = compute_importance_scores(vert)
sorted_indices = np.argsort(-importance)[:keep_count]
# Sort indices to maintain some spatial coherence
sorted_indices = np.sort(sorted_indices)
# Create new vertex data with only kept points
new_vertex_data = vert.data[sorted_indices]
# Create new PlyElement and PlyData
new_vertex_element = PlyElement.describe(new_vertex_data, "vertex")
new_plydata = PlyData([new_vertex_element])
return new_plydata
def convert_ply_to_splat(plydata: PlyData) -> bytes:
"""
Convert PLY data to SPLAT format for the antimatter15 viewer.
Returns the splat data as bytes.
"""
vert = plydata["vertex"]
sorted_indices = np.argsort(
-np.exp(vert["scale_0"] + vert["scale_1"] + vert["scale_2"])
/ (1 + np.exp(-vert["opacity"]))
)
buffer = BytesIO()
for idx in sorted_indices:
v = plydata["vertex"][idx]
position = np.array([v["x"], v["y"], v["z"]], dtype=np.float32)
scales = np.exp(
np.array([v["scale_0"], v["scale_1"], v["scale_2"]], dtype=np.float32)
)
color = np.array([
0.5 + 0.28209479177387814 * v["f_dc_0"],
0.5 + 0.28209479177387814 * v["f_dc_1"],
0.5 + 0.28209479177387814 * v["f_dc_2"],
1 / (1 + np.exp(-v["opacity"])),
])
rot = np.array([v["rot_0"], v["rot_1"], v["rot_2"], v["rot_3"]], dtype=np.float32)
buffer.write(position.tobytes())
buffer.write(scales.tobytes())
buffer.write((color * 255).clip(0, 255).astype(np.uint8).tobytes())
buffer.write(
((rot / np.linalg.norm(rot)) * 128 + 128).clip(0, 255).astype(np.uint8).tobytes()
)
return buffer.getvalue()
def main():
parser = argparse.ArgumentParser(
description="Decimate a Gaussian Splat PLY file and output as PLY or SPLAT format."
)
parser.add_argument(
"input",
type=str,
help="Input PLY file path"
)
parser.add_argument(
"-o", "--output",
type=str,
help="Output file path (default: input_decimated.ply or .splat)"
)
parser.add_argument(
"-r", "--ratio",
type=float,
default=0.5,
help="Ratio of points to keep (0.0-1.0, default: 0.5)"
)
parser.add_argument(
"-f", "--format",
type=str,
choices=["ply", "splat"],
default="ply",
help="Output format: 'ply' or 'splat' (default: ply)"
)
parser.add_argument(
"-v", "--verbose",
action="store_true",
help="Print verbose output"
)
args = parser.parse_args()
# Validate ratio
if not 0.0 < args.ratio <= 1.0:
parser.error("Ratio must be between 0.0 (exclusive) and 1.0 (inclusive)")
# Determine output path
input_path = Path(args.input)
if args.output:
output_path = Path(args.output)
else:
suffix = ".splat" if args.format == "splat" else ".ply"
output_path = input_path.with_stem(f"{input_path.stem}_decimated").with_suffix(suffix)
if args.verbose:
print(f"Loading: {args.input}")
# Load PLY file
plydata = load_gaussian_ply(args.input)
original_count = len(plydata["vertex"].data)
if args.verbose:
print(f"Original point count: {original_count:,}")
# Decimate
decimated_plydata = decimate_ply(plydata, args.ratio)
new_count = len(decimated_plydata["vertex"].data)
if args.verbose:
print(f"Decimated point count: {new_count:,} ({args.ratio * 100:.1f}%)")
# Write output
if args.format == "splat":
splat_data = convert_ply_to_splat(decimated_plydata)
with open(output_path, "wb") as f:
f.write(splat_data)
else:
decimated_plydata.write(str(output_path))
if args.verbose:
print(f"Saved to: {output_path}")
else:
print(f"Decimated {original_count:,} → {new_count:,} points, saved to {output_path}")
if __name__ == "__main__":
main()
|