modax¶
GPU-accelerated ODE solvers for massive ensembles (1-100k) of low-dimensional (<200D) ODE trajectories, built on JAX and Numba-CUDA-MLIR. Applications include: Bayesian parameter inference, uncertainty quantification and the integration of physically uncoupled systems.
Every solver is a hand-written CUDA custom kernel compiled by
Numba-CUDA-MLIR: one CUDA thread
per trajectory, hand-written step kernels with in-kernel LU factorisation,
exposed to JAX as an XLA FFI custom call. That binding makes each solver an
ordinary JAX primitive — jit-traceable, and vmap over a single solve lowers
to one native ensemble launch.
-
Install it, and run the first solve.
-
solve(...), its arguments, and what the callbacks have to look like. -
One
sparsitypattern buys both a coloured Jacobian and a compiled sparse direct solve. -
jax.gradand friends, through a continuous forward-sensitivity solve.
Solvers (modax/)¶
| Method | Type | Use for | File |
|---|---|---|---|
| Tsit5 | Explicit RK (order 5) | Non-stiff systems | tsit5.py |
| Rodas5P | Rosenbrock-W (order 5) | Stiff systems | rodas5P.py |
Rodas5P supports an lu_precision ("fp32"/"fp64") knob: the FP32
factorisation halves shared-memory use without lowering method order, since the
Rosenbrock order conditions hold under an approximate Jacobian.