diff --git a/src/pyrecest/distributions/hypertorus/hypertoroidal_dirac_distribution.py b/src/pyrecest/distributions/hypertorus/hypertoroidal_dirac_distribution.py index 6c15433c1..c015e1d15 100644 --- a/src/pyrecest/distributions/hypertorus/hypertoroidal_dirac_distribution.py +++ b/src/pyrecest/distributions/hypertorus/hypertoroidal_dirac_distribution.py @@ -175,6 +175,20 @@ def trigonometric_moment(self, n: Union[int, int32, int64]): def apply_function(self, f: Callable, function_is_vectorized: bool = True): dist = super().apply_function(f, function_is_vectorized) dist.d = mod(dist.d, 2.0 * pi) + + if dist.d.ndim == 1: + transformed_dim = 1 + elif dist.d.ndim == 2 and dist.d.shape[1] > 0: + transformed_dim = int(dist.d.shape[1]) + else: + raise ValueError( + "Function output must have shape (n,) or (n, dim) with dim > 0." + ) + + if transformed_dim != self.dim: + return HypertoroidalDiracDistribution( + dist.d, dist.w, dim=transformed_dim + ) return dist def to_toroidal_wd(self): diff --git a/tests/distributions/test_hypertoroidal_dirac_apply_function_dimension.py b/tests/distributions/test_hypertoroidal_dirac_apply_function_dimension.py new file mode 100644 index 000000000..6bff14a7c --- /dev/null +++ b/tests/distributions/test_hypertoroidal_dirac_apply_function_dimension.py @@ -0,0 +1,23 @@ +from pyrecest.backend import array +from pyrecest.distributions import HypertoroidalDiracDistribution + + +def test_apply_function_updates_dimension_after_coordinate_reduction(): + distribution = HypertoroidalDiracDistribution( + array( + [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6], + [0.7, 0.8, 0.9], + ] + ), + array([0.2, 0.3, 0.5]), + ) + + transformed = distribution.apply_function(lambda points: points[:, :2]) + + assert distribution.dim == 3 + assert transformed.dim == 2 + assert transformed.d.shape == (3, 2) + assert transformed.w.shape == (3,) + assert transformed.trigonometric_moment(1).shape == (2,)