Getting started¶
Install & run¶
uv sync # CPU
uv sync --extra cuda13 # or --extra cuda12, for GPU
uv run pytest
uv run ruff format && uv run ruff check --fix
The Enzyme-derived Jacobians come from
numba-enzyme-cuda, the
CUDA-enabled fork of numba-enzyme, which is an ordinary PyPI dependency and
carries its own LLVM 15 and Enzyme binaries — nothing has to be built by hand,
and no system LLVM is involved. It provides the numba_enzyme import package,
so upstream numba-enzyme must not be installed alongside it. See
wheels/README.md for what is in the wheel and why.
pip install modax-solvers gets the same set, and there is no system library to
install first: every dependency ships wheels, the AMD ordering included. A GPU
is needed to run a solve.
A GPU is needed to run a solve: every solver is a CUDA kernel. The CPU
install is enough for the tests that exercise the typing and lowering pipeline
without a device (tests/test_examples.py), and for building these docs.
The first solve¶
import jax.numpy as jnp
from modax.rodas5P import solve
# A CUDA-device callback: fixed-size tuples of scalars in and out, `math`
# rather than `numpy`. This is the Robertson problem.
def ode_fn(y, t, p):
y0, y1, y2 = y
k1, k2, k3 = p
f0 = -k1 * y0 + k3 * y1 * y2
f1 = k1 * y0 - k2 * y1 * y1 - k3 * y1 * y2
f2 = k2 * y1 * y1
return f0, f1, f2
y0 = jnp.array([1.0, 0.0, 0.0]) # one initial state ...
params = jnp.tile( # ... and 10k parameter rows
jnp.array([0.04, 3.0e7, 1.0e4]), (10_000, 1)
)
t_span = jnp.logspace(-6, 5, 64)
y = solve(ode_fn, y0, t_span, params) # (10000, 64, 3)
The ensemble dimension comes from whichever of y0 and params is
two-dimensional; both may be, and they must then agree. The result is always
(N, n_save, n_vars).
Where to go next¶
- Calling a solver — the full argument list, and the rules the callbacks obey.
- Writing ODE callbacks — what
numba_cuda_mlirallows, and how a right-hand side that must also run under JAX is written once. - Sparse systems — what a
sparsitypattern is worth on a structured problem. - Gradients — differentiating a solve with respect to
y0andparams.