diff --git a/src/pyrecest/distributions/circle/circular_fourier_distribution.py b/src/pyrecest/distributions/circle/circular_fourier_distribution.py index b267c375e..f05366f2d 100644 --- a/src/pyrecest/distributions/circle/circular_fourier_distribution.py +++ b/src/pyrecest/distributions/circle/circular_fourier_distribution.py @@ -357,8 +357,9 @@ def get_full_c(self): full_c = concatenate([neg_c, self.c]) # Concatenate arrays to get full spectrum return full_c - @staticmethod + @classmethod def from_distribution( + cls, distribution: AbstractCircularDistribution, n: Union[int, int32, int64], transformation: str = "sqrt", @@ -379,7 +380,7 @@ def from_distribution( for k in range(int(n) // 2 + 1) ] ) - fd = CircularFourierDistribution( + fd = cls( c=coeffs, n=int(n), transformation=transformation, @@ -398,14 +399,15 @@ def from_distribution( fvals = sqrt(fvals) else: raise NotImplementedError("Transformation not supported.") - fd = CircularFourierDistribution.from_function_values( + fd = cls.from_function_values( fvals, transformation, store_values_multiplied_by_n ) return fd - @staticmethod + @classmethod def from_function_values( + cls, fvals, transformation: str = "sqrt", store_values_multiplied_by_n: bool = True, @@ -416,7 +418,7 @@ def from_function_values( if not store_values_multiplied_by_n: c = c * (1.0 / n_values) - fd = CircularFourierDistribution( + fd = cls( c=c, transformation=transformation, n=n_values, diff --git a/tests/distributions/test_circular_fourier_factory_subclass_preservation.py b/tests/distributions/test_circular_fourier_factory_subclass_preservation.py new file mode 100644 index 000000000..ff75a7cc0 --- /dev/null +++ b/tests/distributions/test_circular_fourier_factory_subclass_preservation.py @@ -0,0 +1,59 @@ +import unittest + +# pylint: disable=no-name-in-module,no-member +from pyrecest.backend import array +from pyrecest.distributions.circle.circular_dirac_distribution import ( + CircularDiracDistribution, +) +from pyrecest.distributions.circle.circular_fourier_distribution import ( + CircularFourierDistribution, +) +from pyrecest.distributions.circle.von_mises_distribution import VonMisesDistribution +from pyrecest.distributions.conversion import convert_distribution + + +class _CircularFourierSubclass(CircularFourierDistribution): + pass + + +class CircularFourierFactorySubclassPreservationTest(unittest.TestCase): + def test_conversion_factory_preserves_requested_subclass_for_density_source(self): + source = VonMisesDistribution(0.3, 2.0) + + converted = convert_distribution( + source, + _CircularFourierSubclass, + n=9, + transformation="identity", + store_values_multiplied_by_n=False, + ) + + self.assertIsInstance(converted, _CircularFourierSubclass) + + def test_conversion_factory_preserves_requested_subclass_for_dirac_source(self): + source = CircularDiracDistribution(array([0.0, 1.0])) + + converted = convert_distribution( + source, + _CircularFourierSubclass, + n=9, + transformation="identity", + store_values_multiplied_by_n=False, + ) + + self.assertIsInstance(converted, _CircularFourierSubclass) + + def test_function_value_factory_preserves_requested_subclass(self): + function_values = array([1.0, 0.9, 0.8, 0.7, 0.6, 0.7, 0.8, 0.9, 1.0]) + + converted = _CircularFourierSubclass.from_function_values( + function_values, + transformation="identity", + store_values_multiplied_by_n=False, + ) + + self.assertIsInstance(converted, _CircularFourierSubclass) + + +if __name__ == "__main__": + unittest.main()