Geometries

yetanotherspdnet.functions.spd_geometries — one module per Riemannian geometry of the SPD manifold, each with its geodesic, mean, scalar dispersion and, where the batch normalization needs them, exponential and logarithmic maps. The Geometries on the SPD manifold guide compares them and gives their formulas.

Module

Geometry

Mean used by the batch normalization

affine_invariant

affine-invariant (Fisher–Rao)

Karcher mean (iterative)

log_euclidean

log-Euclidean

\(\exp(\overline{\log X})\)

kullback_leibler

left / right Kullback–Leibler

arithmetic / harmonic mean

kullback_leibler_symmetrized

symmetrized Kullback–Leibler

GAH: AI midpoint of harmonic and arithmetic (adaptive: learned point \(t\))

bures_wasserstein

Bures–Wasserstein

fixed-point barycenter

Affine-invariant

autograd path

manual backward

affine_invariant_geodesic()

AffineInvariantGeodesic

Affine-invariant geodesic between two SPD matrices.

affine_invariant_mean_2points()

AffineInvariantMean2Points

Affine-invariant (geometric) mean of two SPD matrices.

AffineInvariantMean()

–

Affine-invariant (geometric) mean computed with fixed-point algorithm

–

AffineInvariantMeanIteration

One iteration of the fixed-point algorithm computing the affine-invariant (geometric) mean

affine_invariant_std_scalar()

AffineInvariantStdScalar

Scalar standard deviation with respect to the affine-invariant distance.

affine_invariant_exp()

AffineInvariantExp

Affine-invariant exponential map on the SPD manifold.

affine_invariant_log()

–

Affine-invariant logarithmic map on the SPD manifold.

affine_invariant_projx()

–

Project matrices onto the SPD manifold by symmetrizing and clamping eigenvalues to [1e-8, 1e8] (as in the reference RResNet implementation, which also bounds the conditioning).

Affine-invariant Riemannian geometry: geodesic, exp/log maps, mean, and standard deviation.

affine_invariant_geodesic(point1, point2, t)[source]

Affine-invariant geodesic between two SPD matrices.

\[\gamma(t) = P_1^{1/2} \big(P_1^{-1/2} P_2 P_1^{-1/2}\big)^{t} P_1^{1/2}\]

with \(t \in [0, 1]\) (\(\gamma(0) = P_1\), \(\gamma(1) = P_2\)). This is the geodesic for the affine-invariant Riemannian metric on the SPD manifold, invariant under congruence transformations \(P \mapsto A P A^\top\) for any invertible \(A\).

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

class AffineInvariantGeodesic(*args, **kwargs)[source]

Affine-invariant geodesic between two batches of SPD matrices.

Computes: point1^{1/2} (point1^{-1/2} point2 point1^{-1/2})^t point1^{1/2}

Supports gradients with respect to point1, point2, and optionally t (when t is a tensor with requires_grad=True).

static forward(ctx, point1, point2, t)[source]

Forward pass of the affine-invariant geodesic between two batches of SPD matrices

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • point1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • point2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • t (float | torch.Tensor) – Parameter on the geodesic path, should be in [0, 1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

static backward(ctx, grad_output)[source]

Backward pass of the affine-invariant geodesic.

Computes gradients with respect to point1, point2, and optionally t.

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to the output

Returns:

  • grad_input1 (torch.Tensor of shape (..., nfeatures, nfeatures) or None) – Gradient of the loss with respect to point1

  • grad_input2 (torch.Tensor of shape (..., nfeatures, nfeatures) or None) – Gradient of the loss with respect to point2

  • grad_t (torch.Tensor of shape () or None) – Gradient of the loss with respect to t (only when t requires grad)

Return type:

tuple[Tensor | None, Tensor | None, Tensor | None]

affine_invariant_mean_2points(point1, point2)[source]

Affine-invariant (geometric) mean of two SPD matrices.

\[G(P_1, P_2) = P_1^{1/2} \big(P_1^{-1/2} P_2 P_1^{-1/2}\big)^{1/2} P_1^{1/2}\]

the midpoint (\(t=1/2\)) of the affine-invariant geodesic between \(P_1\) and \(P_2\) (see affine_invariant_geodesic()).

Parameters:
  • point1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • point2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

Returns:

mean (torch.Tensor of shape (..., nfeatures, nfeatures)) – Geometric means of point1 and point2

Return type:

Tensor

class AffineInvariantMean2Points(*args, **kwargs)[source]

Affine-invariant (geometric) mean of two SPD matrices

static forward(ctx, point1, point2)[source]

Forward pass of the geometric mean of two SPD matrices

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • point1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • point2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

Returns:

mean (torch.Tensor of shape (..., nfeatures, nfeatures)) – Geometric means of point1 and point2

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the geometric mean of two SPD matrices

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to the geometric mean of two SPD matrices

Returns:

  • grad_input1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to point1

  • grad_input2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to point2

Return type:

tuple[Tensor, Tensor]

affine_invariant_mean(data, n_iterations=5)[source]

Affine-invariant (geometric/Fréchet) mean computed with a fixed-point (Karcher flow) algorithm.

Starting from \(M_0 = I\), each iteration \(k\) moves along the average tangent direction at \(M_k\) and retracts back onto the manifold with the affine-invariant exponential map (affine_invariant_exp()):

\[M_{k+1} = M_k^{1/2} \exp\!\left(\eta_k \cdot \frac{1}{N}\sum_{i=1}^{N} \log\big(M_k^{-1/2} P_i M_k^{-1/2}\big) \right) M_k^{1/2}\]

with step size \(\eta_k = 0.95^k\). This converges to the unique minimizer of \(\sum_i d_{AI}(M, P_i)^2\) for the affine-invariant distance \(d_{AI}\).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • n_iterations (int) – Number of iterations to perform to estimate the geometric mean, by default 5

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – SPD matrix

Return type:

Tensor

class AffineInvariantMeanIteration(*args, **kwargs)[source]

One iteration of the fixed-point algorithm computing the affine-invariant (geometric) mean

static forward(ctx, mean_iterate, data, stepsize)[source]

Forward pass of one iteration of the fixed-point algorithm for the affine-invariant (geometric) mean

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • mean_iterate (torch.Tensor of shape (n_features, n_features)) – Current iterate of the affine-invariant mean

  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • stepsize (float) – step-size to stabilize the fixed-point algorithm

Returns:

mean_iterate_new (torch.Tensor of shape (n_features, n_features)) – New iterate of the affine-invariant mean

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of one iteration of the fixed-point algorithm for the affine-invariant (geometric) mean

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape (nfeatures, nfeatures)) – Gradient of the loss with respect to the new iterate of the affine-invariant mean

