arXiv2026
Astronomers increasingly fit their models with gradient-based methods, such as Hamiltonian Monte Carlo, which need the derivatives of the model with respect to its parameters. In many analyses a prediction requires solving a small system of ordinary differential equations (ODEs) for thousands to millions of parameter sets, and this ensemble of integrations often sets the cost of the analysis. We introduce to the astronomical community GRADSOLVE, a JAX library that moves this computation to graphics processing units (GPUs): it integrates each member of the ensemble in its own GPU thread with its own step size and returns the derivatives with respect to the parameters in the same pass. We measure the speed-up it provides in three examples from different fields of astronomy, stellar orbits in the Galactic potential, the expansion history of a dark-energy cosmology and the spin precession of binary black holes, with every code held to the same accuracy requirement. For a million trajectories GRADSOLVE runs 8.5 to 1500 times faster than the serial CPU code of each example running on all 128 cores of a CPU, and several orders of magnitude faster than on one core. On the same GPU it is also 11 to 15 times faster than DIFFRAX, the state-of-the-art ODE library in JAX, in all three examples. With GRADSOLVE a million such integrations take seconds or less on one GPU, so gradient-based analyses that need ensembles of this size become routine. The code is publicly available at https://github.com/ECLIPSE-AI4Science/gradsolve.