Source code for yetanotherspdnet.functions.m_estimators

r"""Robust covariance estimation: sample covariance and M-estimators.

An M-estimator of scatter is a fixed point of

.. math::

    \Sigma = F(\Sigma) = \frac{1}{n} \sum_{i=1}^{n}
        u\big(x_i^\top \Sigma^{-1} x_i\big)\, x_i x_i^\top

for a weight function :math:`u` that down-weights samples with a large
Mahalanobis distance: Tyler (:math:`u(q) = p/q`), Student-t
(:math:`u(q) = (p + \nu)/(\nu + q)`) or Huber. The sample covariance matrix
corresponds to :math:`u \equiv 1`.

Two gradient paths are available, as elsewhere in the library:

- :func:`m_estimator` unrolls the fixed-point iterations and lets autograd
  differentiate through them (memory grows with the number of iterations);
- :class:`MEstimator` differentiates the fixed point implicitly: the backward
  solves the adjoint equation :math:`w = g + J_\Sigma^\top w` by iteration and
  returns :math:`J_X^\top w`, with memory independent of the number of
  iterations.

Tyler's weight is scale-invariant (:math:`F(c\Sigma) = cF(\Sigma)`), so its
fixed point is only defined up to scale: use it with ``normalize="trace"`` or
``normalize="determinant"``, which is applied at every iteration and pins the
scale down.
"""

import math
from collections.abc import Callable
from functools import partial

import torch
from torch.autograd import Function