Returns:

  • grad_input_mean (torch.Tensor of shape (nfeatures, nfeatures)) – Gradient of the loss with respect to the current iterate of the affine-invariant mean

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – Gradient of the loss with respect to the data at the current iterate

Return type:

tuple[Tensor, Tensor, None]

AffineInvariantMean(data, n_iterations=5)[source]

Affine-invariant (geometric) mean computed with fixed-point algorithm

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • n_iterations (int) – Number of iterations to perform to estimate the geometric mean. Default is 10

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – SPD matrix

Return type:

Tensor

affine_invariant_std_scalar(data, reference_point)[source]

Scalar standard deviation with respect to the affine-invariant distance.

\[\sigma = \sqrt{\frac{1}{N}\sum_{i=1}^{N} \big\lVert \log\big(G^{-1/2} P_i G^{-1/2}\big) \big\rVert_F^2}\]

where \(G\) is the reference point (typically the affine-invariant mean) — equivalently, the norm of \(\mathrm{Log}_G(P_i)\) under the affine-invariant metric (affine_invariant_log()).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

class AffineInvariantStdScalar(*args, **kwargs)[source]

Scalar standard deviation with respect to the affine-invariant distance

static forward(ctx, data, reference_point)[source]

Forward pass of the scalar standard deviation with respect to the affine-invariant distance

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

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the scalar standard deviation with respect to the affine-invariant distance

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – gradient of the loss with respect to the input reference point

Return type:

tuple[Tensor, Tensor]

affine_invariant_exp(base, tangent)[source]

Affine-invariant exponential map on the SPD manifold.

\[\mathrm{Exp}_X(V) = X^{1/2} \exp\big(X^{-1/2} V X^{-1/2}\big) X^{1/2}\]

Maps a tangent vector \(V\) at base point \(X\) (a symmetric matrix) to a point on the SPD manifold, by following the geodesic from \(X\) in direction \(V\) for unit time. Inverse of affine_invariant_log().

Parameters:
  • base (torch.Tensor of shape (..., n, n)) – Base point(s) on the SPD manifold

  • tangent (torch.Tensor of shape (..., n, n)) – Tangent vector(s) at base (symmetric matrices)

Returns:

result (torch.Tensor of shape (..., n, n)) – Point(s) on the SPD manifold

Return type:

Tensor

class AffineInvariantExp(*args, **kwargs)[source]

Affine-invariant exponential map with manual backward.

Exp_X(V) = X^{1/2} expm(X^{-1/2} V X^{-1/2}) X^{1/2}

static forward(ctx, base, tangent)[source]

Forward pass of the affine-invariant exponential map.

Parameters:
  • ctx (context) – Context for saving tensors for backward

  • base (torch.Tensor of shape (..., n, n)) – Base point(s) on the SPD manifold

  • tangent (torch.Tensor of shape (..., n, n)) – Tangent vector(s) at base (symmetric matrices)

Returns:

result (torch.Tensor of shape (..., n, n)) – Point(s) on the SPD manifold

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the affine-invariant exponential map.

Parameters:
  • ctx (context) – Context with saved tensors

  • grad_output (torch.Tensor of shape (..., n, n)) – Gradient w.r.t. the output

Returns:

Return type:

tuple[Tensor | None, Tensor | None]

affine_invariant_log(base, point)[source]

Affine-invariant logarithmic map on the SPD manifold.

\[\mathrm{Log}_X(Y) = X^{1/2} \log\big(X^{-1/2} Y X^{-1/2}\big) X^{1/2}\]

Maps a point \(Y\) on the manifold to a tangent vector at base \(X\). Inverse of affine_invariant_exp().

Parameters:
  • base (torch.Tensor of shape (..., n, n)) – Base point(s) on the SPD manifold

  • point (torch.Tensor of shape (..., n, n)) – Point(s) on the SPD manifold

Returns:

tangent (torch.Tensor of shape (..., n, n)) – Tangent vector(s) at base (symmetric matrices)

Return type:

Tensor

affine_invariant_projx(data)[source]

Project matrices onto the SPD manifold by symmetrizing and clamping eigenvalues to [1e-8, 1e8] (as in the reference RResNet implementation, which also bounds the conditioning).

Parameters:

data (torch.Tensor of shape (..., n, n)) – Batch of matrices

Returns:

