MCPcopy Create free account
hub / github.com/JasonLSC/GSCodec_Studio / test_proj

Function test_proj

tests/test_basic.py:131–173  ·  view source on GitHub ↗
(test_data, camera_model: Literal["pinhole", "ortho", "fisheye"])

Source from the content-addressed store, hash-verified

129@pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device")
130@pytest.mark.parametrize("camera_model", ["pinhole", "ortho", "fisheye"])
131def test_proj(test_data, camera_model: Literal["pinhole", "ortho", "fisheye"]):
132 from gsplat.cuda._torch_impl import _persp_proj, _ortho_proj, _fisheye_proj
133 from gsplat.cuda._wrapper import proj, quat_scale_to_covar_preci, world_to_cam
134
135 torch.manual_seed(42)
136
137 Ks = test_data["Ks"]
138 viewmats = test_data["viewmats"]
139 height = test_data["height"]
140 width = test_data["width"]
141
142 covars, _ = quat_scale_to_covar_preci(test_data["quats"], test_data["scales"])
143 means, covars = world_to_cam(test_data["means"], covars, viewmats)
144 means.requires_grad = True
145 covars.requires_grad = True
146
147 # forward
148 means2d, covars2d = proj(means, covars, Ks, width, height, camera_model)
149 if camera_model == "ortho":
150 _means2d, _covars2d = _ortho_proj(means, covars, Ks, width, height)
151 elif camera_model == "fisheye":
152 _means2d, _covars2d = _fisheye_proj(means, covars, Ks, width, height)
153 elif camera_model == "pinhole":
154 _means2d, _covars2d = _persp_proj(means, covars, Ks, width, height)
155 else:
156 assert_never(camera_model)
157
158 torch.testing.assert_close(means2d, _means2d, rtol=1e-4, atol=1e-4)
159 torch.testing.assert_close(covars2d, _covars2d, rtol=1e-1, atol=3e-2)
160
161 # backward
162 v_means2d = torch.randn_like(means2d)
163 v_covars2d = torch.randn_like(covars2d)
164 v_means, v_covars = torch.autograd.grad(
165 (means2d * v_means2d).sum() + (covars2d * v_covars2d).sum(),
166 (means, covars),
167 )
168 _v_means, _v_covars = torch.autograd.grad(
169 (_means2d * v_means2d).sum() + (_covars2d * v_covars2d).sum(),
170 (means, covars),
171 )
172 torch.testing.assert_close(v_means, _v_means, rtol=1e-2, atol=1e-2)
173 torch.testing.assert_close(v_covars, _v_covars, rtol=1e-1, atol=1e-1)
174
175
176@pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device")

Callers

nothing calls this directly

Calls 6

world_to_camFunction · 0.90
projFunction · 0.90
_ortho_projFunction · 0.90
_fisheye_projFunction · 0.90
_persp_projFunction · 0.90

Tested by

no test coverage detected