()
| 14 | |
| 15 | @pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device") |
| 16 | def test_png_compression(): |
| 17 | from gsplat.compression import PngCompression |
| 18 | |
| 19 | torch.manual_seed(42) |
| 20 | |
| 21 | # Prepare Gaussians |
| 22 | N = 100000 |
| 23 | splats = torch.nn.ParameterDict( |
| 24 | { |
| 25 | "means": torch.randn(N, 3), |
| 26 | "scales": torch.randn(N, 3), |
| 27 | "quats": torch.randn(N, 4), |
| 28 | "opacities": torch.randn(N), |
| 29 | "sh0": torch.randn(N, 1, 3), |
| 30 | "shN": torch.randn(N, 24, 3), |
| 31 | "features": torch.randn(N, 128), |
| 32 | } |
| 33 | ).to(device) |
| 34 | compress_dir = "/tmp/gsplat/compression" |
| 35 | |
| 36 | compression_method = PngCompression() |
| 37 | # run compression and save the compressed files to compress_dir |
| 38 | compression_method.compress(compress_dir, splats) |
| 39 | # decompress the compressed files |
| 40 | splats_c = compression_method.decompress(compress_dir) |
| 41 | |
| 42 | |
| 43 | if __name__ == "__main__": |
no test coverage detected