projected (torch.Tensor of shape (..., n, n)) – Batch of SPD matrices

Return type:

Tensor

Log-Euclidean

autograd path

manual backward

LogEuclideanGeodesic()

–

Log-Euclidean geodesic between two batches of SPD matrices

LogEuclideanMean()

–

Log-Euclidean mean

log_euclidean_std_scalar()

LogEuclideanStdScalar

Scalar standard deviation with respect to the Log-Euclidean distance.

Log-Euclidean Riemannian geometry: geodesic, mean, and standard deviation.

log_euclidean_geodesic(point1, point2, t)[source]

Log-Euclidean geodesic between two SPD matrices.

\[\gamma(t) = \exp\big((1-t)\,\log(P_1) + t\,\log(P_2)\big)\]

where \(\log\) and \(\exp\) are the matrix logarithm and exponential (logm_SPD(), expm_symmetric()). This is simply the Euclidean geodesic in the tangent space obtained by taking the matrix log, mapped back to the SPD manifold.

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

LogEuclideanGeodesic(point1, point2, t)[source]

Log-Euclidean geodesic between two batches of SPD matrices

Parameters:
  • point1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • point2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

log_euclidean_mean(data)[source]

Log-Euclidean mean of a batch of SPD matrices.

\[\bar{P} = \exp\!\left(\frac{1}{N}\sum_{i=1}^{N} \log(P_i)\right)\]

i.e. the arithmetic mean of the matrix logarithms, mapped back to the SPD manifold with the matrix exponential. Unlike the affine-invariant or Bures-Wasserstein means, this has a closed form (no fixed-point iteration needed).

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Log-Euclidean mean

Return type:

Tensor

LogEuclideanMean(data)[source]

Log-Euclidean mean

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Log-Euclidean mean

Return type:

Tensor

log_euclidean_std_scalar(data, reference_point)[source]

Scalar standard deviation with respect to the Log-Euclidean distance.

\[\sigma = \sqrt{\frac{1}{N}\sum_{i=1}^{N} \lVert \log(P_i) - \log(G) \rVert_F^2}\]

where \(G\) is the reference point (typically the Log-Euclidean mean) and \(\lVert \cdot \rVert_F\) is the Frobenius norm — the Log-Euclidean distance is simply the Euclidean distance between matrix logarithms.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

class LogEuclideanStdScalar(*args, **kwargs)[source]

Scalar standard deviation with respect to the Log-Euclidean distance

static forward(ctx, data, reference_point)[source]

Forward pass of the scalar standard deviation with respect to the Log-Euclidean distance

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

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the scalar standard deviation with respect to the affine-invariant distance

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – gradient of the loss with respect to the input reference point

Return type:

tuple[Tensor, Tensor]

Kullback–Leibler (arithmetic and harmonic)

autograd path

manual backward

euclidean_geodesic()

EuclideanGeodesic

Euclidean geodesic between two symmetric matrices.

arithmetic_mean()

ArithmeticMean

Arithmetic (Euclidean) mean of a batch of symmetric matrices.

left_kullback_leibler_std_scalar()

LeftKullbackLeiblerStdScalar

Scalar standard deviation with respect to the left Kullback-Leibler divergence.

harmonic_curve()

HarmonicCurve

Curve for adaptive harmonic mean computation.

harmonic_mean()

HarmonicMean

Harmonic mean of a batch of SPD matrices.

right_kullback_leibler_std_scalar()

RightKullbackLeiblerStdScalar

Scalar standard deviation with respect to the right Kullback-Leibler divergence (docstring previously said “left” — copy-paste bug, fixed).

Base geometries: Euclidean geodesic, arithmetic mean, harmonic mean/curve, and KL divergence.

euclidean_geodesic(point1, point2, t)[source]

Euclidean geodesic between two symmetric matrices.

\[\gamma(t) = (1-t)\,P_1 + t\,P_2, \quad t \in [0, 1]\]
Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

Return type:

Tensor

class EuclideanGeodesic(*args, **kwargs)[source]

Euclidean geodesic of a batch of symmetric matrices

static forward(ctx, point1, point2, t)[source]

