diff --git a/src/pyrecest/distributions/circle/circular_dirac_distribution.py b/src/pyrecest/distributions/circle/circular_dirac_distribution.py index aadae91f3..51ad183ba 100644 --- a/src/pyrecest/distributions/circle/circular_dirac_distribution.py +++ b/src/pyrecest/distributions/circle/circular_dirac_distribution.py @@ -31,9 +31,9 @@ def __init__(self, d, w=None): if self.d.shape != self.w.shape: raise ValueError("The shapes of d and w should match.") - @staticmethod + @classmethod def from_distribution( - distribution: AbstractCircularDistribution, n_particles: int | None = None + cls, distribution: AbstractCircularDistribution, n_particles: int | None = None ): """Create a circular Dirac approximation from a circular distribution.""" if not isinstance(distribution, AbstractCircularDistribution): @@ -54,14 +54,12 @@ def from_distribution( if bool(weight_scale > 0.0): weights = weights / weight_scale weights = weights / backend_sum(weights) - return CircularDiracDistribution(get_grid(), weights) + return cls(get_grid(), weights) if n_particles is None: raise ValueError("n_particles is required for sampling-based conversion.") - n_particles = HypertoroidalDiracDistribution._validate_particle_count( - n_particles - ) - return CircularDiracDistribution( + n_particles = cls._validate_particle_count(n_particles) + return cls( distribution.sample(n_particles), ones(n_particles) / n_particles ) diff --git a/src/pyrecest/distributions/circle/circular_grid_distribution.py b/src/pyrecest/distributions/circle/circular_grid_distribution.py index 959b8ca68..4bdf2a0fe 100644 --- a/src/pyrecest/distributions/circle/circular_grid_distribution.py +++ b/src/pyrecest/distributions/circle/circular_grid_distribution.py @@ -140,17 +140,17 @@ def pdf(self, xs, use_sinc=False, sinc_repetitions=5): return self._pdf_via_sinc(xs, sinc_repetitions) return self._pdf_via_fourier(xs) - @staticmethod - def from_distribution(distribution, no_of_gridpoints, enforce_pdf_nonnegative=True): - return CircularGridDistribution.from_function( + @classmethod + def from_distribution(cls, distribution, no_of_gridpoints, enforce_pdf_nonnegative=True): + return cls.from_function( distribution.pdf, no_of_gridpoints, enforce_pdf_nonnegative, ) - @staticmethod - def from_function(fun, no_of_gridpoints, enforce_pdf_nonnegative=True): + @classmethod + def from_function(cls, fun, no_of_gridpoints, enforce_pdf_nonnegative=True): no_of_gridpoints = _validate_no_of_gridpoints(no_of_gridpoints) grid_points = linspace(0.0, 2.0 * pi, no_of_gridpoints, endpoint=False) grid_values = array(fun(grid_points)) - return CircularGridDistribution(grid_values, enforce_pdf_nonnegative) + return cls(grid_values, enforce_pdf_nonnegative) diff --git a/src/pyrecest/distributions/nonperiodic/linear_dirac_distribution.py b/src/pyrecest/distributions/nonperiodic/linear_dirac_distribution.py index 993a6b0a8..37b537f38 100644 --- a/src/pyrecest/distributions/nonperiodic/linear_dirac_distribution.py +++ b/src/pyrecest/distributions/nonperiodic/linear_dirac_distribution.py @@ -122,18 +122,18 @@ def plot(self, *args, **kwargs): raise ValueError("Plotting not supported for this dimension") plt.show() - @staticmethod - def from_distribution(distribution, n_particles=None, n_samples=None, n=None): - particle_count = LinearDiracDistribution._resolve_particle_count( + @classmethod + def from_distribution(cls, distribution, n_particles=None, n_samples=None, n=None): + particle_count = cls._resolve_particle_count( n_particles=n_particles, n_samples=n_samples, n=n, ) samples = distribution.sample(particle_count) - return LinearDiracDistribution(samples, ones(particle_count) / particle_count) + return cls(samples, ones(particle_count) / particle_count) - @staticmethod - def _resolve_particle_count(n_particles=None, n_samples=None, n=None): + @classmethod + def _resolve_particle_count(cls, n_particles=None, n_samples=None, n=None): from ..conversion import ConversionError specified_counts = [ @@ -146,8 +146,7 @@ def _resolve_particle_count(n_particles=None, n_samples=None, n=None): ) particle_counts = [ - LinearDiracDistribution._validate_particle_count(value) - for value in specified_counts + cls._validate_particle_count(value) for value in specified_counts ] if len(set(particle_counts)) != 1: raise ConversionError( diff --git a/src/pyrecest/distributions/se2_dirac_distribution.py b/src/pyrecest/distributions/se2_dirac_distribution.py index 88c5a4a32..31f52f1d6 100644 --- a/src/pyrecest/distributions/se2_dirac_distribution.py +++ b/src/pyrecest/distributions/se2_dirac_distribution.py @@ -74,8 +74,8 @@ def mean(self): """ return self.hybrid_mean() - @staticmethod - def from_distribution(distribution, n_particles): + @classmethod + def from_distribution(cls, distribution, n_particles): """Create an SE2DiracDistribution by sampling from a given distribution. Parameters @@ -100,9 +100,9 @@ def from_distribution(distribution, n_particles): ) if distribution.bound_dim != 1 or distribution.lin_dim != 2: raise ValueError("distribution must have bound_dim=1 and lin_dim=2") - n_particles = SE2DiracDistribution._validate_particle_count(n_particles) + n_particles = cls._validate_particle_count(n_particles) - return SE2DiracDistribution( + return cls( distribution.sample(n_particles), ones(n_particles) / n_particles, ) diff --git a/src/pyrecest/distributions/se3_dirac_distribution.py b/src/pyrecest/distributions/se3_dirac_distribution.py index 59c1b2b0c..b4db3c648 100644 --- a/src/pyrecest/distributions/se3_dirac_distribution.py +++ b/src/pyrecest/distributions/se3_dirac_distribution.py @@ -34,16 +34,16 @@ def mean(self): m = self.hybrid_mean() return m - @staticmethod - def from_distribution(distribution, n_particles): + @classmethod + def from_distribution(cls, distribution, n_particles): if not isinstance(distribution, AbstractSE3Distribution): raise TypeError( "distribution must be an instance of AbstractSE3Distribution" ) - n_particles = SE3DiracDistribution._validate_particle_count(n_particles) + n_particles = cls._validate_particle_count(n_particles) - ddist = SE3DiracDistribution( + ddist = cls( distribution.sample(n_particles), 1 / n_particles * ones(n_particles), ) diff --git a/tests/distributions/test_dirac_factory_subclass_preservation.py b/tests/distributions/test_dirac_factory_subclass_preservation.py new file mode 100644 index 000000000..f644f4141 --- /dev/null +++ b/tests/distributions/test_dirac_factory_subclass_preservation.py @@ -0,0 +1,86 @@ +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_grid_distribution import ( + CircularGridDistribution, +) +from pyrecest.distributions.circle.von_mises_distribution import VonMisesDistribution +from pyrecest.distributions.conversion import convert_distribution +from pyrecest.distributions.nonperiodic.linear_dirac_distribution import ( + LinearDiracDistribution, +) +from pyrecest.distributions.se2_dirac_distribution import SE2DiracDistribution +from pyrecest.distributions.se3_dirac_distribution import SE3DiracDistribution + + +class _LinearDiracSubclass(LinearDiracDistribution): + pass + + +class _CircularDiracSubclass(CircularDiracDistribution): + pass + + +class _CircularGridSubclass(CircularGridDistribution): + pass + + +class _SE2DiracSubclass(SE2DiracDistribution): + pass + + +class _SE3DiracSubclass(SE3DiracDistribution): + pass + + +class DiracFactorySubclassPreservationTest(unittest.TestCase): + def test_linear_conversion_factory_preserves_requested_subclass(self): + source = LinearDiracDistribution(array([0.0, 1.0])) + + converted = convert_distribution( + source, _LinearDiracSubclass, n_particles=2 + ) + + self.assertIsInstance(converted, _LinearDiracSubclass) + + def test_circular_conversion_factory_preserves_requested_subclass(self): + source = CircularDiracDistribution(array([0.0, 1.0])) + + converted = convert_distribution( + source, _CircularDiracSubclass, n_particles=2 + ) + + self.assertIsInstance(converted, _CircularDiracSubclass) + + def test_circular_grid_conversion_factory_preserves_requested_subclass(self): + source = VonMisesDistribution(0.3, 2.0) + + converted = convert_distribution( + source, _CircularGridSubclass, no_of_gridpoints=9 + ) + + self.assertIsInstance(converted, _CircularGridSubclass) + + def test_se2_conversion_factory_preserves_requested_subclass(self): + source = SE2DiracDistribution(array([[0.0, 1.0, 2.0]])) + + converted = convert_distribution(source, _SE2DiracSubclass, n_particles=2) + + self.assertIsInstance(converted, _SE2DiracSubclass) + + def test_se3_conversion_factory_preserves_requested_subclass(self): + source = SE3DiracDistribution( + array([[1.0, 0.0, 0.0, 0.0, 1.0, 2.0, 3.0]]) + ) + + converted = convert_distribution(source, _SE3DiracSubclass, n_particles=2) + + self.assertIsInstance(converted, _SE3DiracSubclass) + + +if __name__ == "__main__": + unittest.main()