spd_learn.functional.airm_geodesic#

spd_learn.functional.airm_geodesic(A, B, t)[source]#

Geodesic interpolation on the SPD manifold under the Affine-Invariant Riemannian Metric (AIRM).

The AIRM endows the SPD manifold with a geometry that is invariant under congruence transformations \(P \mapsto WPW^\top\) for any invertible matrix \(W\). The geodesic (shortest path) between two SPD matrices \(A\) and \(B\) is given by:

\[\gamma(t) = A^{1/2} (A^{-1/2} B A^{-1/2})^t A^{1/2}\]

for \(t \in [0, 1]\).

Parameters:
  • A (torch.Tensor) – Starting point SPD matrices with shape (…, n, n).

  • B (torch.Tensor) – End point SPD matrices with shape (…, n, n).

  • t (float) – Interpolation parameter. For t = 0, returns A. For t = 1, returns B. For t = 0.5, returns the geodesic midpoint (Riemannian mean of two matrices).

Returns:

Interpolated SPD matrices on the geodesic with shape (…, n, n).

Return type:

torch.Tensor

Notes

The geodesic midpoint at \(t = 0.5\) corresponds to the matrix geometric mean \(A \# B = A^{1/2}(A^{-1/2}BA^{-1/2})^{1/2}A^{1/2}\), which is the unique positive definite solution to the Riccati equation \(XA^{-1}X = B\).

Examples

>>> import torch
>>> from spd_learn.functional.metrics import airm_geodesic
>>> A = torch.eye(3)
>>> B = 4 * torch.eye(3)
>>> mid = airm_geodesic(A, B, 0.5)
>>> print(f"Midpoint diagonal: {torch.diag(mid)}")
Midpoint diagonal: tensor([2., 2., 2.])

See also

airm_distance()

Computes the geodesic distance under AIRM.

exp_map_airm()

Riemannian exponential map under AIRM.

bures_wasserstein_geodesic()

Geodesic under Bures-Wasserstein metric.

log_cholesky_geodesic()

Geodesic under Log-Cholesky metric.

References

See [Pennec et al., 2006], [Bhatia, 2007] for more details.