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
12 changes: 5 additions & 7 deletions src/pyrecest/distributions/circle/circular_dirac_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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
)

Expand Down
12 changes: 6 additions & 6 deletions src/pyrecest/distributions/circle/circular_grid_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand All @@ -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(
Expand Down
8 changes: 4 additions & 4 deletions src/pyrecest/distributions/se2_dirac_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
)
Expand Down
8 changes: 4 additions & 4 deletions src/pyrecest/distributions/se3_dirac_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
)
Expand Down
86 changes: 86 additions & 0 deletions tests/distributions/test_dirac_factory_subclass_preservation.py
Original file line number Diff line number Diff line change
@@ -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()
Loading