Source code for yetanotherspdnet.functions.spd_linalg

"""Core SPD matrix linear algebra: eigendecomposition, matrix functions, congruence, and vectorization."""

from collections.abc import Callable

import torch
from torch.autograd import Function

from yetanotherspdnet.functions.scalar_functions import (
    inv,
    inv_scaled_softplus,
    inv_scaled_softplus_derivative,
    inv_sqrt,
    inv_sqrt_derivative,
    scaled_softplus,
    scaled_softplus_derivative,
    sqrt_derivative,
)


[docs] def symmetrize(data: torch.Tensor) -> torch.Tensor: r""" Symmetrize a tensor along the last two dimensions. .. math:: \operatorname{sym}(A) = \frac{1}{2}\big(A + A^\top\big) Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of square matrices Returns ------- sym_data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetrized matrices """ return torch.real(0.5 * (data + data.transpose(-1, -2)))
# ----------------------- # Vectorization operators # -----------------------
[docs] def vec_batch(data: torch.Tensor) -> torch.Tensor: """ Vectorize a batch of tensors along last two dimensions Parameters ---------- data : torch.Tensor of shape (..., n_rows, n_columns) Batch of matrices Returns ------- data_vec : torch.Tensor of shape (..., n_rows*n_columns) Batch of vectorized matrices """ return data.reshape(*data.shape[:-2], -1)
[docs] def unvec_batch(data_vec: torch.Tensor, n_rows: int) -> torch.Tensor: """ Unvectorize a batch of tensors along last dimension Parameters ---------- data_vec : torch.Tensor of shape (..., n_rows*n_columns) Batch of vectorized matrices n_rows : int Number of rows of the matrices Returns ------- data : torch.Tensor of shape (..., n_rows, n_columns) Batch of matrices """ return data_vec.reshape(*data_vec.shape[:-1], n_rows, -1)
[docs] class VecBatch(Function): """ Vectorize a batch of matrices along last two dimensions. Matrices are assumed to be symmetric (for backward) """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the vectorization of a batch of matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- vec_data : torch.Tensor of shape (..., n_features ** 2) """ ctx.n_rows = data.shape[-2] return vec_batch(data)
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the vectorization of a batch of symmetric matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features ** 2) Gradient of the loss with respect to vectorized input batch of matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of matrices """ return symmetrize(unvec_batch(grad_output, ctx.n_rows))
[docs] def vech_batch(data: torch.Tensor) -> torch.Tensor: """ Vectorize the lower triangular part of a batch of square matrices Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of matrices Returns ------- data_vech : torch.Tensor of shape (..., n_features*(n_features+1)//2) Batch of vectorized matrices """ indices = torch.tril_indices(*data.shape[-2:]) return data[..., indices[0], indices[1]]
[docs] def unvech_batch(data_vech: torch.Tensor, n_features: int) -> torch.Tensor: """ Unvectorize a batch of tensors along last dimension, assuming that matrices are symmetric Parameters ---------- X_vech : torch.Tensor of shape (..., n_features*(n_features+1)//2) Batch of vectorized matrices n_features : int number of features Returns ------- X : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices """ indices_l = torch.tril_indices(n_features, n_features) indices_u = torch.triu_indices(n_features, n_features) data = torch.zeros( *data_vech.shape[:-1], n_features, n_features, dtype=data_vech.dtype, device=data_vech.device, ) # fill lower triangular data[..., indices_l[0], indices_l[1]] = data_vech # fill upper triangular data[..., indices_u[0], indices_u[1]] = data.transpose(-1, -2)[ ..., indices_u[0], indices_u[1] ] data = symmetrize(data) return data
[docs] class VechBatch(Function): """ Half vectorize a batch of symmetric matrices along last two dimensions """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the half vectorization of a batch of matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- vech_data : torch.Tensor of shape (..., n_features*(n_features+1) // 2) """ ctx.n_features = data.shape[-1] return vech_batch(data)
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the vectorization of a batch of matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features*(n_features+1) // 2) Gradient of the loss with respect to half vectorized input batch of symmetric matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of symmetric matrices """ return unvech_batch(grad_output, ctx.n_features)
# ------------------------- # Operations on eigenvalues # -------------------------
[docs] def eigh_operation( eigvals: torch.Tensor, eigvecs: torch.Tensor, operation: Callable ) -> torch.Tensor: r""" Applies a function on the eigenvalues of a batch of symmetric matrices. .. math:: f(A) = V \operatorname{diag}\big(f(\lambda_1), \dots, f(\lambda_n)\big) V^\top given the eigendecomposition :math:`A = V \operatorname{diag}(\lambda) V^\top`. This is the core primitive behind every matrix function in this module (:func:`sqrtm_SPD`, :func:`inv_sqrtm_SPD`, :func:`powm_SPD`, :func:`logm_SPD`, :func:`expm_symmetric`, …): each just picks a different scalar ``operation`` applied eigenvalue-wise. Parameters ---------- eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of the corresponding batch of symmetric matrices eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of the corresponding batch of symmetric matrices operation : Callable Function to apply on eigenvalues Returns ------- result : torch.Tensor of shape (..., n_features, n_features) Resulting symmetric matrices with operation applied to eigenvalues """ eigvals, eigvecs = ( torch.real(eigvals), torch.real(eigvecs), ) # to correct eventual numerical errors... _eigvals = operation(eigvals) result = (eigvecs * _eigvals.unsqueeze(-2)) @ eigvecs.transpose(-1, -2) return result
def _aux_eigh_operation_grad( eigvals: torch.Tensor, operation: Callable, operation_deriv: Callable, ) -> torch.Tensor: """ Constructs matrix to be multiplied (Hadamard product) with transformed output gradient to get the input gradient for a function applied on the eigenvalues of a symmetric matrix Parameters ---------- eigvals : torch.Tensor of shape (..., n_features) eigenvalues of the SPD matrices operation : Callable Function to apply on eigenvalues operation_deriv : Callable Derivative of the function to apply on eigenvalues Returns ------- aux_mat : torch.Tensor of shape (..., n_features, n_features) Matrix to be multiplied (Hadamard product) with transformed output gradient """ eigvals_transformed = operation(eigvals) eigvals_deriv = operation_deriv(eigvals) denominator = eigvals.unsqueeze(-1) - eigvals.unsqueeze(-2) numerator = eigvals_transformed.unsqueeze(-1) - eigvals_transformed.unsqueeze(-2) null_denominator = ( torch.abs(denominator) < 1e-6 ) # arbitrary value, probably to be changed to be proper... Ammar ? numerator = torch.where(null_denominator, eigvals_deriv.unsqueeze(-1), numerator) denominator = torch.where(null_denominator, torch.ones_like(numerator), denominator) return numerator / denominator
[docs] def eigh_operation_grad( grad_output: torch.Tensor, eigvals: torch.Tensor, eigvecs: torch.Tensor, operation: Callable, operation_deriv: Callable, ) -> torch.Tensor: """ Computes the backpropagation of the gradient for a function applied on the eigenvalues of a batch of symmetric matrices Parameters ---------- grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the output of the operation on eigenvalues eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of the corresponding batch of symmetric matrices eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of the corresponding batch of symmetric matrices operation : Callable Function to apply on eigenvalues operation_deriv : Callable Derivative of the function to apply on eigenvalues Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of symmetric matrices """ aux_mat = _aux_eigh_operation_grad(eigvals, operation, operation_deriv) middle_term = aux_mat * (eigvecs.transpose(-1, -2) @ grad_output @ eigvecs) return eigvecs @ middle_term @ eigvecs.transpose(-1, -2)
# ------------------ # Sylvester equation # ------------------
[docs] def solve_sylvester_SPD( eigvals: torch.Tensor, eigvecs: torch.Tensor, mat: torch.Tensor ) -> torch.Tensor: r""" Solve Sylvester equations in the context of SPD matrices relying on eigenvalue decomposition. Given :math:`A = V \operatorname{diag}(\lambda) V^\top` (via ``eigvals``, ``eigvecs``), solves :math:`AX + XA = \text{mat}` for :math:`X` in closed form: .. math:: X = V\left[\frac{1}{\lambda_i + \lambda_j} \,(V^\top \,\text{mat}\, V)_{ij}\right]_{ij} V^\top Parameters ---------- eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of a batch of SPD matrices eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of a batch of SPD matrices mat : torch.Tensor of shape (..., n_features, n_features) Batch of matrices on the right side of Sylvester equations. If matrices are symmetric then the solutions will be symmetric. If they are skew-symmetric, then the results will be skew-symmetric. Returns ------- result : torch.Tensor of shape (..., n_features, n_features) Symmetric matrices solutions to Sylvester equations """ # should those corrections be here though ? eigvals, eigvecs = ( torch.real(eigvals), torch.real(eigvecs), ) # to correct eventual numerical errors... K = 1 / (eigvals.unsqueeze(-1) + eigvals.unsqueeze(-2)) middle_term = K * (eigvecs.transpose(-1, -2) @ mat @ eigvecs) return eigvecs @ middle_term @ eigvecs.transpose(-1, -2)
# ---------------------- # SPD matrix square root # ----------------------
[docs] def sqrtm_SPD(data: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Matrix square root of a batch of SPD matrices. .. math:: P^{1/2} = V \operatorname{diag}(\sqrt{\lambda}) V^\top via :func:`eigh_operation` with ``operation=torch.sqrt``. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- sqrtm_data : torch.Tensor of shape (..., n_features, n_features) Matrix square roots of the input batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, torch.sqrt), eigvals, eigvecs
[docs] class SqrtmSPD(Function): """ Matrix square root of a batch of SPD matrices (relies on eigenvalue decomposition) """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the matrix square root of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- sqrtm_data : torch.Tensor of shape (..., n_features, n_features) Matrix square roots of the input batch of SPD matrices """ sqrtm_data, eigvals, eigvecs = sqrtm_SPD(data) ctx.save_for_backward(eigvals, eigvecs) return sqrtm_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the matrix square root of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix square roots of the input batch of SPD matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad( grad_output, eigvals, eigvecs, torch.sqrt, sqrt_derivative )
# ------------------------------ # SPD matrix inverse square root # ------------------------------
[docs] def inv_sqrtm_SPD( data: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Inverse matrix square root of a batch of SPD matrices. .. math:: P^{-1/2} = V \operatorname{diag}(\lambda^{-1/2}) V^\top via :func:`eigh_operation` with ``operation=inv_sqrt``. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- inv_sqrtm_data : torch.Tensor of shape (..., n_features, n_features) Inverse matrix square roots of the input batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, inv_sqrt), eigvals, eigvecs
[docs] class InvSqrtmSPD(Function): """ Matrix inverse square root of a batch of SPD matrices (relies on eigenvalue decomposition) """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the matrix inverse square root of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- inv_sqrtm_data : torch.Tensor of shape (..., n_features, n_features) Matrix square roots of the input batch of SPD matrices """ inv_sqrtm_data, eigvals, eigvecs = inv_sqrtm_SPD(data) ctx.save_for_backward(eigvals, eigvecs) return inv_sqrtm_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the matrix inverse square root of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix square roots of the input batch of SPD matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad( grad_output, eigvals, eigvecs, inv_sqrt, inv_sqrt_derivative )
# ---------------- # SPD matrix power # ----------------
[docs] def powm_SPD( data: torch.Tensor, exponent: float | torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Matrix power of a batch of SPD matrices. .. math:: P^{p} = V \operatorname{diag}(\lambda^{p}) V^\top via :func:`eigh_operation` with ``operation=lambda x: x**exponent``. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices exponent : torch.float Power exponent Returns ------- powm_data : torch.Tensor of shape (..., n_features, n_features) Matrix powers of the input batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) pow_fun = lambda x: torch.pow(x, exponent) return eigh_operation(eigvals, eigvecs, pow_fun), eigvals, eigvecs
[docs] class PowmSPD(Function): """ Matrix power of a batch of SPD matrices (relies on eigenvalue decomposition) """
[docs] @staticmethod def forward( ctx, data: torch.Tensor, exponent: float | torch.Tensor ) -> torch.Tensor: """ Forward pass of the matrix power of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices exponent : torch.float Power exponent Returns ------- powm_data : torch.Tensor of shape (..., n_features, n_features) Matrix powers of the input batch of SPD matrices """ powm_data, eigvals, eigvecs = powm_SPD(data, exponent) ctx.save_for_backward(eigvals, eigvecs, exponent) return powm_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """ Backward pass of the matrix power of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix powers of the input batch of SPD matrices Returns ------- grad_input_data : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices grad_input_exponent : torch.float Gradient of the loss with respect to the power exponent """ eigvals, eigvecs, exponent = ctx.saved_tensors pow_fun = lambda x: torch.pow(x, exponent) pow_deriv = lambda x: exponent * torch.pow(x, exponent - 1) exponent_deriv_fun = lambda x: torch.pow(x, exponent) * torch.log(x) return ( eigh_operation_grad(grad_output, eigvals, eigvecs, pow_fun, pow_deriv), (grad_output @ eigh_operation(eigvals, eigvecs, exponent_deriv_fun)) .diagonal(offset=0, dim1=-1, dim2=-2) .sum() .reshape(exponent.shape), )
# -------------------- # SPD matrix logarithm # --------------------
[docs] def logm_SPD(data: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Matrix logarithm of a batch of SPD matrices. .. math:: \log(P) = V \operatorname{diag}(\log\lambda) V^\top via :func:`eigh_operation` with ``operation=torch.log``. Maps the SPD manifold to the vector space of symmetric matrices — the basis of the Log-Euclidean geometry (see :mod:`~yetanotherspdnet.functions.spd_geometries.log_euclidean`). Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- logm_data : torch.Tensor of shape (..., n_features, n_features) Matrix logarithms of the input batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, torch.log), eigvals, eigvecs
[docs] class LogmSPD(Function): """ Matrix logarithm of a batch of SPD matrices (relies on eigenvalue decomposition) """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the matrix logarithm of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- logm_data : torch.Tensor of shape (..., n_features, n_features) Matrix logarithms of the input batch of SPD matrices """ logm_data, eigvals, eigvecs = logm_SPD(data) ctx.save_for_backward(eigvals, eigvecs) return logm_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the matrix logarithm of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix logarithms of the input batch of SPD matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad(grad_output, eigvals, eigvecs, torch.log, inv)
# ---------------------------- # Symmetric matrix exponential # ----------------------------
[docs] def expm_symmetric( data: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Matrix exponential of a batch of symmetric matrices. .. math:: \exp(S) = V \operatorname{diag}(\exp\lambda) V^\top via :func:`eigh_operation` with ``operation=torch.exp``. The result is always SPD (eigenvalues :math:`\exp\lambda > 0`); inverse of :func:`logm_SPD`. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices Returns ------- expm_data : torch.Tensor of shape (..., n_features, n_features) Matrix exponentials of the input batch of symmetric matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, torch.exp), eigvals, eigvecs
[docs] class ExpmSymmetric(Function): """ Matrix exponential of a batch of symmetric matrices (relies on eigenvalue decomposition) """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the matrix exponential of a batch of symmetric matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices Returns ------- expm_data : torch.Tensor of shape (..., n_features, n_features) Matrix exponentials of the input batch of symmetric matrices """ expm_data, eigvals, eigvecs = expm_symmetric(data) ctx.save_for_backward(eigvals, eigvecs) return expm_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the matrix exponential of a batch of symmetric matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix exponentials of the input batch of symmetric matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of symmetric matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad(grad_output, eigvals, eigvecs, torch.exp, torch.exp)
# ------------------------- # Scaled SoftPlus Symmetric # -------------------------
[docs] def scaled_softplus_symmetric( data: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Scaled matrix SoftPlus of a batch of symmetric matrices. .. math:: f(S) = V \operatorname{diag}\big(\log_2(1 + 2^{\lambda})\big) V^\top via :func:`eigh_operation` with ``operation=`` :func:`~yetanotherspdnet.functions.scalar_functions.scaled_softplus`. Maps any symmetric matrix to an SPD matrix (eigenvalues strictly positive) — used to parametrize BiMap weights or BatchNorm scale so they stay on the SPD manifold under unconstrained optimization. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices Returns ------- softplus_data : torch.Tensor of shape (..., n_features, n_features) Matrix SoftPlus of the input batch of symmetric matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, scaled_softplus), eigvals, eigvecs
[docs] class ScaledSoftPlusSymmetric(Function): """ Scaled matrix SoftPlus of a batch of symmetric matrices. It is scaled so that: f(0) = 1, f(x) -> 0 as x -> -inf and f'(x) -> 1 as x -> +inf """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the scaled matrix SoftPlus of a batch of symmetric matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices Returns ------- softplus_data : torch.Tensor of shape (..., n_features, n_features) Matrix SoftPlus of the input batch of symmetric matrices """ softplus_data, eigvals, eigvecs = scaled_softplus_symmetric(data) ctx.save_for_backward(eigvals, eigvecs) return softplus_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the scaled matrix SoftPlus of a batch of symmetric matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix SoftPlus of the input batch of symmetric matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of symmetric matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad( grad_output, eigvals, eigvecs, scaled_softplus, scaled_softplus_derivative )
# --------------------------- # Inverse Scaled SoftPlus SPD # ---------------------------
[docs] def inv_scaled_softplus_SPD( data: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Inverse scaled SoftPlus of a batch of SPD matrices. .. math:: f^{-1}(P) = V \operatorname{diag}\big(\log_2(2^{\lambda} - 1)\big) V^\top via :func:`eigh_operation`. Inverse of :func:`scaled_softplus_symmetric`. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- inv_softplus_data : torch.Tensor of shape (..., n_features, n_features) Matrix logarithms of the input batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) return eigh_operation(eigvals, eigvecs, inv_scaled_softplus), eigvals, eigvecs
[docs] class InvScaledSoftPlusSPD(Function): """ Matrix inverse scaled SoftPlus of a batch of SPD matrices """
[docs] @staticmethod def forward(ctx, data: torch.Tensor) -> torch.Tensor: """ Forward pass of the matrix inverse scaled SoftPlus of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices Returns ------- inv_softplus_data : torch.Tensor of shape (..., n_features, n_features) Matrix logarithms of the input batch of SPD matrices """ inv_softplus_data, eigvals, eigvecs = inv_scaled_softplus_SPD(data) ctx.save_for_backward(eigvals, eigvecs) return inv_softplus_data
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor: """ Backward pass of the matrix inverse scaled SoftPlus of a batch of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to matrix inverse SoftPlus of the input batch of SPD matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices """ eigvals, eigvecs = ctx.saved_tensors return eigh_operation_grad( grad_output, eigvals, eigvecs, inv_scaled_softplus, inv_scaled_softplus_derivative, )
# ----------------------------------------------------- # ReLu activation function on eigenvalues of SPD matrix # -----------------------------------------------------
[docs] def eigh_relu( data: torch.Tensor, eps: float ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" ReLu activation function on the eigenvalues of SPD matrices. .. math:: \operatorname{ReEig}_\epsilon(P) = V \operatorname{diag}\big( \max(\lambda, \epsilon)\big) V^\top via :func:`eigh_operation`. This is the ReEig layer's core operation (:class:`~yetanotherspdnet.nn.base.ReEig`): clamps small/negative eigenvalues to :math:`\epsilon > 0` to keep the result SPD, the manifold analogue of ReLU rectification. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices eps : float Value for the rectification of the eigenvalues Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices with rectified eigenvalues eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) operation = lambda x: torch.clamp(x, min=eps) return eigh_operation(eigvals, eigvecs, operation), eigvals, eigvecs
[docs] class EighReLu(Function): """ ReLu activation function on the eigenvalues of SPD matrices """
[docs] @staticmethod def forward(ctx, data: torch.Tensor, eps: float) -> torch.Tensor: """ Forward pass of the ReLu activation function on the eigenvalues of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to save tensors for the backward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices eps : float Value for the rectification of the eigenvalues Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices with rectified eigenvalues """ data_transformed, eigvals, eigvecs = eigh_relu(data, eps) ctx.save_for_backward(eigvals, eigvecs) ctx.eps = eps return data_transformed
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple[torch.Tensor, None]: """ Backward pass of the ReLu activation function on the eigenvalues of SPD matrices Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to save tensors for the backward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the output batch of SPD matrices Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices """ eps = ctx.eps eigvals, eigvecs = ctx.saved_tensors operation = lambda x: torch.clamp(x, min=eps) operation_deriv = lambda x: (x > eps).type(x.dtype) return ( eigh_operation_grad( grad_output, eigvals, eigvecs, operation, operation_deriv ), None, )
def _clamped_bias_operation(bias: torch.Tensor, eps: float) -> Callable: """Eigenvalue map of ReEigBias: shift by ``bias`` then clamp to [eps, 1/eps].""" return lambda x: torch.clamp(x + bias, min=eps, max=1 / eps) def _clamped_bias_operation_deriv(bias: torch.Tensor, eps: float) -> Callable: """Derivative of :func:`_clamped_bias_operation` with respect to the eigenvalues.""" return lambda x: ((x + bias > eps) & (x + bias < 1 / eps)).type(x.dtype)
[docs] def eigh_relu_bias( data: torch.Tensor, bias: torch.Tensor, eps: float ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: r""" Eigenvalue rectification with a learnable shift of the eigenvalues. .. math:: \operatorname{ReEigBias}_{\epsilon, b}(P) = V \operatorname{diag}\big( \operatorname{clamp}(\lambda_i + b_i,\ \epsilon,\ 1/\epsilon) \big) V^\top with :math:`\lambda_1 \le \dots \le \lambda_n` the eigenvalues in ascending order (as returned by ``torch.linalg.eigh``) and :math:`b` a bias vector indexed by eigenvalue rank. The upper clamp bounds the condition number of the output by :math:`\epsilon^{-2}`. Unlike :func:`eigh_relu`, this is not a spectral function when two eigenvalues coincide while their biases differ: the output then depends on the arbitrary eigenbasis of the repeated eigenvalue, and gradients with respect to ``data`` are only defined for distinct eigenvalues. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices bias : torch.Tensor of shape (n_features,) Shift added to the (ascending) eigenvalues eps : float Lower clamping value; the upper one is ``1 / eps`` Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices eigvals : torch.Tensor of shape (..., n_features) Eigenvalues of matrices in data eigvecs : torch.Tensor of shape (..., n_features, n_features) Eigenvectors of matrices in data """ eigvals, eigvecs = torch.linalg.eigh(data) operation = _clamped_bias_operation(bias, eps) return eigh_operation(eigvals, eigvecs, operation), eigvals, eigvecs
[docs] class EighReLuBias(Function): """ Eigenvalue rectification with a learnable shift, with a hand-written backward """
[docs] @staticmethod def forward( ctx, data: torch.Tensor, bias: torch.Tensor, eps: float ) -> torch.Tensor: """ Forward pass of :func:`eigh_relu_bias` Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to save tensors for the backward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of symmetric matrices bias : torch.Tensor of shape (n_features,) Shift added to the (ascending) eigenvalues eps : float Lower clamping value; the upper one is ``1 / eps`` Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices """ data_transformed, eigvals, eigvecs = eigh_relu_bias(data, bias, eps) ctx.save_for_backward(eigvals, eigvecs, bias) ctx.eps = eps return data_transformed
[docs] @staticmethod def backward( ctx, grad_output: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor, None]: """ Backward pass of :func:`eigh_relu_bias` The gradient with respect to ``data`` is the Daleckii-Krein formula of :func:`eigh_operation_grad`. Only the eigenvalues depend on the bias, so its gradient is the diagonal of :math:`V^\top G V` masked by the derivative of the clamp, summed over the batch dimensions. Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to save tensors for the backward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the output batch Returns ------- grad_input : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch grad_bias : torch.Tensor of shape (n_features,) Gradient of the loss with respect to the bias """ eps = ctx.eps eigvals, eigvecs, bias = ctx.saved_tensors operation = _clamped_bias_operation(bias, eps) operation_deriv = _clamped_bias_operation_deriv(bias, eps) grad_input = eigh_operation_grad( grad_output, eigvals, eigvecs, operation, operation_deriv ) rotated = eigvecs.transpose(-1, -2) @ symmetrize(grad_output) @ eigvecs grad_eigvals = torch.diagonal(rotated, dim1=-2, dim2=-1) * operation_deriv( eigvals ) grad_bias = grad_eigvals.reshape(-1, grad_eigvals.shape[-1]).sum(dim=0) return grad_input, grad_bias, None
# ----------------------------------- # Various congruences of SPD matrices # -----------------------------------
[docs] def congruence_SPD(data: torch.Tensor, matrix: torch.Tensor) -> torch.Tensor: r""" Congruence of a batch of SPD matrices with an SPD matrix. .. math:: P' = A\, P\, A (here :math:`A` is itself SPD, hence symmetric, so :math:`A^\top = A` and there is no separate transpose). Congruence by an SPD matrix preserves the SPD manifold. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices matrix : torch.Tensor of shape (n_features, n_features) SPD matrix Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Transformed batch of SPD matrices """ return matrix @ data @ matrix
[docs] class CongruenceSPD(Function): """ Congruence of a batch of SPD matrices with an SPD matrix """
[docs] @staticmethod def forward(ctx, data: torch.Tensor, matrix: torch.Tensor) -> torch.Tensor: """ Forward pass of the congruence of a batch of SPD matrices with an SPD matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices matrix : torch.Tensor of shape (n_features, n_features) SPD matrix Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Transformed batch of SPD matrices """ ctx.save_for_backward(data, matrix) return congruence_SPD(data, matrix)
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """ Backward pass of the congruence of a batch of SPD matrices with an SPD matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the batch of transformed SPD matrices Returns ------- grad_input_data : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices grad_input_bias : torch.Tensor of shape (n_features, n_features) Gradient of the loss with respect to the SPD matrix used for congruence """ data, matrix = ctx.saved_tensors # return matrix @ grad_output @ matrix, torch.sum( # 2 * symmetrize(grad_output @ matrix @ data), dim=tuple(range(data.ndim - 2)) # ) return ( matrix @ grad_output @ matrix, 2 * symmetrize(torch.einsum("...ik,kl,...lj->ij", grad_output, matrix, data)), )
[docs] def whitening(data: torch.Tensor, matrix: torch.Tensor) -> torch.Tensor: r""" Whitening of a batch of SPD matrices with an SPD matrix. .. math:: P' = A^{-1/2}\, P\, A^{-1/2} i.e. :func:`congruence_SPD` with :math:`A^{-1/2}` (see :func:`inv_sqrtm_SPD`) — transforms data so that :math:`A` itself maps to the identity, the SPD analogue of standardizing by the covariance. Parameters ---------- data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices matrix : torch.Tensor of shape (n_features, n_features) SPD matrix Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Transformed batch of SPD matrices """ inv_sqrtm_matrix, _, _ = inv_sqrtm_SPD(matrix) return congruence_SPD(data, inv_sqrtm_matrix)
[docs] class Whitening(Function): """ Whitening of a batch of SPD matrices with an SPD matrix """
[docs] @staticmethod def forward(ctx, data: torch.Tensor, matrix: torch.Tensor) -> torch.Tensor: """ Forward pass of the whitening of a batch of SPD matrices with an SPD matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_features, n_features) Batch of SPD matrices matrix : torch.Tensor of shape (n_features, n_features) SPD matrix Returns ------- data_transformed : torch.Tensor of shape (..., n_features, n_features) Transformed batch of SPD matrices """ inv_sqrtm_matrix, eigvals_matrix, eigvecs_matrix = inv_sqrtm_SPD(matrix) ctx.save_for_backward(data, eigvals_matrix, eigvecs_matrix, inv_sqrtm_matrix) return congruence_SPD(data, inv_sqrtm_matrix)
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """ Backward pass of the whitening of a batch of SPD matrices with an SPD matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the batch of whitened SPD matrices Returns ------- grad_input_data : torch.Tensor of shape (..., n_features, n_features) Gradient of the loss with respect to the input batch of SPD matrices grad_input_matrix : torch.Tensor of shape (n_features, n_features) Gradient of the loss with respect to the SPD matrix used for whitening """ data, eigvals_matrix, eigvecs_matrix, inv_sqrtm_matrix = ctx.saved_tensors grad_input_data = inv_sqrtm_matrix @ grad_output @ inv_sqrtm_matrix # syl_right = -torch.sum( # 2 * symmetrize(grad_input_data @ data @ inv_sqrtm_matrix), # dim=tuple(range(data.ndim - 2)), # ) syl_right = -2 * symmetrize( torch.einsum("...ik,...kl,lj->ij", grad_input_data, data, inv_sqrtm_matrix) ) grad_input_matrix = solve_sylvester_SPD( torch.sqrt(eigvals_matrix), eigvecs_matrix, syl_right ) return grad_input_data, grad_input_matrix
[docs] def congruence_rectangular(data: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: r""" Forward pass of the congruence of a batch of SPD matrices with a (full-rank) rectangular matrix. .. math:: P' = W^\top P\, W, \qquad W \in \mathbb{R}^{n_{in} \times n_{out}} with :math:`n_{in} \geq n_{out}`. This is the BiMap layer's core operation (:class:`~yetanotherspdnet.nn.base.BiMap`): a dimension reduction that keeps the result SPD as long as :math:`W` has full column rank. Parameters ---------- data : torch.Tensor of shape (..., n_in, n_in) Batch of SPD matrices weight : torch.Tensor of shape (n_in, n_out) Rectangular matrix (e.g., weights), n_in > n_out is expected Returns ------- data_transformed : torch.Tensor of shape (..., n_out, n_out) Transformed batch of SPD matrices """ assert weight.shape[-2] >= weight.shape[-1], ( "weight must reduce the dimension of data, i.e., n_in >= n_out" ) return weight.transpose(-1, -2) @ data @ weight
[docs] class CongruenceRectangular(Function): """ Congruence of a batch of SPD matrices with a (full-rank) rectangular matrix """
[docs] @staticmethod def forward(ctx, data: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: """ Forward pass of the congruence of a batch of SPD matrices with a (full-rank) rectangular matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass data : torch.Tensor of shape (..., n_in, n_in) Batch of SPD matrices weight : torch.Tensor of shape (n_out, n_in) Rectangular matrix (e.g., weights), n_in > n_out is expected Returns ------- data_transformed : torch.Tensor of shape (..., n_out, n_out) Transformed batch of SPD matrices """ ctx.save_for_backward(data, weight) return congruence_rectangular(data, weight)
[docs] @staticmethod def backward(ctx, grad_output: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """ Backward pass of the congruence of a batch of SPD matrices with a (full-rank) rectangular matrix Parameters ---------- ctx : torch.autograd.function._ContextMethodMixin Context object to retrieve tensors saved during the forward pass grad_output : torch.Tensor of shape (..., n_out, n_out) Gradient of the loss with respect to the batch of transformed SPD matrices Returns ------- grad_input_data : torch.Tensor of shape (..., n_in, n_in) Gradient of the loss with respect to the input batch of SPD matrices grad_input_W : torch.Tensor of shape (n_in, n_out) Gradient of the loss with respect to the (full-rank) rectangular matrix W """ data, weight = ctx.saved_tensors grad_input_data = weight @ grad_output @ weight.transpose(-1, -2) # grad_input_W = 2 * torch.sum( # grad_output @ W @ data, dim=tuple(range(data.ndim - 2)) # ) grad_input_W = 2 * torch.einsum("...ik,kl,...lj->ij", data, weight, grad_output) return grad_input_data, grad_input_W