diff --git a/src/pyrecest/distributions/circle/von_mises_distribution.py b/src/pyrecest/distributions/circle/von_mises_distribution.py index 78be46f9c..be87a0aa1 100644 --- a/src/pyrecest/distributions/circle/von_mises_distribution.py +++ b/src/pyrecest/distributions/circle/von_mises_distribution.py @@ -260,7 +260,7 @@ 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) @@ -268,7 +268,7 @@ def convolve(self, vm2: "VonMisesDistribution") -> "VonMisesDistribution": 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 = ( @@ -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. @@ -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: diff --git a/tests/distributions/test_von_mises_subclass_preservation.py b/tests/distributions/test_von_mises_subclass_preservation.py new file mode 100644 index 000000000..6f9639f4e --- /dev/null +++ b/tests/distributions/test_von_mises_subclass_preservation.py @@ -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)