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
14 changes: 7 additions & 7 deletions src/pyrecest/distributions/circle/von_mises_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -260,15 +260,15 @@ def multiply(self, vm2: "VonMisesDistribution") -> "VonMisesDistribution":
scale = maximum(abs(C), abs(S))
safe_scale = where(scale > 0.0, scale, 1.0)
kappa_ = scale * sqrt((C / safe_scale) ** 2 + (S / safe_scale) ** 2)
return VonMisesDistribution(mu_, kappa_)
return type(self)(mu_, kappa_)

def convolve(self, vm2: "VonMisesDistribution") -> "VonMisesDistribution":
mu_ = mod(self.mu + vm2.mu, 2.0 * pi)
t = VonMisesDistribution.besselratio(
0, self.kappa
) * VonMisesDistribution.besselratio(0, vm2.kappa)
kappa_ = VonMisesDistribution.besselratio_inverse(0, t)
return VonMisesDistribution(mu_, kappa_)
return type(self)(mu_, kappa_)

def entropy(self):
result = (
Expand Down Expand Up @@ -303,8 +303,8 @@ def trigonometric_moment_analytic(self, n: int):

return m

@staticmethod
def from_moment(m):
@classmethod
def from_moment(cls, m):
"""
Obtain a VM distribution from a given first trigonometric moment.

Expand All @@ -314,12 +314,12 @@ def from_moment(m):
Returns:
vm (VMDistribution): Distribution obtained by moment matching.
"""
kappa_ = VonMisesDistribution.besselratio_inverse(0, abs(m))
if VonMisesDistribution._as_float_scalar(kappa_, "kappa") == 0.0:
kappa_ = cls.besselratio_inverse(0, abs(m))
if cls._as_float_scalar(kappa_, "kappa") == 0.0:
mu_ = array(0.0)
else:
mu_ = mod(arctan2(imag(m), real(m)), 2.0 * pi)
vm = VonMisesDistribution(mu_, kappa_)
vm = cls(mu_, kappa_)
return vm

def __str__(self) -> str:
Expand Down
35 changes: 35 additions & 0 deletions tests/distributions/test_von_mises_subclass_preservation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
# pylint: disable=no-name-in-module,no-member
from pyrecest.backend import array
from pyrecest.distributions import VonMisesDistribution


class _VonMisesSubclass(VonMisesDistribution):
pass


def test_from_moment_preserves_requested_subclass_for_regular_and_uniform_results():
source = VonMisesDistribution(array(0.3), array(2.0))

regular = _VonMisesSubclass.from_moment(source.trigonometric_moment(1))
uniform = _VonMisesSubclass.from_moment(array(0.0 + 0.0j))

assert isinstance(regular, _VonMisesSubclass)
assert isinstance(uniform, _VonMisesSubclass)


def test_multiply_preserves_left_subclass():
left = _VonMisesSubclass(array(0.2), array(1.5))
right = VonMisesDistribution(array(1.1), array(0.7))

multiplied = left.multiply(right)

assert isinstance(multiplied, _VonMisesSubclass)


def test_convolve_preserves_left_subclass_for_regular_and_uniform_results():
left = _VonMisesSubclass(array(0.2), array(1.5))
regular = VonMisesDistribution(array(1.1), array(0.7))
uniform = VonMisesDistribution(array(1.1), array(0.0))

assert isinstance(left.convolve(regular), _VonMisesSubclass)
assert isinstance(left.convolve(uniform), _VonMisesSubclass)
Loading