Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 7 additions & 8 deletions src/pyrecest/distributions/circle/wrapped_normal_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`."""
Expand All @@ -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)),
)
Expand All @@ -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."
)
Expand All @@ -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):
Expand Down
43 changes: 43 additions & 0 deletions tests/distributions/test_wrapped_normal_subclass_preservation.py
Original file line number Diff line number Diff line change
@@ -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)
Loading