From 7f2a130d7bb1644bd1d1341c8cc8a7ef8e3a1b38 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Tue, 4 Aug 2026 19:32:44 +0800 Subject: [PATCH 1/2] Validate Gaussian mixture parameter shapes --- .../nonperiodic/gaussian_mixture.py | 24 ++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py b/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py index f392b1a14e..8a1243835b 100644 --- a/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py +++ b/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py @@ -66,8 +66,30 @@ def covariance(self): def mixture_parameters_to_gaussian_parameters( means, covariance_matrices, weights=None ): + means = array(means) + covariance_matrices = array(covariance_matrices) + + if means.ndim == 1: + n_components = means.shape[0] + state_dim = 1 + elif means.ndim == 2: + n_components, state_dim = means.shape + else: + raise ValueError( + "means must have shape (n_components,) or (n_components, state_dim)" + ) + if n_components == 0 or state_dim == 0: + raise ValueError("means must contain at least one nonempty component mean") + + expected_covariance_shape = (state_dim, state_dim, n_components) + if covariance_matrices.shape != expected_covariance_shape: + raise ValueError( + "covariance_matrices must have shape " + f"{expected_covariance_shape}, got {covariance_matrices.shape}" + ) + if weights is None: - weights = ones(means.shape[0]) / means.shape[0] + weights = ones(n_components) / n_components else: weights = array(weights) if weights.ndim == 0: From bc8d5788ac222a1c0b72187c01d43a13bb3e2c7b Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Tue, 4 Aug 2026 19:32:58 +0800 Subject: [PATCH 2/2] Test Gaussian mixture parameter shape validation --- .../test_gaussian_mixture_parameter_shapes.py | 54 +++++++++++++++++++ 1 file changed, 54 insertions(+) create mode 100644 tests/distributions/test_gaussian_mixture_parameter_shapes.py diff --git a/tests/distributions/test_gaussian_mixture_parameter_shapes.py b/tests/distributions/test_gaussian_mixture_parameter_shapes.py new file mode 100644 index 0000000000..b722f8ab57 --- /dev/null +++ b/tests/distributions/test_gaussian_mixture_parameter_shapes.py @@ -0,0 +1,54 @@ +import unittest + +import numpy as np +import numpy.testing as npt + +# pylint: disable=no-name-in-module +from pyrecest.backend import array, to_numpy +from pyrecest.distributions.nonperiodic.gaussian_mixture import GaussianMixture + + +class GaussianMixtureParameterShapeTest(unittest.TestCase): + def test_rejects_missing_component_covariance(self): + means = array([0.0, 2.0]) + covariance_matrices = array([[[1.0]]]) + + with self.assertRaisesRegex( + ValueError, + "covariance_matrices must have shape", + ): + GaussianMixture.mixture_parameters_to_gaussian_parameters( + means, + covariance_matrices, + array([0.25, 0.75]), + ) + + def test_rejects_covariance_without_component_axis(self): + means = array([[0.0, 0.0], [1.0, 1.0]]) + covariance_matrices = array([[1.0, 0.0], [0.0, 1.0]]) + + with self.assertRaisesRegex( + ValueError, + "covariance_matrices must have shape", + ): + GaussianMixture.mixture_parameters_to_gaussian_parameters( + means, + covariance_matrices, + array([0.5, 0.5]), + ) + + def test_matching_component_shapes_preserve_moment_matching(self): + mean, covariance = ( + GaussianMixture.mixture_parameters_to_gaussian_parameters( + array([0.0, 2.0]), + array([[[1.0, 3.0]]]), + array([0.25, 0.75]), + ) + ) + + npt.assert_allclose(to_numpy(mean), np.array([1.5])) + npt.assert_allclose(to_numpy(covariance), np.array([[3.25]])) + + +if __name__ == "__main__": + unittest.main()