diff --git a/src/pyrecest/utils/metrics.py b/src/pyrecest/utils/metrics.py index 958b802e5..fc9ca392b 100644 --- a/src/pyrecest/utils/metrics.py +++ b/src/pyrecest/utils/metrics.py @@ -770,6 +770,19 @@ def _validate_positive_semidefinite(matrix: np.ndarray, name: str) -> None: def _as_covariance_matrix(value: ArrayLike, name: str) -> np.ndarray: matrix = _as_numeric_array(value, name) _validate_square_matrix(matrix, name) + if not np.all(np.isfinite(matrix)): + raise ValueError(f"{name} must contain only finite values") + matrix_transpose = matrix.T + matrix_scale = np.maximum( + np.maximum(np.abs(matrix), np.abs(matrix_transpose)), + 1.0, + ) + with np.errstate(over="ignore", under="ignore", invalid="ignore"): + relative_asymmetry = np.abs( + matrix / matrix_scale - matrix_transpose / matrix_scale + ) + if np.any(relative_asymmetry > 1e-12): + raise ValueError(f"{name} must be symmetric") matrix = _symmetrize(matrix) _validate_positive_semidefinite(matrix, name) return matrix diff --git a/tests/test_metrics_wasserstein_validation.py b/tests/test_metrics_wasserstein_validation.py index 1307e9d09..7b722318f 100644 --- a/tests/test_metrics_wasserstein_validation.py +++ b/tests/test_metrics_wasserstein_validation.py @@ -35,6 +35,25 @@ def test_extent_wasserstein_rejects_indefinite_extents(self): ): extent_wasserstein_distance(valid, indefinite) + def test_gaussian_wasserstein_rejects_nonsymmetric_covariances(self): + mean = np.zeros(2) + valid = np.eye(2) + nonsymmetric = np.array([[2.0, 1.0], [0.0, 2.0]]) + + with self.assertRaisesRegex(ValueError, "covariance1 must be symmetric"): + gaussian_wasserstein_distance(mean, nonsymmetric, mean, valid) + with self.assertRaisesRegex(ValueError, "covariance2 must be symmetric"): + gaussian_wasserstein_distance(mean, valid, mean, nonsymmetric) + + def test_extent_wasserstein_rejects_nonsymmetric_extents(self): + valid = np.eye(2) + nonsymmetric = np.array([[2.0, 1.0], [0.0, 2.0]]) + + with self.assertRaisesRegex(ValueError, "estimated_extent must be symmetric"): + extent_wasserstein_distance(nonsymmetric, valid) + with self.assertRaisesRegex(ValueError, "reference_extent must be symmetric"): + extent_wasserstein_distance(valid, nonsymmetric) + if __name__ == "__main__": unittest.main()