Skip to content

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.

  • Getting started

    Install it, and run the first solve.

  • Calling a solver

    solve(...), its arguments, and what the callbacks have to look like.

  • Sparse systems

    One sparsity pattern buys both a coloured Jacobian and a compiled sparse direct solve.

  • Gradients

    jax.grad and 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.