Skip to content
Closed
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
24 changes: 23 additions & 1 deletion src/pyrecest/distributions/nonperiodic/gaussian_mixture.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
54 changes: 54 additions & 0 deletions tests/distributions/test_gaussian_mixture_parameter_shapes.py
Original file line number Diff line number Diff line change
@@ -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()
Loading