(test_data, camera_model: Literal["pinhole", "ortho", "fisheye"])
| 129 | @pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device") |
| 130 | @pytest.mark.parametrize("camera_model", ["pinhole", "ortho", "fisheye"]) |
| 131 | def 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") |
nothing calls this directly
no test coverage detected