GRADSOLVE: fast exact gradients for ODE ensembles on GPUs
cs.MS, cs.DC, cs.LG, cs.NA, math.NA
Submitted: 2026-09-02
Updated: 2026-09-02
Comments: 38 pages, 12 figures. GRADSOLVE available at https://github.com/ECLIPSE-AI4Science/gradsolve
Code: https://github.com/ECLIPSE-AI4Science/gradsolve
License: http://creativecommons.org/licenses/by/4.0/
The gist: Ordinary differential equations (ODEs) underlie models in science and engineering, and many applications need derivatives of their solutions with respect to parameters.
Terminology
Abstract
Ordinary differential equations (ODEs) underlie models in science and engineering, and many applications need derivatives of their solutions with respect to parameters. Ensembles of independent trajectories suit graphics processing units (GPUs), but current GPU software forces a trade-off: the fastest ensemble solvers cannot be differentiated in reverse mode at the speed they solve, and the solvers built for differentiation solve more slowly. No single tool has yet offered a reverse-mode gradient at the speed of a fused-kernel solve. We present GRADSOLVE, an open-source JAX library for solving and reverse-mode differentiating low-dimensional ODE ensembles on NVIDIA GPUs. It records the steps an adaptive solver accepts and differentiates a fixed-step replay of them; the returned gradient is the exact discrete adjoint of those steps, the same derivative Diffrax returns by default, obtained more cheaply from a fixed-length chain than from an adaptive loop. It targets ensembles differentiated many times against one recorded mesh, keeps Diffrax as a fallback, and supports explicit and Rosenbrock integrators. Used as a solver, GRADSOLVE's forward-only kernel ran 2.8x faster than DiffEqGPU.jl; used for gradients, once a record exists, it computed them 5.6-14.1x faster than Diffrax's checkpointed adjoint at matched forward-state accuracy across three GPU generations, the advantage narrowing on large ensembles and, on stiff systems, down to parity at tight accuracy. GRADSOLVE is released at https://github.com/ECLIPSE-AI4Science/gradsolve.
Sources
- Automatic differentiation in machine learning: a survey
- Neural Ordinary Differential Equations
- On Neural Differential Equations
- Equinox: neural networks in JAX via callable PyTrees and filtered transformations
- Adam: A Method for Stochastic Optimization
- torchode: A Parallel ODE Solver for PyTorch
- A Comparison of Automatic Differentiation and Continuous Sensitivity Analysis for Derivatives of Differential Equation Solutions
- Discretize-Optimize vs. Optimize-Discretize for Time-Series Regression and Continuous Normalizing Flows
- PyTorch: An Imperative Style, High-Performance Deep Learning Library
- Universal Differential Equations for Scientific Machine Learning
- Lineax: unified linear solves and linear least-squares in JAX and Equinox
- Optimistix: modular optimisation in JAX and Equinox
- Adaptive Checkpoint Adjoint Method for Gradient Estimation in Neural ODE
- MALI: A memory efficient and reverse accurate integrator for Neural ODEs
Related papers
- Scen-Opt: A Scenario Optimization Toolbox for Data-Driven Convex Programming
- The Pauli Lightcone: Information-Theoretic Error Mitigation Beyond the Autocorrelation
- FP8 is All You Need (Part 2): Full-FP64 3-D FFT on FP8-Generation Tensor CoresThe Integer-Epilogue Wall and the Minimal Hardware That Would Remove It
- Learning to Optimize by Differentiable Programming
- Tensor Network Kernel Machines: A JAX Framework for Machine Learning and Nonlinear System Identification