GeoJAX
GeoJAX is a JAX-native toolkit for Riemannian geometry and manifold optimization. Geometry objects provide manifold primitives, models provide JAX scalar cost functions, and class-style solvers combine the two.
[!IMPORTANT] GeoJAX is alpha software. The 0.1 series deliberately favors a coherent scientific API over backward compatibility.
Highlights
- Exact geodesic operations where closed forms are available, with machine-readable retraction and transport proxies elsewhere.
- Matrix manifolds, Lie groups, hyperbolic spaces, shape spaces, low-rank models, and arbitrary pytree products.
- JAX autodifferentiation, JIT-compatible geometry primitives, composable batch helpers, and pytree-safe optimization state.
- Executable tutorial pages that place mathematical discussion, code, output, and figures in one document.
Installation
GeoJAX requires Python 3.11 or newer. The human-facing project name is
GeoJAX; the Python package, repository slug, and PyPI distribution are all
lowercase geojax. Install it with:
PyPI
python -m pip install geojax
GitHub source
git clone https://github.com/kisungyou/geojax.git
cd geojax
python -m pip install .
For development, replace the final command with
python -m pip install -e ".[dev,docs,examples]".
Quick Start
The following problem minimizes a Rayleigh quotient on the unit sphere. Its
solution is a dominant eigenvector of A.
import jax.numpy as jnp
from geojax.geometry import Sphere
from geojax.optimization import ConjugateGradient, Minimize
M = Sphere(size=3)
A = jnp.array(
[[3.0, 1.0, 0.0],
[1.0, 2.0, 0.0],
[0.0, 0.0, 0.5]]
)
problem = Minimize(
M=M,
cost=lambda x: -x @ A @ x,
solver=ConjugateGradient(verbosity=0),
key=0,
)
x_hat, final_cost, history = problem.solve()
Sphere(size=3) is the unit sphere embedded in R^3. Minimize obtains the
ambient derivative with JAX, converts it through the geometry, and delegates
the iteration to the selected solver.
JAX Transformations
Geometry instances are static configuration objects. Their numerical protocol
methods accept array or pytree arguments and compose with jax.jit; random
generation takes explicit PRNG keys, and exp_batch, log_batch, and
dist_batch compose jax.vmap with JIT compilation. GeoJAX also supports
JAX-derived gradients and Hessian-vector products on array and Product states.
Solver solve() methods are deliberately Python drivers rather than
whole-solver JIT kernels. They perform stopping checks, callbacks, timing,
line-search control flow, and conversion of diagnostics to Python scalars.
Costs, derivative callbacks, and geometry operations used inside those drivers
may still be independently JIT compiled.
Scientific Scope
The public geometry namespace includes Euclidean, sphere, oblique, simplex, hyperbolic, torus, Grassmann, Stiefel, generalized orthogonality, Lie-group, SPD, fixed-rank, elliptope, spectrahedron, correlation, Kendall-shape, and pytree product geometries.
The public solver set is:
SteepestDescentConjugateGradientTrustRegionsBarzilaiBorweinLBFGSNewtonCGParticleSwarmNelderMeadAdaptiveRegularizationCubicsGaussNewtonandLevenbergMarquardtforLeastSquaresStochasticGradientforFiniteSumAlternatingGradientforProductgeometries
Gradient solvers share public fixed-step, Armijo, adaptive Armijo, and strong-Wolfe line-search strategies.
See the geometry guide, optimization guide, and executable tutorials for the mathematical and computational conventions.
Documentation
The documentation website is published at kisungyou.com/geojax.
Install the documentation dependencies and build the complete site locally:
python -m pip install -e ".[docs,examples]"
make website
python -m http.server 8000 --directory site
The build executes every tutorial, fails on Sphinx warnings, and audits the
generated HTML for malformed mathematics and broken local references. Open
http://127.0.0.1:8000 after the server starts.
Development
python -m pip install -e ".[dev,docs,examples]"
make test
make test-float32
ruff check geojax tests
python -m build
python -m twine check dist/*
Before a release, make test-matrix exercises the supported Python versions,
the declared dependency floor, current dependencies, and both JAX precision
modes. See the
testing guide
for the exact matrix.
The geometry protocol and optimization protocol describe the contracts expected from new implementations. Maintainers can follow the release checklist for manual TestPyPI and PyPI publication.
Citation
Academic users can cite the project using
CITATION.cff.
Release history is recorded in
CHANGELOG.md.
License
GeoJAX is released under the
MIT License.
Licenses and attribution for documentation data and fonts are recorded in
THIRD_PARTY_NOTICES.md.
Metadata
Release files for geojax 0.1.1
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| geojax-0.1.1.tar.gz | 3.4 MB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| geojax-0.1.1-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 3.5 MB
Release files / geojax-0.1.1.tar.gz
| Download URL | geojax-0.1.1.tar.gz |
|---|---|
| Size | 3.4 MB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
da034e38ce020abfcbc1da58a7b3f04faadc829ef906493fc78a55e4035feb9b
|
|
BLAKE2b-256 checksum How to use checksums |
92b05f61b0d5ae25d6ab42af1e990d5543b7acebff079194e565a0bb3a729f86
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/6.2.0 CPython/3.13.11
|
Release files / geojax-0.1.1-py3-none-any.whl
| Download URL | geojax-0.1.1-py3-none-any.whl |
|---|---|
| Size | 101.8 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
75476d99b8194465d838f15d6830b7e972181f700eb3ba7aa27f2a6c315744f7
|
|
BLAKE2b-256 checksum How to use checksums |
1748c78cd27122c09f0cf0b69a6b6a97f2cbed7854c097de65fc72db4344d28c
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/6.2.0 CPython/3.13.11
|