"""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