Skip to content

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_mlir allows, and how a right-hand side that must also run under JAX is written once.
  • Sparse systems — what a sparsity pattern is worth on a structured problem.
  • Gradients — differentiating a solve with respect to y0 and params.