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
3 changes: 3 additions & 0 deletions .jules/bolt.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
## 2025-05-15 - [Optimization of batched covariance functions]
**Learning:** Replacing intermediate diagonal matrix creation and full matrix multiplication with broadcasting scaling (`V * D[..., np.newaxis, :]`) significantly improves performance by reducing memory allocation and operation complexity from $O(N^3)$ to $O(N^2)$ for the scaling step. Batched diagonal matrix creation can also be optimized using advanced indexing instead of loops.
**Action:** Always look for diagonal matrix multiplications and replace them with vectorized broadcasting when applicable. For batched operations, avoid loops and `np.concatenate` in favor of pre-allocated arrays and advanced indexing.
23 changes: 12 additions & 11 deletions src/eegprep/plugins/clean_rawdata/private/covariance.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,50 +34,51 @@
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:
if not dims:
return np.diag(M)
res = np.zeros((*dims, N, N), dtype=M.dtype)
idx = np.arange(N)
res[..., idx, idx] = 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
Loading