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 (Fisher–Rao) |
Karcher mean (iterative) |
|
log-Euclidean |
\(\exp(\overline{\log X})\) |
|
left / right Kullback–Leibler |
arithmetic / harmonic mean |
|
symmetrized Kullback–Leibler |
GAH: AI midpoint of harmonic and arithmetic (adaptive: learned point \(t\)) |
|
Bures–Wasserstein |
fixed-point barycenter |
Affine-invariant¶
autograd path |
manual backward |
|
|---|---|---|
Affine-invariant geodesic between two SPD matrices. |
||
Affine-invariant (geometric) mean of two SPD matrices. |
||
– |
Affine-invariant (geometric) mean computed with fixed-point algorithm |
|
– |
One iteration of the fixed-point algorithm computing the affine-invariant (geometric) mean |
|
Scalar standard deviation with respect to the affine-invariant distance. |
||
Affine-invariant exponential map on the SPD manifold. |
||
– |
Affine-invariant logarithmic map on the SPD manifold. |
|
– |
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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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 passpoint1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricespoint2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricest (
float | torch.Tensor) – Parameter on the geodesic path, should be in [0, 1]
- Returns:
point (
torch.Tensorofshape (...,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 passgrad_output (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to the output
- Returns:
grad_input1 (
torch.Tensorofshape (...,nfeatures,nfeatures)orNone) – Gradient of the loss with respect to point1grad_input2 (
torch.Tensorofshape (...,nfeatures,nfeatures)orNone) – Gradient of the loss with respect to point2grad_t (
torch.Tensorofshape ()orNone) – Gradient of the loss with respect to t (only when t requires grad)
- Return type:
- 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.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricespoint2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matrices
- Returns:
mean (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Geometric means of point1 and point2- Return type:
- 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 passpoint1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricespoint2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matrices
- Returns:
mean (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Geometric means of point1 and point2- Return type:
- 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 passgrad_output (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to the geometric mean of two SPD matrices
- Returns:
grad_input1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to point1grad_input2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to point2
- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axesn_iterations (
int) – Number of iterations to perform to estimate the geometric mean, by default 5
- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – SPD matrix- Return type:
- 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 passmean_iterate (
torch.Tensorofshape (n_features,n_features)) – Current iterate of the affine-invariant meandata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axesstepsize (
float) – step-size to stabilize the fixed-point algorithm
- Returns:
mean_iterate_new (
torch.Tensorofshape (n_features,n_features)) – New iterate of the affine-invariant mean- Return type:
- 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 passgrad_output (
torch.Tensorofshape (nfeatures,nfeatures)) – Gradient of the loss with respect to the new iterate of the affine-invariant mean
- Returns:
grad_input_mean (
torch.Tensorofshape (nfeatures,nfeatures)) – Gradient of the loss with respect to the current iterate of the affine-invariant meangrad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – Gradient of the loss with respect to the data at the current iterate
- Return type:
- AffineInvariantMean(data, n_iterations=5)[source]¶
Affine-invariant (geometric) mean computed with fixed-point algorithm
- Parameters:
data (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axesn_iterations (
int) – Number of iterations to perform to estimate the geometric mean. Default is 10
- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – SPD matrix- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – gradient of the loss with respect to the input reference point
- Return type:
- 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.Tensorofshape (...,n,n)) – Base point(s) on the SPD manifoldtangent (
torch.Tensorofshape (...,n,n)) – Tangent vector(s) at base (symmetric matrices)
- Returns:
result (
torch.Tensorofshape (...,n,n)) – Point(s) on the SPD manifold- Return type:
- 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 backwardbase (
torch.Tensorofshape (...,n,n)) – Base point(s) on the SPD manifoldtangent (
torch.Tensorofshape (...,n,n)) – Tangent vector(s) at base (symmetric matrices)
- Returns:
result (
torch.Tensorofshape (...,n,n)) – Point(s) on the SPD manifold- Return type:
- static backward(ctx, grad_output)[source]¶
Backward pass of the affine-invariant exponential map.
- Parameters:
ctx (
context) – Context with saved tensorsgrad_output (
torch.Tensorofshape (...,n,n)) – Gradient w.r.t. the output
- Returns:
grad_base (
torch.Tensorofshape (...,n,n)orNone)grad_tangent (
torch.Tensorofshape (...,n,n)orNone)
- Return type:
- 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.Tensorofshape (...,n,n)) – Base point(s) on the SPD manifoldpoint (
torch.Tensorofshape (...,n,n)) – Point(s) on the SPD manifold
- Returns:
tangent (
torch.Tensorofshape (...,n,n)) – Tangent vector(s) at base (symmetric matrices)- Return type:
- 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.Tensorofshape (...,n,n)) – Batch of matrices- Returns:
projected (
torch.Tensorofshape (...,n,n)) – Batch of SPD matrices- Return type:
Log-Euclidean¶
autograd path |
manual backward |
|
|---|---|---|
– |
Log-Euclidean geodesic between two batches of SPD matrices |
|
– |
Log-Euclidean mean |
|
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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- LogEuclideanGeodesic(point1, point2, t)[source]¶
Log-Euclidean geodesic between two batches of SPD matrices
- Parameters:
point1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricespoint2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axes- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Log-Euclidean mean- Return type:
- LogEuclideanMean(data)[source]¶
Log-Euclidean mean
- Parameters:
data (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axes- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Log-Euclidean mean- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – gradient of the loss with respect to the input reference point
- Return type:
Kullback–Leibler (arithmetic and harmonic)¶
autograd path |
manual backward |
|
|---|---|---|
Euclidean geodesic between two symmetric matrices. |
||
Arithmetic (Euclidean) mean of a batch of symmetric matrices. |
||
Scalar standard deviation with respect to the left Kullback-Leibler divergence. |
||
Curve for adaptive harmonic mean computation. |
||
Harmonic mean of a batch of SPD matrices. |
||
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.Tensorofshape (...,n_features,n_features)) – Symmetric matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matrices- Return type:
- 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 passpoint1 (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matrices- Return type:
- 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 passgrad_output (
torch.Tensorofshape (...,n_features,n_features))
- Returns:
grad_input1 (
torch.Tensorofshape (...,n_features,n_features))grad_input2 (
torch.Tensorofshape (...,n_features,n_features))
- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of symmetric matrices. The mean is computed along … axes- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Arithmetic mean- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of symmetric matrices. The mean is computed along … axes
- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Arithmetic mean- Return type:
- 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 passgrad_output (
torch.Tensorofshape (n_features,n_features)) – Gradient of the loss with respect to the arithmetic mean
- Returns:
grad_input (
torch.Tensorofshape (...,n_features,n_features)) – Gradient of the loss with respect to the input batch of symmetric matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – gradient of the loss with respect to the input reference point
- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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 passpoint1 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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 passgrad_output (
torch.Tensorofshape (...,n_features,n_features))
- Returns:
grad_input1 (
torch.Tensorofshape (...,n_features,n_features))grad_input2 (
torch.Tensorofshape (...,n_features,n_features))
- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axes- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Harmonic mean- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axes
- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Harmonic mean- Return type:
- 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 passgrad_output (
torch.Tensorofshape (n_features,n_features)) – Gradient of the loss with respect to the harmonic mean
- Returns:
grad_input (
torch.Tensorofshape (...,n_features,n_features)) – Gradient of the loss with respect to the input batch of SPD matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – gradient of the loss with respect to the input reference point
- Return type:
Symmetrized Kullback–Leibler (GAH)¶
autograd path |
manual backward |
|
|---|---|---|
– |
Curve corresponding to the geometric mean of the Euclidean geodesic and the harmonic curve |
|
– |
Geometric mean of the arithmetic and harmonic means |
|
Scalar standard deviation with respect to the symmetrized Kullback-Leibler (Jeffreys) divergence. |
||
Adaptive (AdaptiveGAH) geodesic between harmonic (point1) and arithmetic (point2) means, with a learnable position :math: |
||
– |
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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- GeometricEuclideanHarmonicCurve(point1, point2, t)[source]¶
Curve corresponding to the geometric mean of the Euclidean geodesic and the harmonic curve
- Parameters:
point1 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – parameter on the path, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axesreturn_arithmetic_harmonic (
bool, optional) – Whether to also return arithmetic and harmonic means (for adptative mean update reasons), by default False
- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Geometric mean of the arithmetic and harmonic means- Return type:
- GeometricArithmeticHarmonicMean(data)[source]¶
Geometric mean of the arithmetic and harmonic means
- Parameters:
data (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axes- Returns:
mean (
torch.Tensorofshape (n_features,n_features)) – Geometric mean of the arithmetic and harmonic means- Return type:
- 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()andright_kullback_leibler_std_scalar()(the \(\log\det\) terms cancel out).- Parameters:
data (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the output of the scalar standard deviation Function
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – gradient of the loss with respect to the input reference point
- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matrices (typically harmonic mean)point2 (
torch.Tensorofshape (...,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.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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 passpoint1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matrices (harmonic mean)point2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – SPD matrices (arithmetic mean)t (
torch.Tensorofshape ()) – Learnable parameter on the geodesic, should be in [0,1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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 passgrad_output (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to the output
- Returns:
grad_input1 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to point1 (harmonic mean)grad_input2 (
torch.Tensorofshape (...,nfeatures,nfeatures)) – Gradient of the loss with respect to point2 (arithmetic mean)grad_t (
torch.Tensorofshape ()) – Gradient of the loss with respect to t
- Return type:
- 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()). Unlikegeometric_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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axest (
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.Tensorofshape (n_features,n_features)) – Adaptive geometric mean of the arithmetic and harmonic means- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices. The mean is computed along … axest (
torch.Tensorofshape ()) – 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.Tensorofshape (n_features,n_features)) – Adaptive geometric mean of the arithmetic and harmonic means- Return type:
Bures–Wasserstein¶
autograd path |
manual backward |
|
|---|---|---|
– |
Squared Bures-Wasserstein distance between SPD matrices. |
|
– |
Logarithmic map at the identity under Bures-Wasserstein geometry. |
|
– |
Exponential map at the identity under Bures-Wasserstein geometry. |
|
– |
Logarithmic map at a general base point under BW geometry. |
|
– |
Solution :math: |
|
– |
Parallel transport from source to the identity under BW geometry. |
|
– |
Parallel transport from the identity to target under BW geometry. |
|
– |
BW transport from the identity, :math: |
|
– |
Bures-Wasserstein geodesic (closed-form 2-sample weighted mean). |
|
– |
Bures-Wasserstein barycenter (Frechet mean) – CamelCase wrapper. |
|
Scalar standard deviation under the Bures-Wasserstein distance. |
||
– |
Center SPD data by parallel-transporting from the barycenter to the identity. |
|
– |
Scale centered SPD data at the identity. |
|
– |
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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices
- Returns:
dist_sq (
torch.Tensorofshape (...)) – Squared BW distances- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matrices- Returns:
S (
torch.Tensorofshape (...,n_features,n_features)) – Symmetric matrices in the tangent space at I- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Symmetric matrices in the tangent space at I- Returns:
X (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matricesbase (
torch.Tensorofshape (...,n_features,n_features)or(n_features,n_features)) – Base point (SPD matrix)
- Returns:
tangent (
torch.Tensorofshape (...,n_features,n_features)) – Tangent vectors at base- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matrices \(B\) (broadcast againstrhs)rhs (
torch.Tensorofshape (...,n_features,n_features)) – Right-hand sides \(V\)
- Returns:
solution (
torch.Tensorofshape (...,n_features,n_features)) – Solutions \(Z\)- Return type:
- 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
vjpfunction.)It must accept a context
ctxas the first argument, followed by as many outputs as theforward()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 toforward(). 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_gradas a tuple of booleans representing whether each input needs gradient. E.g.,backward()will havectx.needs_input_grad[0] = Trueif the first input toforward()needs gradient computed w.r.t. the output.
- 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.Tensorofshape (...,n_features,n_features)) – Tangent vectors at sourcesource (
torch.Tensorofshape (...,n_features,n_features)or(n_features,n_features)) – Source point (SPD matrix)
- Returns:
transported (
torch.Tensorofshape (...,n_features,n_features)) – Tangent vectors at identity- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Tangent vectors at identitytarget (
torch.Tensorofshape (...,n_features,n_features)or(n_features,n_features)) – Target point (SPD matrix)
- Returns:
transported (
torch.Tensorofshape (...,n_features,n_features)) – Tangent vectors at target- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Tangent vectors \(S\) at the identitytarget (
torch.Tensorofshape (n_features,n_features)) – Target point \(G\) (SPD)
- Returns:
transported (
torch.Tensorofshape (...,n_features,n_features)) – Tangent vectors at target- Return type:
- 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
vjpfunction.)It must accept a context
ctxas the first argument, followed by as many outputs as theforward()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 toforward(). 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_gradas a tuple of booleans representing whether each input needs gradient. E.g.,backward()will havectx.needs_input_grad[0] = Trueif the first input toforward()needs gradient computed w.r.t. the output.
- 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.Tensorofshape (...,n_features,n_features)) – SPD matricespoint2 (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricest (
float | torch.Tensor) – Interpolation parameter in [0, 1]
- Returns:
point (
torch.Tensorofshape (...,n_features,n_features)) – SPD matrices on the geodesic- Return type:
- 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.Tensorofshape (...,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.Tensorofshape (n_features,n_features)) – BW barycenter- Return type:
- 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.Tensorofshape (...,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.Tensorofshape (n_features,n_features)) – BW barycenter- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – Scalar standard deviation- Return type:
- 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 passdata (
torch.Tensorofshape (...,n_features,n_features)) – Batch of SPD matricesreference_point (
torch.Tensorofshape (n_features,n_features)) – SPD matrix (some kind of mean of data)
- Returns:
scalar_std (
torch.Tensorofshape ()) – Scalar standard deviation- Return type:
- 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 passgrad_output (
torch.Tensorofshape ()) – Gradient of the loss with respect to the scalar std
- Returns:
grad_input_data (
torch.Tensorofshape (...,n_features,n_features)) – Gradient of the loss with respect to the input datagrad_input_reference_point (
torch.Tensorofshape (n_features,n_features)) – Gradient of the loss with respect to the reference point
- Return type:
- 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(), andbures_wasserstein_exp_identity()— the BatchNorm analogue of subtracting the mean, but on the SPD manifold.- Parameters:
data (
torch.Tensorofshape (...,n_features,n_features)) – SPD matricesbarycenter (
torch.Tensorofshape (n_features,n_features)) – BW barycenter
- Returns:
centered (
torch.Tensorofshape (...,n_features,n_features)) – Centered SPD matrices (around the identity)- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – Centered SPD matrices (around the identity)variance (
torch.Tensorofshape ()) – Frechet varianceshift (
torch.Tensorofshape ()) – Learnable scaling parametereps (
float) – Small constant for numerical stability, by default 1e-5
- Returns:
scaled (
torch.Tensorofshape (...,n_features,n_features)) – Scaled SPD matrices- Return type:
- 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.Tensorofshape (...,n_features,n_features)) – SPD matrices around the identitybias_point (
torch.Tensorofshape (n_features,n_features)) – Learned bias (SPD matrix)
- Returns:
biased (
torch.Tensorofshape (...,n_features,n_features)) – Biased SPD matrices- Return type: