Skip to content
Open
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
44 changes: 22 additions & 22 deletions src/pyrecest/distributions/nonperiodic/gaussian_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
54 changes: 54 additions & 0 deletions tests/distributions/test_gaussian_factory_subclass_preservation.py
Original file line number Diff line number Diff line change
@@ -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()
Loading