diff --git a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py index 448e10d1b..38d58cd18 100644 --- a/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py +++ b/src/pyrecest/distributions/nonperiodic/gaussian_distribution.py @@ -384,7 +384,7 @@ def marginalize_out(self, dimensions): new_C = self.C[remaining_indices][ :, remaining_indices ] # Instead of np.ix_ for interface compatibility - return GaussianDistribution(new_mu, new_C, check_validity=False) + return type(self)(new_mu, new_C, check_validity=False) def sample(self, n): """Draw ``n`` random samples with shape ``(n, dim)``.""" diff --git a/tests/distributions/test_gaussian_distribution_subclass_marginalization.py b/tests/distributions/test_gaussian_distribution_subclass_marginalization.py new file mode 100644 index 000000000..2ce395342 --- /dev/null +++ b/tests/distributions/test_gaussian_distribution_subclass_marginalization.py @@ -0,0 +1,27 @@ +"""Regression test for Gaussian subclass preservation during marginalization.""" + +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 TestGaussianDistributionSubclassMarginalization(unittest.TestCase): + def test_marginalize_out_preserves_subclass(self): + derived = _DerivedGaussian( + array([1.0, 2.0]), + array([[2.0, 0.5], [0.5, 3.0]]), + ) + + result = derived.marginalize_out(1) + + self.assertIsInstance(result, _DerivedGaussian) + self.assertEqual(result.dim, 1) + + +if __name__ == "__main__": + unittest.main()