Mathematical foundations#
GeoJAX represents a Riemannian geometry by an object that describes points,
tangent vectors, a metric, and the maps needed to move between them. The public
classes follow a capability-qualified contract. ManifoldProtocol contains the
retraction-based operations required by optimizers. GeometryProtocol is a
structural interface adding the common exp, log, and dist names, but
satisfying it does not by itself assert that those operations are globally
exact. operation_kind carries that mathematical status at runtime.
GeometryMixin supplies batching and common helpers. Standard geometric
definitions follow do Carmo [1992]; the computational interface
is informed by
Absil et al. [2008] and Boumal [2023].
Let \(\mathcal M\) be a smooth manifold. At \(x\in\mathcal M\), its tangent space is \(T_x\mathcal M\) and its Riemannian metric is
The same manifold can carry more than one metric. For example, GeoJAX provides log-Euclidean and affine-invariant metrics on positive-definite matrices. A geometry object therefore specifies both the set of valid points and the metric-dependent operations on that set.
Representation and dimension#
Three attributes describe the numerical representation.
Member |
Mathematical meaning |
|---|---|
|
Constructor-level description of one point, such as |
|
Array shape of one unbatched point; for |
|
Intrinsic dimension \(\dim(\mathcal M)=\dim(T_x\mathcal M)\) |
The representation may use more numbers than the intrinsic dimension. A point
on \(S^{n-1}\), for example, is stored as \(n\) coordinates constrained by
\(x^\top x=1\), so shape == (n,) while dim == n - 1.
Product is the exception to the single-array representation: its factors,
points, and tangent vectors can be matching tuples, lists, dictionaries, or
nested JAX pytrees.
Points and tangent vectors#
The validation and projection members separate ambient arrays from geometric objects.
belongs(x) and project(z)#
belongs(x) tests the defining constraints of the manifold. Conceptually,
project(z) maps ambient data to a valid point,
where \(\mathcal E\) is the numerical embedding space. It is intended for
initialization and numerical repair. Unless the geometry explicitly says so,
project should not be interpreted as the nearest-point projection for a
particular metric.
is_tangent(x, u) and tangent_project(x, a)#
is_tangent(x, u) checks whether \(u\in T_x\mathcal M\).
tangent_project(x, a) maps an ambient vector to a tangent vector,
For an embedded manifold defined as a level set \(\mathcal M=\{x:F(x)=0\}\),
For quotient geometries, the returned tangent representation may instead be a
horizontal vector in a chosen representative space. This is how Grassmann
uses \(n\times k\) matrices while representing subspaces rather than frames.
Metric operations#
inner(x, u, v) evaluates \(g_x(u,v)\). norm(x, u) is induced by that metric:
lincomb(x, a, u, b, v, ...) forms a tangent-space linear combination,
The generic implementation projects the result back to \(T_x\mathcal M\) to remove numerical drift. Product geometries perform the same operation leaf by leaf.
Exponential, logarithm, and distance#
For \(u\in T_x\mathcal M\), let \(\gamma_u\) be the geodesic satisfying
The exponential map follows that geodesic for one unit of time:
exp(x, u) evaluates this map. The logarithm is a local inverse:
When operation_kind("log") == "exact", log(x, y) returns one selected
shortest initial velocity where that choice is well defined. The corresponding
geodesic distance is
squared_dist(x, y) evaluates \(d(x,y)^2\) without differentiating the final
square root. It is the preferred primitive for smooth losses and Fréchet
objectives because
is well defined even though the derivative of \(d(x,y)\) itself is not defined
at coincidence. Individual geometries use direct coordinate, angle, or
spectral formulas where those are more stable than squaring dist.
Globally, a logarithm can be nonunique or undefined at a cut locus. Antipodal sphere points and Grassmann subspaces with a principal angle \(\pi/2\) are important examples. A closed formula is also not available for every metric: the two Stiefel classes expose convergence information for their iterative endpoint-shooting logarithms. Class-specific behavior is documented in the Geometry guide.
Exact operations and proxies#
Many useful matrix manifolds have efficient retractions but no practical
closed geodesic formulas. GeoJAX follows the Manopt convention: optimization
depends on retr, invretr, and transport, while geodesic operations are
optional mathematically [Boumal et al., 2014]. For a coherent compositional API, retraction-only
classes still expose exp, log, and dist as compatibility aliases, with
machine-readable capability metadata:
Query |
Result |
|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
For a retraction-only class,
The last expression is local and need not be symmetric, so it must not be reported as a geodesic distance. This distinction lets generic optimizers and product manifolds compose uniformly without overstating the available geometry.
"numerical-local" is different from "proxy". It means that an exact
exponential is available and numerical endpoint shooting seeks a local inverse,
but convergence does not certify that the returned geodesic is globally
shortest. The Stiefel and generalized Stiefel logarithms use this status.
Retractions and means#
retr(x, u, t) maps a tangent step back to the manifold. A retraction
\(R_x:T_x\mathcal M\to\mathcal M\) satisfies
It agrees with the exponential to first order and may be cheaper to evaluate. The default GeoJAX implementation uses \(R_x(tu)=\operatorname{Exp}_x(tu)\); individual geometries can override it. Retractions and compatible vector transports are the central abstraction for manifold optimization [Absil et al., 2008, Boumal, 2023].
invretr(x, y) is a selected local inverse of the retraction. It satisfies
\(R_x(R_x^{-1}(y))\approx y\) near \(x\) and provides the displacement needed by
some derivative-free or quasi-Newton constructions. It is not automatically a
Riemannian logarithm.
pair_mean(x, y) applies the available named maps:
It is the selected geodesic midpoint only when exp and log are exact and the
minimizing logarithm is unique. With proxy or numerical-local operations it is
a midpoint-like local construction instead. A sample Fréchet mean minimizes
\(\sum_i w_i d(x,x_i)^2\) and is generally an optimization problem
[Fréchet, 1948, Karcher, 1977].
Transport#
Tangent vectors at different points belong to different vector spaces.
transport(x, y, u) maps
For exact parallel transport along a geodesic, the Levi-Civita connection preserves the metric:
Transport is used by conjugate-gradient and quasi-Newton methods to compare directions constructed at successive iterates.
The protocol intentionally says transport, not parallel transport. Exact
Levi-Civita transport is available when a stable closed formula is known. An
isometric vector transport may be used otherwise, provided the class documents
that distinction. SPDBuresWasserstein follows the latter convention because
general Bures-Wasserstein parallel transport is obtained from a differential
equation rather than a closed endpoint formula.
Matrix Lie groups#
A matrix Lie group has algebraic maps in addition to its Riemannian geometry. For a Lie-algebra element \(A\), the group exponential is the matrix exponential
This need not equal the Riemannian exponential. They coincide on
SpecialOrthogonal because its metric is bi-invariant. On
SpecialEuclidean, GeoJAX uses the canonical product metric on rotations and
translations; its Riemannian geodesic has a straight translation component,
whereas the group exponential couples angular and translational velocity.
Accordingly, the class exposes both exp/log and
group_exp/group_log. The distinction between Lie-group and Riemannian
exponentials is reviewed by Hall [2015].
Autodiff and Riemannian derivatives#
Suppose a cost \(f:\mathcal M\to\mathbb R\) is differentiated through its ambient array representation. Its Riemannian gradient is defined by
egrad_to_rgrad(x, egrad) converts the ambient Euclidean gradient into this
metric-dual tangent vector. On an isometrically embedded manifold this is often
an orthogonal tangent projection; for a non-Euclidean metric, additional
metric factors are required.
ehess_to_rhess(x, egrad, ehess_vec, u) converts an ambient Hessian-vector
product into the Riemannian Hessian action
including connection or embedding-curvature terms when a geometry supplies
them. The mixin fallback only tangent-projects the ambient Hessian-vector
product and has operation_kind("ehess_to_rhess") == "projection". GeoJAX
second-order solvers reject that fallback: use a geometry advertising "exact"
or supply rhess_vec. The same rule applies to the directional derivative of a
user-supplied Riemannian gradient through operation_kind("rgrad_jvp"). See the
optimization guide for the support table.
Random and batched operations#
random_point(key, sample_shape=()) returns points with shape
sample_shape + M.shape. random_tangent(key, x, scale=...) samples an ambient
random vector and maps it to \(T_x\mathcal M\) according to the geometry. These
routines are reproducible because randomness is controlled by explicit JAX
keys.
Core pointwise protocol operations accept points shaped
batch_shape + M.shape. Scalar-valued pointwise operations return
batch_shape, while point and tangent operations preserve the event
dimensions. NumPy-style broadcasting applies to compatible leading shapes.
Reducers such as sample means document their reduction axes separately.
Product applies the pointwise contract leafwise.
Event axes are part of the manifold definition, not broadcast dimensions.
Consequently, belongs and is_tangent return False for malformed event
shapes, while constructors such as project and tangent_project raise
ValueError. For correctly shaped ambient data, including zero or
rank-deficient matrices, project guarantees a finite point accepted by
belongs.
GeometryMixin retains convenience names for fixed-base collections:
These methods delegate to the natively batched operations. Users may also
compose the methods with jax.vmap when a transformation makes the mapped axis
explicit.
Differentiability at numerical singularities#
Closed geometric formulas often contain removable expressions such as \(\sin(r)/r\), \(\sinh(r)/r\), or spectral divided differences with repeated eigenvalues. GeoJAX evaluates their analytic limits and supplies custom JAX derivatives where naïve autodiff would otherwise differentiate an undefined intermediate eigenbasis. Tests cover zero tangents, coincident points, and repeated SPD spectra under both float32 and float64.
This policy does not conceal genuine singularities. Each geometry documents its
cut-locus convention: sphere, Grassmann, and rotation logarithms return
nonfinite values where no unique branch is selected, while Torus deliberately
uses its half-open angular representation to choose one of the two directions
at a component difference of \(\pi\). Quotient constructions can likewise select
a representative through a documented alignment rule.
Equivariant embeddings#
An embedding \(j:\mathcal M\to\mathbb R^N\) is equivariant when a group action on \(\mathcal M\) is carried to a compatible action in the embedding space. Its differential maps tangent vectors by
The ambient Euclidean metric can be pulled back as
An embedding also supplies extrinsic operations such as chordal distance and
projection of an ambient mean back to the manifold. SphereExtrinsic uses the
identity embedding, while GrassmannProjection uses \(j([X])=XX^\top\).
Product manifolds#
For \(\mathcal M=\mathcal M_1\times\cdots\times\mathcal M_r\), tangent vectors split as \(u=(u_1,\ldots,u_r)\) and GeoJAX uses the direct-sum metric
When every factor distance is exact, consequently,
and projection, exponential, logarithm, transport, gradient conversion, and
sampling act independently on the leaves of the matching pytree. Product
capability metadata is exact only when every factor advertises exactness.
Construction rejects leaves that do not satisfy GeometryProtocol; absent
third-party capability metadata is never interpreted as a mathematical
guarantee.
References#
P.-A. Absil, Robert Mahony, and Rodolphe Sepulchre. Optimization Algorithms on Matrix Manifolds. Princeton University Press, 2008. doi:10.1515/9781400830244.
Nicolas Boumal. An Introduction to Optimization on Smooth Manifolds. Cambridge University Press, 2023. doi:10.1017/9781009166164.
Nicolas Boumal, Bamdev Mishra, P.-A. Absil, and Rodolphe Sepulchre. Manopt, a Matlab toolbox for optimization on manifolds. Journal of Machine Learning Research, 15(42):1455–1459, 2014. URL: https://jmlr.org/papers/v15/boumal14a.html.
Manfredo P. do Carmo. Riemannian Geometry. Birkhäuser, 1992. doi:10.1007/978-1-4757-2201-7.
Maurice Fréchet. Les éléments aléatoires de nature quelconque dans un espace distancié. Annales de l'Institut Henri Poincaré, 10(4):215–310, 1948. URL: https://www.numdam.org/item/AIHP_1948__10_4_215_0/.
Brian C. Hall. Lie Groups, Lie Algebras, and Representations: An Elementary Introduction. Volume 222 of Graduate Texts in Mathematics. Springer, 2 edition, 2015. doi:10.1007/978-3-319-13467-3.
Hermann Karcher. Riemannian center of mass and mollifier smoothing. Communications on Pure and Applied Mathematics, 30(5):509–541, 1977. doi:10.1002/cpa.3160300502.