# ----------------
# Weight functions
# ----------------
[docs] def tyler_function(quadratic: torch.Tensor, n_features: int) -> torch.Tensor: r""" Tyler weight :math:`u(q) = p / q`. Parameters ---------- quadratic : torch.Tensor of shape (..., n_samples) Squared Mahalanobis distances :math:`q_i = x_i^\top \Sigma^{-1} x_i` n_features : int Dimension :math:`p` of the samples Returns ------- weights : torch.Tensor of shape (..., n_samples) Sample weights """ return n_features / quadratic
[docs] def student_function( quadratic: torch.Tensor, n_features: int, nu: float ) -> torch.Tensor: r""" Student-t weight :math:`u(q) = (p + \nu) / (\nu + q)`. Maximum-likelihood weight for a multivariate Student-t distribution with :math:`\nu` degrees of freedom; tends to the sample covariance when :math:`\nu \to \infty` and to Tyler's weight when :math:`\nu \to 0`. Parameters ---------- quadratic : torch.Tensor of shape (..., n_samples) Squared Mahalanobis distances n_features : int Dimension :math:`p` of the samples nu : float Degrees of freedom Returns ------- weights : torch.Tensor of shape (..., n_samples) Sample weights """ return (n_features + nu) / (nu + quadratic)
[docs] def huber_function(quadratic: torch.Tensor, delta: float, beta: float) -> torch.Tensor: r""" Huber weight: :math:`u(q) = 1/\beta` if :math:`q \le \delta`, else :math:`\delta / (\beta q)`. Parameters ---------- quadratic : torch.Tensor of shape (..., n_samples) Squared Mahalanobis distances delta : float Threshold above which samples are down-weighted beta : float Scaling factor (makes the estimator consistent for Gaussian data) Returns ------- weights : torch.Tensor of shape (..., n_samples) Sample weights """ return torch.where(quadratic <= delta, 1 / beta, delta / (beta * quadratic))
# ------------- # Normalization # -------------
[docs] def normalize_trace(data: torch.Tensor) -> torch.Tensor: r""" Scale SPD matrices so that :math:`\operatorname{tr}(\Sigma) = p`. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) SPD matrices Returns ------- normalized : torch.Tensor of shape (..., n_features, n_features) SPD matrices with trace equal to ``n_features`` """ trace = torch.diagonal(data, dim1=-2, dim2=-1).sum(dim=-1) return data.shape[-1] * data / trace[..., None, None]
[docs] def normalize_determinant(data: torch.Tensor) -> torch.Tensor: r""" Scale SPD matrices so that :math:`\det(\Sigma) = 1`. The determinant is computed through its logarithm to avoid overflow. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) SPD matrices Returns ------- normalized : torch.Tensor of shape (..., n_features, n_features) SPD matrices with unit determinant """ logdet = torch.linalg.slogdet(data).logabsdet return data * torch.exp(-logdet / data.shape[-1])[..., None, None]
_NORMALIZATIONS: dict[str, Callable] = { "trace": normalize_trace, "determinant": normalize_determinant, } def _get_normalization(normalize: str | None) -> Callable | None: if normalize is None: return None if normalize not in _NORMALIZATIONS: raise ValueError( f"normalize must be None or one of {list(_NORMALIZATIONS)}, got {normalize}" ) return _NORMALIZATIONS[normalize] # ----------------- # Sample covariance # -----------------
[docs] def sample_covariance( data: torch.Tensor, assume_centered: bool = False ) -> torch.Tensor: r""" Sample covariance matrix (SCM). .. math:: \hat\Sigma = \frac{1}{n'} \sum_{i=1}^{n} (x_i - \bar x)(x_i - \bar x)^\top with :math:`n' = n - 1` when the data are centered here, :math:`n'= n` (and :math:`\bar x = 0`) when ``assume_centered``. Parameters ---------- data : torch.Tensor of shape (..., n_samples, n_features) Samples assume_centered : bool, optional Whether the data are already centered. Default is False Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Sample covariance matrices """ n_samples = data.shape[-2] if not assume_centered: data = data - data.mean(dim=-2, keepdim=True) n_samples = n_samples - 1 covariance = data.transpose(-2, -1) @ data / n_samples return 0.5 * (covariance + covariance.transpose(-2, -1))
# ------------- # M-estimators # -------------
[docs] def m_estimator_step( covariance: torch.Tensor, data: torch.Tensor, weight_function: Callable, shrinkage: float | None = None, normalize: Callable | None = None, ) -> torch.Tensor: r""" One fixed-point iteration :math:`\Sigma \mapsto F(\Sigma)` of an M-estimator. .. math:: F(\Sigma) = \frac{1}{n} \sum_{i=1}^{n} u\big(x_i^\top \Sigma^{-1} x_i\big)\, x_i x_i^\top followed, if requested, by the shrinkage :math:`\beta F(\Sigma) + (1 - \beta) I` and by a normalization. The Mahalanobis distances are computed with a Cholesky factorization. Parameters ---------- covariance : torch.Tensor of shape (..., n_features, n_features) Current estimate data : torch.Tensor of shape (..., n_samples, n_features) Centered samples weight_function : Callable Weight :math:`u`, called on the tensor of squared distances of shape ``(..., n_samples)`` shrinkage : float | None, optional Shrinkage coefficient :math:`\beta \in (0, 1]` towards the identity. Default is None (no shrinkage) normalize : Callable | None, optional Normalization applied to the result (e.g. :func:`normalize_trace`). Default is None Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Updated estimate """ cholesky = torch.linalg.cholesky(covariance) whitened = torch.linalg.solve_triangular( cholesky, data.transpose(-2, -1), upper=False ) # L^{-1} x_i as columns quadratic = (whitened**2).sum(dim=-2) weights = weight_function(quadratic) weighted = data * weights.unsqueeze(-1) updated = weighted.transpose(-2, -1) @ data / data.shape[-2] updated = 0.5 * (updated + updated.transpose(-2, -1)) if shrinkage is not None: eye = torch.eye(data.shape[-1], dtype=data.dtype, device=data.device) updated = shrinkage * updated + (1 - shrinkage) * eye if normalize is not None: updated = normalize(updated) return updated
def _initial_covariance(data: torch.Tensor, init: torch.Tensor | None) -> torch.Tensor: n_features = data.shape[-1] batch_shape = data.shape[:-2] if init is None: init = torch.eye(n_features, dtype=data.dtype, device=data.device) if init.shape[-1] != n_features: raise ValueError( f"init of size {tuple(init.shape)} incompatible with data " f"of size {tuple(data.shape)}" ) return init.expand(*batch_shape, n_features, n_features) def _relative_change(new: torch.Tensor, old: torch.Tensor) -> torch.Tensor: """Largest relative Frobenius change over the batch.""" return (torch.linalg.matrix_norm(new - old) / torch.linalg.matrix_norm(old)).max()
[docs] def m_estimator( data: torch.Tensor, weight_function: Callable, n_iterations: int = 30, tol: float = 1e-6, assume_centered: bool = False, init: torch.Tensor | None = None, shrinkage: float | None = None, normalize: str | None = None, ) -> torch.Tensor: r""" M-estimator of scatter by fixed-point iterations (autograd path). Iterates :func:`m_estimator_step` from ``init`` until the largest relative Frobenius change over the batch falls below ``tol`` or ``n_iterations`` is reached. Gradients flow through the unrolled iterations. Parameters ---------- data : torch.Tensor of shape (..., n_samples, n_features) Samples weight_function : Callable Weight :math:`u` of squared Mahalanobis distances, e.g. ``functools.partial(student_function, n_features=p, nu=3.0)`` n_iterations : int, optional Maximum number of iterations. Default is 30 tol : float, optional Stopping tolerance on the relative change. Default is 1e-6 assume_centered : bool, optional Whether the data are already centered. Default is False init : torch.Tensor of shape (n_features, n_features) or (..., n_features, n_features), optional Initial estimate. Default is the identity shrinkage : float | None, optional Shrinkage coefficient towards the identity, applied at each iteration. Default is None normalize : str | None, optional ``"trace"``, ``"determinant"`` or None, applied at each iteration. Required for scale-invariant weights such as Tyler's. Default is None Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Estimated scatter matrices """ normalization = _get_normalization(normalize) if not assume_centered: data = data - data.mean(dim=-2, keepdim=True) covariance = _initial_covariance(data, init) for _ in range(n_iterations): updated = m_estimator_step( covariance, data, weight_function, shrinkage, normalization ) converged = _relative_change(updated.detach(), covariance.detach()) < tol covariance = updated if converged: break return covariance
[docs] class MEstimator(Function): r""" M-estimator of scatter with an implicit (fixed-point) backward. The forward pass iterates without building a graph. At the fixed point :math:`\Sigma^\star = F(\Sigma^\star, X)`, the implicit function theorem gives :math:`\partial L / \partial X = J_X^\top w` where :math:`w` solves :math:`w = g + J_\Sigma^\top w` (:math:`g` the incoming gradient). The adjoint equation is solved by fixed-point iteration, which converges at the rate of the forward iteration since :math:`J_\Sigma` has spectral radius below one at an attracting fixed point. The vector-Jacobian products of one step are obtained with autograd. Use as ``MEstimator.apply(data, weight_function, n_iterations, tol, assume_centered, init, shrinkage, normalize)`` with the arguments of :func:`m_estimator`. """
[docs] @staticmethod def forward( ctx, data: torch.Tensor, weight_function: Callable, n_iterations: int = 30, tol: float = 1e-6, assume_centered: bool = False, init: torch.Tensor | None = None, shrinkage: float | None = None, normalize: str | None = None, ) -> torch.Tensor: """ Forward pass: fixed-point iterations without graph Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to save tensors for the backward pass data, weight_function, n_iterations, tol, assume_centered, init, shrinkage, normalize See :func:`m_estimator` Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Estimated scatter matrices """ with torch.no_grad(): covariance = m_estimator( data, weight_function, n_iterations, tol, assume_centered, init, shrinkage, normalize, ) ctx.save_for_backward(data, covariance) ctx.weight_function = weight_function ctx.n_iterations = n_iterations ctx.tol = tol ctx.assume_centered = assume_centered ctx.shrinkage = shrinkage ctx.normalization = _get_normalization(normalize) return covariance
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple: """ Backward pass: implicit differentiation at the fixed point Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object with the saved tensors grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the estimate Returns ------- grad_data : torch.Tensor of shape (..., n_samples, n_features) Gradient of the loss with respect to the samples; None for the other arguments """ data, covariance = ctx.saved_tensors with torch.enable_grad(): data_leaf = data.detach().requires_grad_(True) covariance_leaf = covariance.detach().requires_grad_(True) centered = ( data_leaf if ctx.assume_centered else data_leaf - data_leaf.mean(dim=-2, keepdim=True) ) fixed_point = m_estimator_step( covariance_leaf, centered, ctx.weight_function, ctx.shrinkage, ctx.normalization, ) grad_output = 0.5 * (grad_output + grad_output.transpose(-2, -1)) adjoint = grad_output for _ in range(ctx.n_iterations): (vjp_covariance,) = torch.autograd.grad( fixed_point, covariance_leaf, adjoint, retain_graph=True ) updated = grad_output + vjp_covariance converged = ( _relative_change(updated, adjoint) if torch.linalg.matrix_norm(adjoint).max() > 0 else torch.tensor(0.0) ) < ctx.tol adjoint = updated if converged: break (grad_data,) = torch.autograd.grad(fixed_point, data_leaf, adjoint) return grad_data, None, None, None, None, None, None, None
[docs] def tyler_estimator( data: torch.Tensor, n_iterations: int = 30, tol: float = 1e-6, assume_centered: bool = False, normalize: str = "trace", use_autograd: bool = False, ) -> torch.Tensor: r""" Tyler's M-estimator of scatter, normalized to fix its scale. Parameters ---------- data : torch.Tensor of shape (..., n_samples, n_features) Samples (at least ``n_features + 1`` of them) n_iterations : int, optional Maximum number of iterations. Default is 30 tol : float, optional Stopping tolerance. Default is 1e-6 assume_centered : bool, optional Whether the data are already centered. Default is False normalize : str, optional ``"trace"`` or ``"determinant"``. Default is ``"trace"`` use_autograd : bool, optional Unrolled autograd path (True) or implicit backward (False, default) Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Tyler estimates """ if normalize is None: raise ValueError("Tyler's estimator is scale-invariant: normalize is required") weight = partial(tyler_function, n_features=data.shape[-1]) if use_autograd: return m_estimator( data, weight, n_iterations, tol, assume_centered, normalize=normalize ) return MEstimator.apply( data, weight, n_iterations, tol, assume_centered, None, None, normalize )
[docs] def student_estimator( data: torch.Tensor, nu: float, n_iterations: int = 30, tol: float = 1e-6, assume_centered: bool = False, use_autograd: bool = False, ) -> torch.Tensor: r""" Student-t M-estimator of scatter. Parameters ---------- data : torch.Tensor of shape (..., n_samples, n_features) Samples nu : float Degrees of freedom (:math:`\nu > 0`) n_iterations : int, optional Maximum number of iterations. Default is 30 tol : float, optional Stopping tolerance. Default is 1e-6 assume_centered : bool, optional Whether the data are already centered. Default is False use_autograd : bool, optional Unrolled autograd path (True) or implicit backward (False, default) Returns ------- covariance : torch.Tensor of shape (..., n_features, n_features) Student-t estimates """ if not nu > 0 or math.isinf(nu): raise ValueError(f"nu must be a finite positive number, got {nu}") weight = partial(student_function, n_features=data.shape[-1], nu=nu) if use_autograd: return m_estimator(data, weight, n_iterations, tol, assume_centered) return MEstimator.apply( data, weight, n_iterations, tol, assume_centered, None, None, None )