diff --git a/src/pyrecest/distributions/cart_prod/gauss_von_mises_distribution.py b/src/pyrecest/distributions/cart_prod/gauss_von_mises_distribution.py index 3f2359e89..82364e007 100644 --- a/src/pyrecest/distributions/cart_prod/gauss_von_mises_distribution.py +++ b/src/pyrecest/distributions/cart_prod/gauss_von_mises_distribution.py @@ -234,7 +234,7 @@ def pdf(self, xs): ) if single_point: - return float(p[0]) + return p[0] return array(p) def mode(self): diff --git a/tests/distributions/test_gauss_von_mises_pytorch_autograd.py b/tests/distributions/test_gauss_von_mises_pytorch_autograd.py new file mode 100644 index 000000000..4fc7b27bb --- /dev/null +++ b/tests/distributions/test_gauss_von_mises_pytorch_autograd.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import pytest + +import pyrecest.backend +from pyrecest.distributions.cart_prod.gauss_von_mises_distribution import ( + GaussVonMisesDistribution, +) + +torch = pytest.importorskip("torch") + +pytestmark = pytest.mark.skipif( + pyrecest.backend.__backend_name__ != "pytorch", + reason="PyTorch backend regression", +) + + +def test_single_point_pdf_preserves_pytorch_autograd() -> None: + distribution = GaussVonMisesDistribution( + mu=2.0, + P=1.3, + alpha=3.0, + beta=0.0, + Gamma=0.001, + kappa=0.7, + ) + point = torch.tensor( + [0.8, 1.4], + dtype=torch.float64, + requires_grad=True, + ) + + density = distribution.pdf(point) + + assert torch.is_tensor(density) + assert density.ndim == 0 + assert density.dtype == point.dtype + assert density.device == point.device + assert density.requires_grad + assert torch.isfinite(density) + assert density > 0.0 + + density.backward() + + assert point.grad is not None + assert torch.all(torch.isfinite(point.grad)) + assert torch.all(point.grad != 0.0)