diff --git a/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py b/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py index f392b1a14..f3e05ea8a 100644 --- a/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py +++ b/src/pyrecest/distributions/nonperiodic/gaussian_mixture.py @@ -66,8 +66,42 @@ def covariance(self): def mixture_parameters_to_gaussian_parameters( means, covariance_matrices, weights=None ): + means = array(means) + if means.ndim == 0: + means = reshape(means, (1, 1)) + elif means.ndim == 1: + means = reshape(means, (-1, 1)) + elif means.ndim != 2: + raise ValueError( + "means must have shape (n_components, dim) or be a scalar/1D " + "sequence of one-dimensional component means" + ) + + n_components, dim = means.shape + covariance_matrices = array(covariance_matrices) + expected_shape = (dim, dim, n_components) + + if covariance_matrices.ndim == 3: + if covariance_matrices.shape != expected_shape: + raise ValueError( + "covariance_matrices must have shape " + f"{expected_shape}, got {covariance_matrices.shape}" + ) + elif n_components == 1 and covariance_matrices.shape == (dim, dim): + covariance_matrices = reshape(covariance_matrices, expected_shape) + elif dim == 1 and covariance_matrices.shape == (n_components,): + covariance_matrices = reshape(covariance_matrices, expected_shape) + elif dim == 1 and n_components == 1 and covariance_matrices.ndim == 0: + covariance_matrices = reshape(covariance_matrices, expected_shape) + else: + raise ValueError( + "covariance_matrices must have shape " + f"{expected_shape}; a single ({dim}, {dim}) matrix is only " + "accepted for one component" + ) + 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: @@ -79,7 +113,9 @@ def mixture_parameters_to_gaussian_parameters( mu, C_from_means = LinearDiracDistribution.weighted_samples_to_mean_and_cov( means, weights ) - C_from_cov = sum(covariance_matrices * reshape(weights, (1, 1, -1)), axis=2) + C_from_cov = sum( + covariance_matrices * reshape(weights, (1, 1, -1)), axis=2 + ) C = C_from_cov + C_from_means return mu, C diff --git a/tests/distributions/test_gaussian_mixture_constructor.py b/tests/distributions/test_gaussian_mixture_constructor.py index f24e962c8..d3b1dfb4c 100644 --- a/tests/distributions/test_gaussian_mixture_constructor.py +++ b/tests/distributions/test_gaussian_mixture_constructor.py @@ -36,6 +36,28 @@ def test_parameter_conversion_accepts_sequence_weights(self): np.testing.assert_allclose(to_numpy(mean), [1.5]) np.testing.assert_allclose(to_numpy(covariance), [[3.25]]) + def test_parameter_conversion_accepts_single_covariance_matrix(self): + means = array([[1.0, -2.0]]) + covariance_matrix = array([[4.0, 1.5], [1.5, 9.0]]) + + mean, covariance = ( + GaussianMixture.mixture_parameters_to_gaussian_parameters( + means, covariance_matrix, [1.0] + ) + ) + + np.testing.assert_allclose(to_numpy(mean), [1.0, -2.0]) + np.testing.assert_allclose(to_numpy(covariance), to_numpy(covariance_matrix)) + + def test_parameter_conversion_rejects_ambiguous_covariance_matrix(self): + means = array([[0.0, 0.0], [1.0, 1.0]]) + covariance_matrix = 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_matrix, [0.5, 0.5] + ) + if __name__ == "__main__": unittest.main()