Forward pass of the Euclidean geodesic

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • point1 (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the Euclidean geodesic

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

Returns:

  • grad_input1 (torch.Tensor of shape (..., n_features, n_features))

  • grad_input2 (torch.Tensor of shape (..., n_features, n_features))

Return type:

tuple[Tensor, Tensor, None]

arithmetic_mean(data)[source]

Arithmetic (Euclidean) mean of a batch of symmetric matrices.

\[\bar{P} = \frac{1}{N}\sum_{i=1}^{N} P_i\]
Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of symmetric matrices. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Arithmetic mean

Return type:

Tensor

class ArithmeticMean(*args, **kwargs)[source]

Arithmetic mean

static forward(ctx, data)[source]

Forward pass of the arithmetic mean 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. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Arithmetic mean

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the arithmetic mean 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 the arithmetic mean

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 type:

Tensor

left_kullback_leibler_std_scalar(data, reference_point)[source]

Scalar standard deviation with respect to the left Kullback-Leibler divergence.

\[\sigma^2 = \frac{1}{N}\sum_{i=1}^{N}\Big[ \operatorname{tr}(G^{-1} P_i) - \log\det(P_i)\Big] + \log\det(G) - n\]

where \(n\) is the matrix dimension and \(G\) is the reference point. Up to a factor 2, this is the average Kullback-Leibler divergence \(\mathrm{KL}\big(\mathcal{N}(0, P_i) \,\|\, \mathcal{N}(0, G)\big)\) between zero-mean Gaussians with the batch’s covariances and the reference covariance — hence “left” (\(G\) is the second argument of the KL divergence).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

class LeftKullbackLeiblerStdScalar(*args, **kwargs)[source]

Scalar standard deviation with respect to the left Kullback-Leibler divergence

static forward(ctx, data, reference_point)[source]

Forward pass of the scalar standard deviation with respect to the left Kullback-Leibler divergence

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the scalar standard deviation with respect to the left Kullback-Leibler divergence

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – gradient of the loss with respect to the input reference point

Return type:

tuple[Tensor, Tensor]

harmonic_curve(point1, point2, t)[source]

Curve for adaptive harmonic mean computation.

\[\gamma(t) = \big((1-t)\,P_1^{-1} + t\,P_2^{-1}\big)^{-1}\]

the harmonic analogue of euclidean_geodesic(): a Euclidean geodesic between the matrix inverses, inverted back.

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

class HarmonicCurve(*args, **kwargs)[source]

Harmonic curve of two batches of SPD matrices

static forward(ctx, point1, point2, t)[source]

Forward pass of the harmonic curve

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the harmonic curve

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

Returns:

  • grad_input1 (torch.Tensor of shape (..., n_features, n_features))

  • grad_input2 (torch.Tensor of shape (..., n_features, n_features))

Return type:

tuple[Tensor, Tensor, None]

harmonic_mean(data)[source]

Harmonic mean of a batch of SPD matrices.

\[\bar{P} = \left(\frac{1}{N}\sum_{i=1}^{N} P_i^{-1}\right)^{-1}\]

the inverse of the arithmetic mean of the matrix inverses.

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Harmonic mean

Return type:

Tensor

class HarmonicMean(*args, **kwargs)[source]

Harmonic mean

static forward(ctx, data)[source]

Forward pass of the harmonic mean 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. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Harmonic mean

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the harmonic mean 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 the harmonic mean

Returns:

grad_input (torch.Tensor of shape (..., n_features, n_features)) – Gradient of the loss with respect to the input batch of SPD matrices

Return type:

Tensor

right_kullback_leibler_std_scalar(data, reference_point)[source]

Scalar standard deviation with respect to the right Kullback-Leibler divergence (docstring previously said “left” — copy-paste bug, fixed).

\[\sigma^2 = \frac{1}{N}\sum_{i=1}^{N}\Big[ \operatorname{tr}(P_i^{-1} G) + \log\det(P_i)\Big] - \log\det(G) - n\]

Up to a factor 2, this is the average reverse Kullback-Leibler divergence \(\mathrm{KL}\big(\mathcal{N}(0, G) \,\|\, \mathcal{N}(0, P_i)\big)\) — the roles of \(G\) and \(P_i\) are swapped relative to left_kullback_leibler_std_scalar().

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

class RightKullbackLeiblerStdScalar(*args, **kwargs)[source]

Scalar standard deviation with respect to the left Kullback-Leibler divergence

static forward(ctx, data, reference_point)[source]

Forward pass of the scalar standard deviation with respect to the right Kullback-Leibler divergence

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the scalar standard deviation with respect to the right Kullback-Leibler divergence

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – gradient of the loss with respect to the input reference point

Return type:

tuple[Tensor, Tensor]

Symmetrized Kullback–Leibler (GAH)

autograd path

manual backward

GeometricEuclideanHarmonicCurve()

–

Curve corresponding to the geometric mean of the Euclidean geodesic and the harmonic curve

GeometricArithmeticHarmonicMean()

–

Geometric mean of the arithmetic and harmonic means

symmetrized_kullback_leibler_std_scalar()

SymmetrizedKullbackLeiblerStdScalar

Scalar standard deviation with respect to the symmetrized Kullback-Leibler (Jeffreys) divergence.

adaptive_geometric_arithmetic_harmonic_geodesic()

AdaptiveGeometricArithmeticHarmonicGeodesic

Adaptive (AdaptiveGAH) geodesic between harmonic (point1) and arithmetic (point2) means, with a learnable position :math:t on the path.

AdaptiveGeometricArithmeticHarmonicMean()

–

Adaptive geometric mean of the arithmetic and harmonic means with learnable t (using custom autograd Function for efficient gradients)

Symmetrized KL geometry: GAH curves/means and adaptive geodesic with learnable parameter.

geometric_euclidean_harmonic_curve(point1, point2, t)[source]

Curve corresponding to the geometric mean of the Euclidean geodesic and the harmonic curve (the GAH — Geometric-Arithmetic-Harmonic — curve).

\[\gamma(t) = G\big(E(t),\ H(t)\big), \qquad E(t) = (1-t)P_1 + tP_2, \qquad H(t) = \big((1-t)P_1^{-1} + tP_2^{-1}\big)^{-1}\]

where \(E\) is the Euclidean geodesic (euclidean_geodesic()), \(H\) is the harmonic curve (harmonic_curve()), and \(G\) is the affine-invariant geometric mean of two points (affine_invariant_mean_2points()).

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

GeometricEuclideanHarmonicCurve(point1, point2, t)[source]

Curve corresponding to the geometric mean of the Euclidean geodesic and the harmonic curve

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

geometric_arithmetic_harmonic_mean(data)[source]

Geometric mean of the arithmetic and harmonic means of a batch of SPD matrices (GAH mean).

\[\bar{P}_{GAH} = G\big(\bar{P}_{arith},\ \bar{P}_{harm}\big)\]

with \(\bar{P}_{arith}\) the arithmetic mean (arithmetic_mean()), \(\bar{P}_{harm}\) the harmonic mean (harmonic_mean()), and \(G\) the affine-invariant geometric mean of two points.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • return_arithmetic_harmonic (bool, optional) – Whether to also return arithmetic and harmonic means (for adptative mean update reasons), by default False

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Geometric mean of the arithmetic and harmonic means

Return type:

Tensor

GeometricArithmeticHarmonicMean(data)[source]

Geometric mean of the arithmetic and harmonic means

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Geometric mean of the arithmetic and harmonic means

Return type:

Tensor

symmetrized_kullback_leibler_std_scalar(data, reference_point)[source]

Scalar standard deviation with respect to the symmetrized Kullback-Leibler (Jeffreys) divergence.

\[\sigma^2 = \frac{1}{N}\sum_{i=1}^{N} \frac{\operatorname{tr}(G^{-1}P_i) + \operatorname{tr}(P_i^{-1}G)}{2} - n\]

the average Jeffreys divergence between the batch’s covariances and the reference point \(G\) — symmetrizing left_kullback_leibler_std_scalar() and right_kullback_leibler_std_scalar() (the \(\log\det\) terms cancel out).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

class SymmetrizedKullbackLeiblerStdScalar(*args, **kwargs)[source]

Scalar standard deviation with respect to the symmetrized Kullback-Leibler divergence

static forward(ctx, data, reference_point)[source]

Forward pass of the scalar standard deviation with respect to the symmetrized Kullback-Leibler divergence

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

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the scalar standard deviation with respect to the symmetrized Kullback-Leibler divergence

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – gradient of the loss with respect to the input reference point

Return type:

tuple[Tensor, Tensor]

adaptive_geometric_arithmetic_harmonic_geodesic(point1, point2, t)[source]

Adaptive (AdaptiveGAH) geodesic between harmonic (point1) and arithmetic (point2) means, with a learnable position \(t\) on the path.

\[\gamma(t) = P_1^{1/2} \big(P_1^{-1/2} P_2 P_1^{-1/2}\big)^{t} P_1^{1/2}\]

This is mathematically identical to affine_invariant_geodesic() — the point of this alias is that here \(t\) is treated as a learnable parameter (interpolating between the harmonic mean at \(t=0\) and the arithmetic mean at \(t=1\)) rather than a fixed schedule value.

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices (typically harmonic mean)

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices (typically arithmetic mean)

  • t (float | torch.Tensor) – parameter on the path, should be in [0,1]. t=0.5 gives the standard geometric mean.

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

class AdaptiveGeometricArithmeticHarmonicGeodesic(*args, **kwargs)[source]

Adaptive geodesic between harmonic and arithmetic means with learnable parameter t.

This is mathematically identical to the affine-invariant geodesic. It delegates to AffineInvariantGeodesic and includes gradient with respect to t for learning.

static forward(ctx, point1, point2, t)[source]

Forward pass of the adaptive GAH geodesic.

Delegates to AffineInvariantGeodesic.

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • point1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices (harmonic mean)

  • point2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – SPD matrices (arithmetic mean)

  • t (torch.Tensor of shape ()) – Learnable parameter on the geodesic, should be in [0,1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the adaptive GAH geodesic.

Uses the same gradient logic as AffineInvariantGeodesic, with the addition of grad_t using broadcasting pattern.

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to the output

Returns:

  • grad_input1 (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to point1 (harmonic mean)

  • grad_input2 (torch.Tensor of shape (..., nfeatures, nfeatures)) – Gradient of the loss with respect to point2 (arithmetic mean)

  • grad_t (torch.Tensor of shape ()) – Gradient of the loss with respect to t

Return type:

tuple[Tensor, Tensor, Tensor]

adaptive_geometric_arithmetic_harmonic_mean(data, t)[source]

Adaptive geometric mean (AdaptiveGAH) of the arithmetic and harmonic means of a batch of SPD matrices, with a learnable parameter \(t\).

\[\bar{P}_{t} = \gamma_{AI}\big(\bar{P}_{harm}, \bar{P}_{arith}, t\big)\]

where \(\gamma_{AI}\) is the affine-invariant geodesic (adaptive_geometric_arithmetic_harmonic_geodesic()). Unlike geometric_arithmetic_harmonic_mean() (fixed at \(t=0.5\)), \(t\) here can be learned, letting the model pick where between the harmonic and arithmetic means the effective “center” of BatchNorm sits.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • t (float | torch.Tensor) – parameter on the geodesic between harmonic (t=0) and arithmetic (t=1) means. t=0.5 gives the standard geometric mean of arithmetic and harmonic means.

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Adaptive geometric mean of the arithmetic and harmonic means

Return type:

Tensor

AdaptiveGeometricArithmeticHarmonicMean(data, t)[source]

Adaptive geometric mean of the arithmetic and harmonic means with learnable t (using custom autograd Function for efficient gradients)

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along … axes

  • t (torch.Tensor of shape ()) – Learnable parameter on the geodesic between harmonic (t=0) and arithmetic (t=1) means. t=0.5 gives the standard geometric mean of arithmetic and harmonic means.

Returns:

mean (torch.Tensor of shape (n_features, n_features)) – Adaptive geometric mean of the arithmetic and harmonic means

Return type:

Tensor

Bures–Wasserstein

autograd path

manual backward

bures_wasserstein_distance_squared()

–

Squared Bures-Wasserstein distance between SPD matrices.

bures_wasserstein_log_identity()

–

Logarithmic map at the identity under Bures-Wasserstein geometry.

bures_wasserstein_exp_identity()

–

Exponential map at the identity under Bures-Wasserstein geometry.

bures_wasserstein_log()

–

Logarithmic map at a general base point under BW geometry.

–

LyapunovSolveSPD

Solution :math:Z of :math:BZ + ZB = V for SPD :math:B, with an implicit backward.

bures_wasserstein_parallel_transport_to_identity()

–

Parallel transport from source to the identity under BW geometry.

bures_wasserstein_parallel_transport_from_identity()

–

Parallel transport from the identity to target under BW geometry.

–

ParallelTransportFromIdentityBW

BW transport from the identity, :math:\Gamma_{I\to G} = \mathcal{A}_G^{1/2}, with a backward that stays finite for repeated eigenvalues of :math:G.

bures_wasserstein_geodesic()

–

Bures-Wasserstein geodesic (closed-form 2-sample weighted mean).

BuresWassersteinMean()

–

Bures-Wasserstein barycenter (Frechet mean) – CamelCase wrapper.

bures_wasserstein_std_scalar()

BuresWassersteinStdScalar

Scalar standard deviation under the Bures-Wasserstein distance.

bures_wasserstein_center()

–

Center SPD data by parallel-transporting from the barycenter to the identity.

bures_wasserstein_scale()

–

Scale centered SPD data at the identity.

bures_wasserstein_bias()

–

Bias SPD data by parallel-transporting from the identity to bias_point.

Bures-Wasserstein geometry: geodesic, mean, standard deviation, and transport maps.

Implements the Bures-Wasserstein (BW) metric, geodesic, Frechet mean (barycenter), scalar standard deviation, and related operations (log/exp maps, parallel transport) on the manifold of symmetric positive definite matrices.

References

[1] Bhatia, Jain, Lim. “On the Bures-Wasserstein distance between positive

definite matrices.” Expositiones Mathematicae, 2019.

[2] Kobler et al. “Controlling the Fréchet Variance Improves Batch

Normalization on the Symmetric Positive Definite Manifold.” CVPR, 2022.

bures_wasserstein_distance_squared(point1, point2)[source]

Squared Bures-Wasserstein distance between SPD matrices.

\[d_{BW}^2(X_1, X_2) = \operatorname{tr}(X_1) + \operatorname{tr}(X_2) - 2\operatorname{tr}\!\Big(\big(X_1^{1/2} X_2 X_1^{1/2}\big)^{1/2}\Big)\]

This is the (squared) 2-Wasserstein distance between the zero-mean Gaussian distributions with covariances \(X_1\) and \(X_2\).

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Returns:

dist_sq (torch.Tensor of shape (...)) – Squared BW distances

Return type:

Tensor

bures_wasserstein_log_identity(X)[source]

Logarithmic map at the identity under Bures-Wasserstein geometry.

\[\mathrm{Log}_I(X) = 2\big(X^{1/2} - I\big)\]
Parameters:

X (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Returns:

S (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices in the tangent space at I

Return type:

Tensor

bures_wasserstein_exp_identity(S)[source]

Exponential map at the identity under Bures-Wasserstein geometry.

\[\mathrm{Exp}_I(S) = \left(I + \frac{S}{2}\right)^2\]

Inverse of bures_wasserstein_log_identity().

Parameters:

S (torch.Tensor of shape (..., n_features, n_features)) – Symmetric matrices in the tangent space at I

Returns:

X (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

Return type:

Tensor

bures_wasserstein_log(X, base)[source]

Logarithmic map at a general base point under BW geometry.

\[\mathrm{Log}_B(X) = (XB)^{1/2} + (BX)^{1/2} - 2B\]

computed via the identity \((BX)^{1/2} = B^{1/2} (B^{1/2} X B^{1/2})^{1/2} B^{-1/2}\) and \((XB)^{1/2} = \big[(BX)^{1/2}\big]^\top\).

Parameters:
  • X (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • base (torch.Tensor of shape (..., n_features, n_features) or (n_features, n_features)) – Base point (SPD matrix)

Returns:

tangent (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at base

Return type:

Tensor

class LyapunovSolveSPD(*args, **kwargs)[source]

Solution \(Z\) of \(BZ + ZB = V\) for SPD \(B\), with an implicit backward.

Differentiating the equation gives \(B\,dZ + dZ\,B = dV - (dB\,Z + Z\,dB)\), so with \(W\) solving \(BW + WB = \bar Z\): \(\bar V = W\) and \(\bar B = -(WZ + ZW)\). Only solves with the eigendecomposition of \(B\) are needed, never the derivative of torch.linalg.eigh, so the gradient stays finite when \(B\) has repeated eigenvalues (e.g. a parameter initialized at the identity).

static forward(ctx, base, rhs)[source]
Parameters:
  • base (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices \(B\) (broadcast against rhs)

  • rhs (torch.Tensor of shape (..., n_features, n_features)) – Right-hand sides \(V\)

Returns:

solution (torch.Tensor of shape (..., n_features, n_features)) – Solutions \(Z\)

Return type:

Tensor

static backward(ctx, grad_output)[source]

Define a formula for differentiating the operation with backward mode automatic differentiation.

This function is to be overridden by all subclasses. (Defining this function is equivalent to defining the vjp function.)

It must accept a context ctx as the first argument, followed by as many outputs as the forward() returned (None will be passed in for non tensor outputs of the forward function), and it should return as many tensors, as there were inputs to forward(). Each argument is the gradient w.r.t the given output, and each returned value should be the gradient w.r.t. the corresponding input. If an input is not a Tensor or is a Tensor not requiring grads, you can just pass None as a gradient for that input.

The context can be used to retrieve tensors saved during the forward pass. It also has an attribute ctx.needs_input_grad as a tuple of booleans representing whether each input needs gradient. E.g., backward() will have ctx.needs_input_grad[0] = True if the first input to forward() needs gradient computed w.r.t. the output.

Return type:

tuple[Tensor, Tensor]

bures_wasserstein_parallel_transport_to_identity(tangent_vec, source)[source]

Parallel transport from source to the identity under BW geometry.

Given \(\text{source} = V \operatorname{diag}(\lambda) V^\top\):

\[\Gamma_{\text{source}\to I}(S) = V\left[ \sqrt{\frac{2}{\lambda_i + \lambda_j}} \,(V^\top S V)_{ij} \right]_{ij} V^\top\]
Parameters:
  • tangent_vec (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at source

  • source (torch.Tensor of shape (..., n_features, n_features) or (n_features, n_features)) – Source point (SPD matrix)

Returns:

transported (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at identity

Return type:

Tensor

bures_wasserstein_parallel_transport_from_identity(tangent_vec, target)[source]

Parallel transport from the identity to target under BW geometry.

Given \(\text{target} = U \operatorname{diag}(\delta) U^\top\):

\[\Gamma_{I\to\text{target}}(S) = U\left[ \sqrt{\frac{\delta_i + \delta_j}{2}} \,(U^\top S U)_{ij} \right]_{ij} U^\top\]

Inverse of bures_wasserstein_parallel_transport_to_identity().

Parameters:
  • tangent_vec (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at identity

  • target (torch.Tensor of shape (..., n_features, n_features) or (n_features, n_features)) – Target point (SPD matrix)

Returns:

transported (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at target

Return type:

Tensor

class ParallelTransportFromIdentityBW(*args, **kwargs)[source]

BW transport from the identity, \(\Gamma_{I\to G} = \mathcal{A}_G^{1/2}\), with a backward that stays finite for repeated eigenvalues of \(G\).

\(\mathcal{A}_G(S) = (GS + SG)/2\) is the Lyapunov operator; in the eigenbasis \(u_i\) of \(G\) it is diagonal with entries \(a_{ij} = (\delta_i + \delta_j)/2\), and the transport multiplies \((U^\top S U)_{ij}\) by \(b_{ij} = \sqrt{a_{ij}}\). Its derivative with respect to \(G\) follows from the Sylvester equation \(\mathcal{B}\,d\mathcal{B} + d\mathcal{B}\,\mathcal{B} = d\mathcal{A}\) between operators (\(\mathcal{B} = \mathcal{A}^{1/2}\)):

\[d\tilde W_{ij} = \frac12 \sum_k \frac{\tilde E_{ik} \tilde S_{kj}}{b_{ij} + b_{kj}} + \frac12 \sum_k \frac{\tilde S_{ik} \tilde E_{kj}}{b_{ij} + b_{ik}}, \qquad \tilde E = U^\top dG\, U,\]

whose denominators are positive, unlike the \(1/(\delta_i - \delta_j)\) terms of the autograd path through torch.linalg.eigh.

static forward(ctx, tangent_vec, target)[source]
Parameters:
  • tangent_vec (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors \(S\) at the identity

  • target (torch.Tensor of shape (n_features, n_features)) – Target point \(G\) (SPD)

Returns:

transported (torch.Tensor of shape (..., n_features, n_features)) – Tangent vectors at target

Return type:

Tensor

static backward(ctx, grad_output)[source]

Define a formula for differentiating the operation with backward mode automatic differentiation.

This function is to be overridden by all subclasses. (Defining this function is equivalent to defining the vjp function.)

It must accept a context ctx as the first argument, followed by as many outputs as the forward() returned (None will be passed in for non tensor outputs of the forward function), and it should return as many tensors, as there were inputs to forward(). Each argument is the gradient w.r.t the given output, and each returned value should be the gradient w.r.t. the corresponding input. If an input is not a Tensor or is a Tensor not requiring grads, you can just pass None as a gradient for that input.

The context can be used to retrieve tensors saved during the forward pass. It also has an attribute ctx.needs_input_grad as a tuple of booleans representing whether each input needs gradient. E.g., backward() will have ctx.needs_input_grad[0] = True if the first input to forward() needs gradient computed w.r.t. the output.

Return type:

tuple[Tensor, Tensor]

bures_wasserstein_geodesic(point1, point2, t)[source]

Bures-Wasserstein geodesic (closed-form 2-sample weighted mean).

\[E_2(X_1, X_2; t) = (1-t)^2 X_1 + t^2 X_2 + t(1-t)\Big[(X_2 X_1)^{1/2} + (X_1 X_2)^{1/2}\Big]\]

where \((X_1 X_2)^{1/2} = X_1^{1/2} (X_1^{1/2} X_2 X_1^{1/2})^{1/2} X_1^{-1/2}\) and \(t \in [0, 1]\).

Parameters:
  • point1 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • point2 (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • t (float | torch.Tensor) – Interpolation parameter in [0, 1]

Returns:

point (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices on the geodesic

Return type:

Tensor

bures_wasserstein_mean(data, n_iterations=1)[source]

Bures-Wasserstein barycenter (Fréchet mean) via fixed-point iteration.

\[G_{k+1} = G_k^{-1/2}\left(\frac{1}{N}\sum_{i=1}^{N} \big(G_k^{1/2} X_i G_k^{1/2}\big)^{1/2}\right)^2 G_k^{-1/2}\]

the unique fixed point of this map is the barycenter minimizing \(\sum_i d_{BW}(G, X_i)^2\) (see bures_wasserstein_distance_squared()). The initial estimate \(G_0\) is the arithmetic mean of data.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along ... axes.

  • n_iterations (int) – Number of fixed-point iterations, by default 1

Returns:

barycenter (torch.Tensor of shape (n_features, n_features)) – BW barycenter

Return type:

Tensor

BuresWassersteinMean(data, n_iterations=1)[source]

Bures-Wasserstein barycenter (Frechet mean) – CamelCase wrapper.

Since the fixed-point iteration composes differentiable operations (eigh, matmul, sqrtm), autograd handles the backward automatically.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices. The mean is computed along ... axes.

  • n_iterations (int) – Number of fixed-point iterations, by default 1

Returns:

barycenter (torch.Tensor of shape (n_features, n_features)) – BW barycenter

Return type:

Tensor

bures_wasserstein_std_scalar(data, reference_point)[source]

Scalar standard deviation under the Bures-Wasserstein distance.

\[\sigma = \sqrt{\frac{1}{N}\sum_{i=1}^{N} d_{BW}^2(G, X_i)}\]

where \(G\) is the reference point (typically the BW barycenter) and \(d_{BW}\) is the Bures-Wasserstein distance (see bures_wasserstein_distance_squared()).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – Scalar standard deviation

Return type:

Tensor

class BuresWassersteinStdScalar(*args, **kwargs)[source]

Scalar standard deviation under the Bures-Wasserstein distance (manual backward).

static forward(ctx, data, reference_point)[source]

Forward pass of the BW scalar standard deviation

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

  • reference_point (torch.Tensor of shape (n_features, n_features)) – SPD matrix (some kind of mean of data)

Returns:

scalar_std (torch.Tensor of shape ()) – Scalar standard deviation

Return type:

Tensor

static backward(ctx, grad_output)[source]

Backward pass of the BW scalar standard deviation

Uses:

d(d_BW^2(B,X)) / dX = I - B^{1/2} (B^{1/2} X B^{1/2})^{-1/2} B^{1/2} d(d_BW^2(B,X)) / dB = I - X^{1/2} (X^{1/2} B X^{1/2})^{-1/2} X^{1/2}

Parameters:
  • ctx (torch.autograd.function._ContextMethodMixin) – Context object to retrieve tensors saved during the forward pass

  • grad_output (torch.Tensor of shape ()) – Gradient of the loss with respect to the scalar std

Returns:

  • grad_input_data (torch.Tensor of shape (..., n_features, n_features)) – Gradient of the loss with respect to the input data

  • grad_input_reference_point (torch.Tensor of shape (n_features, n_features)) – Gradient of the loss with respect to the reference point

Return type:

tuple[Tensor, Tensor]

bures_wasserstein_center(data, barycenter)[source]

Center SPD data by parallel-transporting from the barycenter to the identity.

\[X_{\text{centered}} = \mathrm{Exp}_I\big(\Gamma_{B\to I}(\mathrm{Log}_B(X))\big)\]

composing bures_wasserstein_log(), bures_wasserstein_parallel_transport_to_identity(), and bures_wasserstein_exp_identity() — the BatchNorm analogue of subtracting the mean, but on the SPD manifold.

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices

  • barycenter (torch.Tensor of shape (n_features, n_features)) – BW barycenter

Returns:

centered (torch.Tensor of shape (..., n_features, n_features)) – Centered SPD matrices (around the identity)

Return type:

Tensor

bures_wasserstein_scale(data, variance, shift, eps=1e-05)[source]

Scale centered SPD data at the identity.

\[X_{\text{scaled}} = \mathrm{Exp}_I\!\left( \frac{s}{\sqrt{\text{var} + \epsilon}}\, \mathrm{Log}_I(X)\right) = \left(I + \frac{s}{\sqrt{\text{var} + \epsilon}}\, \big(X^{1/2} - I\big)\right)^2\]

the BatchNorm analogue of dividing by the standard deviation and multiplying by a learnable scale \(s\).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – Centered SPD matrices (around the identity)

  • variance (torch.Tensor of shape ()) – Frechet variance

  • shift (torch.Tensor of shape ()) – Learnable scaling parameter

  • eps (float) – Small constant for numerical stability, by default 1e-5

Returns:

scaled (torch.Tensor of shape (..., n_features, n_features)) – Scaled SPD matrices

Return type:

Tensor

bures_wasserstein_bias(data, bias_point)[source]

Bias SPD data by parallel-transporting from the identity to bias_point.

\[X_{\text{biased}} = \mathrm{Exp}_G\big(\Gamma_{I\to G}(\mathrm{Log}_I(X))\big)\]

where \(\mathrm{Exp}_G(V) = G + V + Z G Z\) with \(G Z + Z G = V\) (see bures_wasserstein_parallel_transport_from_identity()). The BatchNorm analogue of adding a learnable bias \(G\).

Parameters:
  • data (torch.Tensor of shape (..., n_features, n_features)) – SPD matrices around the identity

  • bias_point (torch.Tensor of shape (n_features, n_features)) – Learned bias (SPD matrix)

Returns:

biased (torch.Tensor of shape (..., n_features, n_features)) – Biased SPD matrices

Return type:

Tensor