diff --git a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py index f64225f67..448e10d1b 100644 --- a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py +++ b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py @@ -325,7 +325,7 @@ def multiply(self, other): new_mu = linalg.solve(new_precision, information_vector) new_C = linalg.solve(new_precision, identity) new_C = 0.5 * (new_C + transpose(new_C)) - return GaussianDistribution(new_mu, new_C, check_validity=False) + return type(self)(new_mu, new_C, check_validity=False) def convolve(self, other): """Convolve two independent Gaussian distributions. @@ -343,7 +343,7 @@ def convolve(self, other): _validate_same_dimension(self, other, "convolve") new_mu = self.mu + other.mu new_C = self.C + other.C - return GaussianDistribution(new_mu, new_C, check_validity=False) + return type(self)(new_mu, new_C, check_validity=False) def marginalize_out(self, dimensions): """Return the marginal distribution after dropping dimensions. diff --git a/tests/distributions/test_gaussian_distribution_subclass_operations.py b/tests/distributions/test_gaussian_distribution_subclass_operations.py new file mode 100644 index 000000000..25bc6cf31 --- /dev/null +++ b/tests/distributions/test_gaussian_distribution_subclass_operations.py @@ -0,0 +1,30 @@ +"""Regression tests for Gaussian subclass preservation.""" + +import unittest + +from pyrecest.backend import array +from pyrecest.distributions import GaussianDistribution + + +class _DerivedGaussian(GaussianDistribution): + """Minimal derived Gaussian used to verify operation return types.""" + + +class TestGaussianDistributionSubclassOperations(unittest.TestCase): + def setUp(self): + self.derived = _DerivedGaussian(array([0.0]), array([[2.0]])) + self.other = GaussianDistribution(array([1.0]), array([[3.0]])) + + def test_multiply_preserves_left_hand_subclass(self): + result = self.derived.multiply(self.other) + + self.assertIsInstance(result, _DerivedGaussian) + + def test_convolve_preserves_left_hand_subclass(self): + result = self.derived.convolve(self.other) + + self.assertIsInstance(result, _DerivedGaussian) + + +if __name__ == "__main__": + unittest.main()