diff --git a/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py b/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py index f1757d99ef..1ac0175f1c 100644 --- a/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py +++ b/src/pyrecest/distributions/cart_prod/hypercylindrical_state_space_subdivision_distribution.py @@ -16,6 +16,7 @@ concatenate, full, minimum, + mod, ndim, pi, random, @@ -124,10 +125,9 @@ def pdf(self, xs): # Find nearest grid indices (vectorised toroidal distance) grid = asarray(self.gd.get_grid()) # (n_grid, bound_dim) delta = grid[None, :, :] - x_bound[:, None, :] # (n_eval, n_grid, bound_dim) - abs_delta = abs(delta) - dists = backend_sum( - minimum(abs_delta**2, (2.0 * pi - abs_delta) ** 2), axis=-1 - ) # (n_eval, n_grid) + abs_delta = mod(abs(delta), 2.0 * pi) + wrapped_delta = minimum(abs_delta, 2.0 * pi - abs_delta) + dists = backend_sum(wrapped_delta**2, axis=-1) # (n_eval, n_grid) indices = argmin(dists, axis=1) # (n_eval,) # Evaluate periodic marginal pdf (nearest-neighbor) @@ -303,7 +303,6 @@ def fun_curr(y, _grid_i=grid_i): # Unnormalized conditional to compute the marginal weight cd_unnorm = CustomLinearDistribution(lambda x, fc=fun_curr: fc(x), dim_lin) - integral_val = float( cd_unnorm.integrate( left=array([float(int_range[0])]), diff --git a/tests/distributions/test_hsssd_periodicity.py b/tests/distributions/test_hsssd_periodicity.py new file mode 100644 index 0000000000..e54a67d1c1 --- /dev/null +++ b/tests/distributions/test_hsssd_periodicity.py @@ -0,0 +1,45 @@ +import unittest +from math import pi + +import numpy.testing as npt +from pyrecest.backend import __backend_name__ as backend_name +from pyrecest.backend import array +from pyrecest.distributions.cart_prod.hypercylindrical_state_space_subdivision_distribution import ( + HypercylindricalStateSpaceSubdivisionDistribution, +) +from pyrecest.distributions.hypertorus.hypertoroidal_grid_distribution import ( + HypertoroidalGridDistribution, +) +from pyrecest.distributions.nonperiodic.gaussian_distribution import ( + GaussianDistribution, +) + + +class HypercylindricalSubdivisionPeriodicityTest(unittest.TestCase): + @unittest.skipIf( + backend_name != "numpy", + reason="Not supported on this backend", + ) + def test_pdf_uses_periodic_conditional_selection_across_multiple_turns(self): + grid_distribution = HypertoroidalGridDistribution( + array([1.0, 1.0]), + grid_type="custom", + grid=array([[0.0], [pi]]), + ) + conditional_distributions = [ + GaussianDistribution(array([0.0]), array([[0.25]])), + GaussianDistribution(array([5.0]), array([[0.25]])), + ] + distribution = HypercylindricalStateSpaceSubdivisionDistribution( + grid_distribution, + conditional_distributions, + ) + + reference = distribution.pdf(array([[0.1, 0.0]])) + equivalent = distribution.pdf(array([[0.1 + 4.0 * pi, 0.0]])) + + npt.assert_allclose(equivalent, reference, rtol=0.0, atol=1.0e-12) + + +if __name__ == "__main__": + unittest.main()