InfiniSplat / tests /test_decoder_gsplat.py
PLUS-WAVE's picture
Deploy InfiniSplat ZeroGPU demo
41ff959 verified
Raw
History Blame Contribute Delete
2.2 kB
import pytest
import torch
from src.model.decoder import decoder_gsplat
from src.model.decoder.decoder_gsplat import DecoderGsplat, DecoderGsplatCfg
from src.utils.gaussians import Gaussians3D
def test_decoder_rasterizes_all_views_without_tile_override(monkeypatch) -> None:
calls = []
def fake_rasterization(*args, **kwargs):
calls.append((args, kwargs))
viewmats = args[5]
width = args[7]
height = args[8]
rendering = torch.zeros(1, viewmats.shape[1], height, width, 3)
alpha = torch.zeros(1, viewmats.shape[1], height, width, 1)
return rendering, alpha, {}
monkeypatch.setattr(decoder_gsplat, "rasterization", fake_rasterization)
decoder = DecoderGsplat(
DecoderGsplatCfg(name="gsplat", background_color=[0.0, 0.0, 0.0])
)
gaussians = Gaussians3D(
mean_vectors=torch.zeros(1, 2, 3),
singular_values=torch.ones(1, 2, 3),
quaternions=torch.zeros(1, 2, 4),
colors=torch.zeros(1, 2, 3),
opacities=torch.ones(1, 2),
)
extrinsics = torch.eye(4).reshape(1, 1, 4, 4).repeat(1, 2, 1, 1)
intrinsics = torch.eye(3).reshape(1, 1, 3, 3).repeat(1, 2, 1, 1)
output = decoder(gaussians, extrinsics, intrinsics, image_shape=(4, 6))
assert len(calls) == 1
args, kwargs = calls[0]
assert args[5].shape[1] == 2
assert "tile_size" not in kwargs
assert kwargs["render_mode"] == "RGB"
assert output.shape == (1, 2, 3, 4, 6)
def test_decoder_reports_missing_optional_gsplat(monkeypatch) -> None:
monkeypatch.setattr(decoder_gsplat, "rasterization", None)
decoder = DecoderGsplat(
DecoderGsplatCfg(name="gsplat", background_color=[0.0, 0.0, 0.0])
)
gaussians = Gaussians3D(
mean_vectors=torch.zeros(1, 1, 3),
singular_values=torch.ones(1, 1, 3),
quaternions=torch.zeros(1, 1, 4),
colors=torch.zeros(1, 1, 3),
opacities=torch.ones(1, 1),
)
extrinsics = torch.eye(4).reshape(1, 1, 4, 4)
intrinsics = torch.eye(3).reshape(1, 1, 3, 3)
with pytest.raises(RuntimeError, match="optional `gsplat`"):
decoder(gaussians, extrinsics, intrinsics, image_shape=(4, 6))