From c7fbad693d53107a4c4f2eb6b54d01364a2c6e8d Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Wed, 19 Aug 2026 20:03:17 +0800 Subject: [PATCH 1/2] Preserve Gaussian target subclasses in conversion factory --- .../nonperiodic/gaussian_distribution.py | 44 +++++++++---------- 1 file changed, 22 insertions(+), 22 deletions(-) diff --git a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py index 2fb9ba5dfb..f64225f675 100644 --- a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py +++ b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py @@ -391,8 +391,8 @@ def sample(self, n): n = _validate_positive_sample_count(n) return random.multivariate_normal(mean=self.mu, cov=self.C, size=n) - @staticmethod - def from_distribution(distribution, check_validity=False): + @classmethod + def from_distribution(cls, distribution, check_validity=False): """Approximate or convert another distribution as a Gaussian. Gaussian mixtures are converted with ``to_gaussian``. Other @@ -404,23 +404,23 @@ def from_distribution(distribution, check_validity=False): if isinstance(distribution, GaussianMixture): gaussian = distribution.to_gaussian(check_validity=check_validity) - else: - try: - mean = distribution.mean - covariance = distribution.covariance - except AttributeError as exc: - raise ConversionError( - "GaussianDistribution.from_distribution requires the source " - "distribution to expose mean() and covariance()." - ) from exc - - if callable(mean): - mean = mean() - - if callable(covariance): - covariance = covariance() - - gaussian = GaussianDistribution( - mean, covariance, check_validity=check_validity - ) - return gaussian + if cls is GaussianDistribution: + return gaussian + return cls(gaussian.mu, gaussian.C, check_validity=check_validity) + + try: + mean = distribution.mean + covariance = distribution.covariance + except AttributeError as exc: + raise ConversionError( + "GaussianDistribution.from_distribution requires the source " + "distribution to expose mean() and covariance()." + ) from exc + + if callable(mean): + mean = mean() + + if callable(covariance): + covariance = covariance() + + return cls(mean, covariance, check_validity=check_validity) From 2b886f3886a299f791946548cb75d40b80884f2d Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Wed, 19 Aug 2026 20:03:51 +0800 Subject: [PATCH 2/2] Add Gaussian conversion subclass regressions --- ..._gaussian_factory_subclass_preservation.py | 54 +++++++++++++++++++ 1 file changed, 54 insertions(+) create mode 100644 tests/distributions/test_gaussian_factory_subclass_preservation.py diff --git a/tests/distributions/test_gaussian_factory_subclass_preservation.py b/tests/distributions/test_gaussian_factory_subclass_preservation.py new file mode 100644 index 0000000000..d25e462498 --- /dev/null +++ b/tests/distributions/test_gaussian_factory_subclass_preservation.py @@ -0,0 +1,54 @@ +import unittest + +from pyrecest.backend import allclose, array +from pyrecest.distributions.conversion import convert_distribution +from pyrecest.distributions.nonperiodic.gaussian_distribution import ( + GaussianDistribution, +) +from pyrecest.distributions.nonperiodic.gaussian_mixture import GaussianMixture + + +class GaussianSubclass(GaussianDistribution): + pass + + +class GaussianFactorySubclassPreservationTest(unittest.TestCase): + def test_conversion_to_gaussian_subclass_preserves_requested_type(self): + source = GaussianDistribution( + array([1.0, -2.0]), + array([[2.0, 0.25], [0.25, 1.0]]), + ) + + converted = convert_distribution( + source, + GaussianSubclass, + check_validity=True, + ) + + self.assertIs(type(converted), GaussianSubclass) + self.assertTrue(bool(allclose(converted.mu, source.mu))) + self.assertTrue(bool(allclose(converted.C, source.C))) + + def test_mixture_conversion_to_gaussian_subclass_preserves_requested_type(self): + mixture = GaussianMixture( + [ + GaussianDistribution(array([0.0]), array([[1.0]])), + GaussianDistribution(array([2.0]), array([[3.0]])), + ], + array([0.25, 0.75]), + ) + expected = mixture.to_gaussian(check_validity=True) + + converted = convert_distribution( + mixture, + GaussianSubclass, + check_validity=True, + ) + + self.assertIs(type(converted), GaussianSubclass) + self.assertTrue(bool(allclose(converted.mu, expected.mu))) + self.assertTrue(bool(allclose(converted.C, expected.C))) + + +if __name__ == "__main__": + unittest.main()