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.
ShrinkageModule 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)