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()