From e062d642c0a0caccb4574234b687b3bbd18f0735 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Sun, 23 Aug 2026 15:30:53 +0800 Subject: [PATCH 1/2] Preserve wrapped-normal subclasses in operations --- .../circle/wrapped_normal_distribution.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/src/pyrecest/distributions/circle/wrapped_normal_distribution.py b/src/pyrecest/distributions/circle/wrapped_normal_distribution.py index 4b5626dac3..29173d030a 100644 --- a/src/pyrecest/distributions/circle/wrapped_normal_distribution.py +++ b/src/pyrecest/distributions/circle/wrapped_normal_distribution.py @@ -241,8 +241,7 @@ def multiply_vm_approximation( vm1 = self.to_vm() vm2 = other.to_vm() vm = vm1.multiply(vm2) - wn = vm.to_wn() - return wn + return type(self).from_moment(vm.trigonometric_moment(1)) def multiply_vm(self, other: "WrappedNormalDistribution"): """Backward-compatible alias for :meth:`multiply_vm_approximation`.""" @@ -260,7 +259,7 @@ def convolve( """ if not isinstance(other, WrappedNormalDistribution): raise TypeError("other must be a WrappedNormalDistribution") - return WrappedNormalDistribution( + return type(self)( mod(self.scalar_mu + other.scalar_mu, 2.0 * pi), sqrt(squeeze(self.C) + squeeze(other.C)), ) @@ -287,22 +286,22 @@ def to_dirac5(self): def shift(self, shift_by): shift_by = as_shift_vector(shift_by, 1) - return WrappedNormalDistribution(self.scalar_mu + shift_by[0], self.sigma) + return type(self)(self.scalar_mu + shift_by[0], self.sigma) def to_vm(self) -> VonMisesDistribution: # Convert to Von Mises distribution kappa = self.sigma_to_kappa(self.sigma) return VonMisesDistribution(self.scalar_mu, kappa) - @staticmethod - def from_moment(m) -> "WrappedNormalDistribution": + @classmethod + def from_moment(cls, m) -> "WrappedNormalDistribution": moment = squeeze(array(m)) if ndim(moment) != 0: raise ValueError("First trigonometric moment must be a scalar.") moment_abs = float(abs(moment)) if not isfinite(moment_abs): raise ValueError("First trigonometric moment must be finite.") - if moment_abs > 1.0 + WrappedNormalDistribution._MOMENT_NORM_TOL: + if moment_abs > 1.0 + cls._MOMENT_NORM_TOL: raise ValueError( "First trigonometric moment must have magnitude at most 1." ) @@ -321,7 +320,7 @@ def from_moment(m) -> "WrappedNormalDistribution": mu = mod(angle(moment), 2.0 * pi) sigma = sqrt(-2 * log(moment_abs)) - return WrappedNormalDistribution(mu, sigma) + return cls(mu, sigma) @staticmethod def sigma_to_kappa(sigma): From 9bb3d24404e172ba8b3f9f76576e671debfe0d49 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Sun, 23 Aug 2026 15:30:59 +0800 Subject: [PATCH 2/2] Add wrapped-normal subclass regressions --- ...st_wrapped_normal_subclass_preservation.py | 43 +++++++++++++++++++ 1 file changed, 43 insertions(+) create mode 100644 tests/distributions/test_wrapped_normal_subclass_preservation.py diff --git a/tests/distributions/test_wrapped_normal_subclass_preservation.py b/tests/distributions/test_wrapped_normal_subclass_preservation.py new file mode 100644 index 0000000000..4f87beefb6 --- /dev/null +++ b/tests/distributions/test_wrapped_normal_subclass_preservation.py @@ -0,0 +1,43 @@ +# pylint: disable=no-name-in-module,no-member +from pyrecest.backend import array +from pyrecest.distributions import WrappedNormalDistribution + + +class _WrappedNormalSubclass(WrappedNormalDistribution): + pass + + +def test_from_moment_preserves_requested_subclass(): + source = WrappedNormalDistribution(array(0.3), array(0.8)) + + reconstructed = _WrappedNormalSubclass.from_moment( + source.trigonometric_moment(1) + ) + + assert isinstance(reconstructed, _WrappedNormalSubclass) + + +def test_shift_preserves_runtime_subclass(): + dist = _WrappedNormalSubclass(array(0.3), array(0.8)) + + shifted = dist.shift(0.2) + + assert isinstance(shifted, _WrappedNormalSubclass) + + +def test_convolve_preserves_left_subclass(): + left = _WrappedNormalSubclass(array(0.3), array(0.8)) + right = WrappedNormalDistribution(array(0.4), array(0.6)) + + convolved = left.convolve(right) + + assert isinstance(convolved, _WrappedNormalSubclass) + + +def test_multiply_paths_preserve_left_subclass(): + left = _WrappedNormalSubclass(array(0.3), array(0.8)) + right = WrappedNormalDistribution(array(0.4), array(0.6)) + + assert isinstance(left.multiply(right), _WrappedNormalSubclass) + assert isinstance(left.multiply_vm_approximation(right), _WrappedNormalSubclass) + assert isinstance(left.multiply_vm(right), _WrappedNormalSubclass)