Efficient, Accurate and Stable Gradients for Neural ODEs
Sam McCallum ⋅ James Foster
Abstract
Training Neural ODEs requires backpropagating through an ODE solve. The state-of-the-art backpropagation method is recursive checkpointing that balances recomputation with memory cost. Here, we introduce a class of algebraically reversible ODE solvers that significantly improve upon both the time and memory cost of recursive checkpointing. The reversible solvers presented calculate exact gradients, are high-order and numerically stable -- strictly improving on previous reversible architectures. On scientific modeling and time series classification experiments, reversible solvers reduce the training time by more than $2\times$ while using at least $10\times$ less memory on average than recursive checkpointing.
Chat is not available.
Successful Page Load