spd_learn.functional.shrinkage_covariance#

spd_learn.functional.shrinkage_covariance(X: Tensor, alpha: Tensor, n_chans: int, identity: Tensor | None = None) → Tensor[source]#

Apply shrinkage regularization to covariance matrices.

Computes the shrinkage estimator:

\[\hat{C} = (1 - \alpha) C + \alpha \cdot \frac{\text{tr}(C)}{n} \cdot I_n\]

This convex combination interpolates between the empirical covariance \(C\) and a scaled identity matrix.

Parameters:
  • X (Tensor) – Batch of covariance matrices with shape (…, n_chans, n_chans).

  • alpha (Tensor) – Shrinkage intensity in range [0, 1]. Can be a scalar or broadcastable with X.

  • n_chans (int) – Number of channels (matrix dimension).

  • identity (Tensor, optional) – Pre-computed identity matrix. If None, one is created.

Returns:

Regularized covariance matrices with shape (…, n_chans, n_chans).

Return type:

Tensor

Notes

For \(\alpha = 0\), returns the original covariance. For \(\alpha = 1\), returns a scaled identity matrix.

See also

ledoit_wolf()

Alternative shrinkage formulation.

trace_normalization()

Trace-based normalization.

Shrinkage

Module wrapper for this function.

Examples

>>> import torch
>>> from spd_learn.functional import shrinkage_covariance
>>> X = torch.randn(4, 8, 8)
>>> X = X @ X.mT  # Make SPD
>>> alpha = torch.tensor(0.5)
>>> X_shrunk = shrinkage_covariance(X, alpha, n_chans=8)