diff --git a/src/eegprep/plugins/clean_rawdata/private/covariance.py b/src/eegprep/plugins/clean_rawdata/private/covariance.py index cd646640..d5b5d9ac 100644 --- a/src/eegprep/plugins/clean_rawdata/private/covariance.py +++ b/src/eegprep/plugins/clean_rawdata/private/covariance.py @@ -34,41 +34,40 @@ def diag_nd(M): """Like np.diag, but in case of a ...,N, returns a ...,N,N array of diag matrices.""" *dims, N = M.shape - if dims: - cat = np.concatenate([np.diag(d) for d in M.reshape((-1, N))]) - return np.reshape(cat, dims + [N, N]) - else: - return np.diag(M) + res = np.zeros((*dims, N, N), dtype=M.dtype) + i = np.arange(N) + res[..., i, i] = M + return res def cov_logm(C): """Calculate the matrix logarithm of a covariance matrix or ...,N,N array.""" D, V = np.linalg.eigh(C) - return finite_matmul(finite_matmul(V, diag_nd(np.log(D))), V.swapaxes(-2, -1)) + return finite_matmul(V * np.log(D)[..., np.newaxis, :], V.swapaxes(-2, -1)) def cov_expm(C): """Calculate the matrix exponent of a covariance matrix or ...,N,N array.""" D, V = np.linalg.eigh(C) - return finite_matmul(finite_matmul(V, diag_nd(np.exp(D))), V.swapaxes(-2, -1)) + return finite_matmul(V * np.exp(D)[..., np.newaxis, :], V.swapaxes(-2, -1)) def cov_powm(C, exp): """Calculate a matrix power of a covariance matrix or ...,N,N array.""" D, V = np.linalg.eigh(C) - return finite_matmul(finite_matmul(V, diag_nd(D**exp)), V.swapaxes(-2, -1)) + return finite_matmul(V * (D**exp)[..., np.newaxis, :], V.swapaxes(-2, -1)) def cov_sqrtm(C): """Calculate the matrix square root of a covariance matrix or ...,N,N array.""" D, V = np.linalg.eigh(C) - return finite_matmul(finite_matmul(V, diag_nd(np.sqrt(D))), V.swapaxes(-2, -1)) + return finite_matmul(V * np.sqrt(D)[..., np.newaxis, :], V.swapaxes(-2, -1)) def cov_rsqrtm(C): """Calculate the matrix reciprocal square root of a covariance matrix or ...,N,N array.""" D, V = np.linalg.eigh(C) - return finite_matmul(finite_matmul(V, diag_nd(1.0 / np.sqrt(D))), V.swapaxes(-2, -1)) + return finite_matmul(V * (1.0 / np.sqrt(D))[..., np.newaxis, :], V.swapaxes(-2, -1)) def cov_sqrtm2(C): @@ -76,8 +75,8 @@ def cov_sqrtm2(C): D, V = np.linalg.eigh(C) sqrtD = np.sqrt(D) return ( - finite_matmul(finite_matmul(V, diag_nd(sqrtD)), V.swapaxes(-2, -1)), - finite_matmul(finite_matmul(V, diag_nd(1.0 / sqrtD)), V.swapaxes(-2, -1)), + finite_matmul(V * sqrtD[..., np.newaxis, :], V.swapaxes(-2, -1)), + finite_matmul(V * (1.0 / sqrtD)[..., np.newaxis, :], V.swapaxes(-2, -1)), ) diff --git a/tests/test_utils_covariance.py b/tests/test_utils_covariance.py index ce397cd3..b9504536 100644 --- a/tests/test_utils_covariance.py +++ b/tests/test_utils_covariance.py @@ -1,4 +1,6 @@ import unittest +import warnings + import numpy as np from unittest.mock import patch @@ -391,6 +393,18 @@ def test_positive_definite_preservation(self): class TestEdgeCases(unittest.TestCase): """Test edge cases and numerical stability.""" + def test_singular_matrix_warning_behavior(self): + singular = np.diag([1.0, 0.0]) + + with self.assertWarnsRegex(RuntimeWarning, "divide by zero"): + cov_logm(singular) + + with warnings.catch_warnings(): + warnings.simplefilter("error", RuntimeWarning) + sqrt_result = cov_sqrtm(singular) + + np.testing.assert_array_equal(sqrt_result, singular) + def test_near_singular_matrices(self): """Test operations on near-singular matrices.""" # Create a matrix with very small eigenvalues