Skip to content
23 changes: 11 additions & 12 deletions src/eegprep/plugins/clean_rawdata/private/covariance.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,50 +34,49 @@
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):
"""Calculate the matrix square root, and its reciprocal, for a covariance matrix or ...,N,N array."""
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)),
)


Expand Down
14 changes: 14 additions & 0 deletions tests/test_utils_covariance.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import unittest
import warnings

import numpy as np
from unittest.mock import patch

Expand Down Expand Up @@ -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
Expand Down
Loading