Hypersolvers: Toward Fast Continuous-Depth Models
Michael Poli, Stefano Massaroli, Atsushi Yamashita, Hajime Asama, Jinkyoo Park
Introduction
The framework of neural ordinary differential equations (Neural ODEs) (Chen et al., 2018) has reinvigorated research in continuous deep learning (Zhang et al., 2014), offering new system–theoretic perspectives on neural network architecture design (Greydanus et al., 2019; Bai et al., 2019; Poli et al., 2019; Cranmer et al., 2020) and generative modeling (Grathwohl et al., 2018; Yang et al., 2019). Despite the successes, Neural ODEs have been met with skepticism, as these models are often slow in both training and inference due to heavy numerical solver overheads. These issues are further exacerbated by applications which require extremely accurate numerical solutions to the differential equations, such as physics–inspired neural networks (Raissi et al., 2019) and continuous normalizing flows (CNFs) (Chen et al., 2018).
Common knowledge within the field is that these models appear too slow in their current form for meaningful large-scale or embedded applications. Several attempts have been made to either directly or indirectly address some of these limitations, such as redefining the forward pass as a root finding problem (Bai et al., 2019), introducing ad hoc regularization terms (Finlay et al., 2020; Massaroli et al., 2020a) and augmenting the state to reduce stiffness of the solutions (Dupont et al., 2019; Massaroli et al., 2020b). Unfortunately, these approaches either give up on the Neural ODE formulation altogether, do not reduce computation overhead sufficiently or introduce additional memory requirements. Although there is no shortage of works utilizing Neural ODEs in forecasting or classification tasks (Yıldız et al., 2019; Jia and Benson, 2019; Kidger et al., 2020), current state–of–the–art is limited to offline applications with no constraints on inference time. In particular, high–potential application domains for Neural ODEs such as control and prediction often deal with tight requirements on inference speed and computation e.g robotics (Hester, 2013) that are not currently within reach. For example, a generic state–of–the–art convolutional Neural ODE takes at least an order of magnitudeCompared with an equivalent–performance ResNet. longer to infer the label of a single MNIST image. This inefficiency results in inference passes far too slow for real–time applications.
Model–solver synergy The interplay between Neural ODEs and numerical solvers has largely been overlooked as research on model variants has been predominant, often treating solver choice as a simple hyper–parameter to be tuned based on empirical observations. Here, we argue for the importance of computational scalability outside of specific Neural ODE architectural modifications, and highlight the synergistic combination of model–solver to be a likely candidate for unlocking the full potential of continuous–depth models. Namely, this work attempts to alleviate computational overheads by introducing the paradigm of Neural ODE hypersolvers; these auxiliary neural networks are trained to solve the initial value problem (IVP) emerging from the forward pass of continuous–depth models. Hypersolvers improve on the computation–correctness trade–off provided by traditional numerical solvers, enabling fast and arbitrarily accurate solutions during inference.
Pareto efficiency The trade–off between solution accuracy and computation is one of the oldest and best–studied topics in the numerics literature (Butcher, 2016) and was mentioned in the seminal work (Chen et al., 2018) as a feature of continuous models. Traditional approaches shift additional compute resources into improved accuracy via higher–order adaptive–step methods (Prince and Dormand, 1981). For the most part, the computation–accuracy pareto front determined by traditional methods has been treated as optimal, allowing practitioners its traversal with different solver choices. We provide theoretical and practical results in support of the pareto efficiency of hypersolvers, measured with respect to both number of function evaluations (NFEs) as well as standard indicators of algorithmic complexity. Fig. 2 provides a comparison of hypersolvers and traditional methods.
Inference speed By leveraging Hypersolved Neural ODEs, we obtain significant speedups on common benchmarks for continuous–depth models. In image classification tasks, inference is sped up by at least one order of magnitude. Additionally, the proposed approach is capable of solving continuous normalizing flow (CNF) (Chen et al., 2018; Grathwohl et al., 2018) sampling in few steps with little–to–no degradation of the sample quality as shown in Fig. 4. Moving beyond computational advantages at inference time, the proposed framework is compatible with continual learning (Parisi et al., 2019) or adversarial learning (Ganin et al., 2016) techniques where model and hypersolver are co–designed and jointly optimized. Sec. 6 provides an overview of this peculiar interplay.
Background: Continuous-Depth Models
We start by introducing necessary background on Neural ODE and numerical integration methods.
We consider the following general Neural ODE formulation (Massaroli et al., 2020b)
where is a function performing the state update.
ODE solvers differ in how this map is constructedNumerical solvers which obey to (2) are called explicit solvers. In example, the Euler method is realized by setting . Note that, higher–order solvers compute iteratively in steps where denotes the order of the solver. For example, in a -th order Runge-Kutta (RK) (Runge, 1895) method is computed as
In classic numerical analysis, two type of metrics are often defined, i.e. the local truncation error
representing the error accumulated in a single step, and the global truncation error is
i.e. the error accumulated in the first steps. Note that for a -th order solver and (Butcher, 2016).
Hypersolvers for Neural ODEs
Hypersolvers for Neural ODEs Hypersolvers offers a computational framework for the interplay between Neural ODEs and their numerical solver. The core idea behind hypersolvers is to introduce an additional neural network to approximate the higher–order terms of a given solver, greatly increasing its accuracy while preserving the computational and memory efficiency. The simplest instance of Hypersolved Neural ODEs is based on Euler scheme:
where is a neural network approximating the second–order term of the Euler method. The derivation of the Euler hypersolver comes naturally from the following. Let be the true solution of (1) at and let such that . From the Taylor expansion of the solution around , i.e.
we deduce that the classic Euler scheme corresponds to the first–order truncation of the above. The Euler hypersolver, instead, aims at approximating the second–order term, reducing the local truncation error of the overall scheme, while avoiding to compute and store further evaluations of , as required by higher order schemes, e.g. RK methods.
A general formulation of Hypersolved Neural ODEs can be obtained extending (2). If we assume to be the update step of a -th order solver, then the general -th order Hypersolved Neural ODE is defined as
(5) Software implementation We implemented hypersolver variants of common low–order explicit ODE solvers, designed for compatibility with the TorchDyn (Poli et al., 2020) librarySupporting reproducibility code is at https://github.com/DiffEqML/diffeqml-research/tree/master/hypersolver. The Appendix further includes a PyTorch (Paszke et al., 2017) module implementation.
2 Training hypersolvers
Assume to have available the exact solution of the Neural ODE evaluated at the mesh points , practically obtained through an adaptive–step solver set up with low tolerances. With these solution checkpoints we construct the training set for the DE solver with tuples:
According to the introduced metrics and , we introduce two types of loss functions aimed at improving each of the metrics.
We first start by defining the residual of the solver (2)
which correspond to a scaled local truncation error without the neural correction term . Then, we can consider a loss measuring the discrepancy between the residual terms and the output of :
If is a approximator of , i.e.
then, the local truncation error of the hypersolver is .
The proof and further theoretical insights are reported in the Appendix.
The second type of hypersolvers training aims at containing the global truncation error by minimizing the difference between the exact and approximated solutions in the whole depth domain , i.e.
It should be noted that trajectory and residual fitting can be combined into a single loss term, depending on the application.
Experimental Evaluation
The evaluation protocol is designed to measure hypersolver pareto efficiency, inference time speedups and generalizability across base solvers. We consider the following general benchmarks for Neural ODEs: standard image classification (Dupont et al., 2019; Massaroli et al., 2020b) and density estimation with continuous normalizing flows (CNFs) (Chen et al., 2018; Grathwohl et al., 2018).
We train standard convolutional Neural ODEs with input–layer augmentation (Massaroli et al., 2020b) on MNIST and CIFAR10 datasets. Following this initial optimization step, 2–layer convolutional Euler hypersolvers, HyperEuler, (4) are trained by residual fitting (6) on epochs of the training dataset with solution mesh length set to . As ground–truth labels, we utilize the solutions obtained via dopri5 with absolute and relative tolerances set to on the same data. The objective of this first task is to show that hypersolvers retain their pareto efficiency when applied in high–dimensional data regimes. Additional details on hyperparameter choice and architectures are provided as supplementary material.
Pareto comparison We analyze pareto efficiency of hypersolvers with respect to both ODE ODE solution accuracy and test task classification accuracy. It should be noted that residual fitting does not require task supervision; indeed, test data could be used for hypersolver training. Nonetheless, we decide to use only training data for residual fitting, in order to confirm hypersolver ability to generalize to unseen initial conditions of the Neural ODE.
Multiply–accumulate operations i.e MACs are used as a general algorithmic complexity measure. We opt for MACs instead of number of function evaluations (NFEs) of the Neural ODE vector field since the latter does not take into account computational overheads due to hypersolver network . It should be noted that for these specific architectures, single evaluations of and correspond to GMACs and GMACs, respectively. HyperEuler is able to generalize to different step sizes not seen during training, which involved a steps over an integration interval of s. Such residual training scheme over residuals corresponds to a computational complexity for HyperEuler of GMACs, highlighted in blue in Fig. 3. As shown in the Figure, HyperEuler enjoys pareto optimality over alternative fixed–step methods. The hypersolver is able to generalize to different step sizes not seen during training, outperforming higher–order methods such as midpoint and RK4 at low NFEs. As expected, even though higher–order methods eventually surpass HyperEuler at higher NFEs as predicted by theoretical bounds, the hypersolver retains its pareto optimality over Euler.
Wall–clock speedups We measure wall–clock solution time speedups of various fixed–step methods over dopri5 for image classification Neural ODEs. Here, absolute time refers to average time across batches of the MNIST test set required to solve the Neural ODE with different numerical schemes.
Each method performs the minimum number of steps to preserve total accuracy loss across the test set to less than . As shown in Fig. 4, HyperEuler solves an MNIST Neural ODE roughly times faster than dopri5 and with comparable accuracy, achieving significant speedups even over its base method Euler. Indeed, Euler requires a larger number of steps due to its pareto inefficiency compared to HyperEuler, leading to a slower overall solve. The measurements presented are collected on a single V100 GPU.
Generalization across base solvers We verify hypersolver capability to generalize across different base solvers of the same order. We consider the general family of second–order explicit methods parametrized by (Süli and Mayers, 2003) as shown in Fig. 6. Employing a parametrizing family for second–order methods instead of specific instances such as midpoint or Heun allows for an analysis of gradual generalization performance as is tuned away from its value corresponding to the chosen base solver. In particular we consider as midpoint, recovered by , as the base solver for the corresponding hypersolver.
Fig. 6 shows average terminal MAPE solution error of MNIST Neural ODEs solved with both various methods as well as a single HyperMidpoint. As with the previous experiments, the error is computed over dopri5 solutions, and averaged across test data batches. HyperMidpoint is then evaluated, without finetuning, by swapping its base solver with other members of the family. The hypersolver generalizes to different base solvers, preserving its pareto efficiency over the entire –family.
2 Lightweight Density Estimation
We consider sampling in the FFJORD (Grathwohl et al., 2018) variant of continuous normalizing flows (Chen et al., 2018) as an additional task to showcase hypersolver performance. We train CNFs closely following the setup of Grathwohl et al. (2018). Then, we optimize two–layer, second–order Heun hypersolvers, HyperHeun, with residuals obtained against dopri5 with absolute tolerance and relative tolerance . The striking result highlighted in Fig. 7 is that with as little as two NFEs, Hypersolved CNFs provide samples that are as accurate as those obtained through the much more computationally expensive dopri5.
Related Work
There is a long line of research leveraging the universal approximation capabilities of neural networks for solving differential equations. A recurrent theme of the existing work (Lagaris et al., 1997, 1998; Li-ying et al., 2007; Li and Li, 2013; Mall and Chakraverty, 2013; Raissi et al., 2018; Qin et al., 2019) is direct utilization of noiseless analytical solutions and evaluations in low dimensional settings. Application specific attempts (Xing and McCue, 2010; Breen et al., 2019; Fang et al., 2020) provide empirical evidence in support of the earlier work, though the approximation task is still cast as a gradient–matching regression problem on noiseless labels. Deep neural network base solvers have also been used in the distributed parameters setting for PDEs (Han et al., 2018; Magill et al., 2018; Weinan and Yu, 2018; Raissi, 2018; Piscopo et al., 2019; Both et al., 2019; Khoo and Ying, 2019; Winovich et al., 2019; Raissi et al., 2019). Techniques to use neural networks for fast simulation of physical systems have been explored in (Grzeszczuk et al., 1998; James and Fatahalian, 2003; Sanchez-Gonzalez et al., 2020). More recent advances involving symbolic regressions include (Winovich et al., 2019; Regazzoni et al., 2019; Long et al., 2019).
The hypersolver approach is different in several key aspects. To the best knowledge of the authors, this represents the first example where neural network solvers show both consistent and significant pareto efficiency improvements over traditional solvers in high–dimensional settings. The performance advantages are demonstrated in the absence of analytic solutions and are supported by theoretical guarantees, ultimately yielding large inference speedups of practical relevance for Neural ODEs.
After seminal research (Sonoda and Murata, 2017; Lu et al., 2017; Chang et al., 2017; Hauser and Ray, 2017; Chen et al., 2018) uncovered and strengthened the connection bewteen ResNets and ODE discretizations, a variety of architecture and objective specific adjustments have been made to the vanilla formulation. The above allow, for example, to accomodate irregular observations in sequence data (Demeester, 2019) or inherit beneficial properties from the corresponding numerical methods (Zhu et al., 2018). Although these approaches share some structural similarities with the Hypersolved formulation (4), the objective is drastically different. Indeed, such models are optimized for task–specific metrics without concern about preserving ODE properties, or developing a synergistic connection between model and solver.
Discussion
Hypersolvers can be leveraged beyond the inference step of continuous–depth models. Here, we provide avenues of further development of the framework.
The source of the computational (and memory) overheads caused by the use of hypersolver is indeed represented by the evaluation of at each solver step. Nonetheless, this overhead (e.g. in terms of multiply–accumulate operations, MACs) decreases as the solver order increases. In fact, in a th order solver where should be evaluated times, is evaluated only once. Let MACf, MACg be indicators of algorithmic complexity of and , respectively. We have that the relative overhead (in terms of MACs) Or is
and O for Thus, the experiments on pareto efficiency and wall–clock speedup using HyperEuler showcased in Sec. 4.1 should be regarded as worst–case scenario, i.e. the most expensive computational–wise.
In this work, we focus on developing hypersolvers as enhancements to fixed–step explicit methods for Neural ODEs. Although this approach is already effective during inference, hypersolvers are not constrained to this setting. Indeed, the proposed framework can be used to systematically blend learning models and numerical solvers beyond the fixed–step, explicit case. In principle, we could employ hypersolvers into predictor–corrector scheme where we may learn higher–order terms of either the (explicit) predictor or the (implicit) corrector, effectively reducing the overall truncation error. Similarly, adaptive stepping might be achieved by augmenting, in example, the Dormand–Prince (dopri5) scheme. dopri5 uses six NFEs to calculate fourth- and fifth-order Runge–Kutta solutions and obtain the error estimate for step adaptation. Here, we could substitute RK5 with an HyperRK4 and/or train a NN to perform the adaptation given the error estimate.
Speeding up continuous–depth model training with hypersolvers involves additional challenges. In particular, it is necessary to ensure that the hypersolver network remains a approximator of residuals across training iterations. A theoretical toolkit to tackle such a task may be offered by continual learning (Parisi et al., 2019).
Consider the problem of approximating the solution of a Neural ODE at training iteration having optimized the hypersolver on flows generated by the model at the previous training step . This setting involves a certifiably smooth transition between tasks that is directly controlled by the learning rate , leading to the following result
Let the model parameters be updated according to the gradient-based optimizer step to minimize a loss function and let be Lipsichitz w.r.t. . Then,
being the Lipschitz constant.
By leveraging the above result, or pretraining the hypersolver on a sufficiently large collection of dynamics, it might be possible to construct a training procedure for Neural ODEs which maximizes hypersolver reuse across training iterations. Similar to other application areas such as language processing (Howard and Ruder, 2018; Devlin et al., 2018), we envision pretraining techniques to play a fundamental part in the search for easy–to–train continuous–depth models.
Hypersolver and Neural ODE training can be carried out jointly during optimization for the main task. Beyond numerical accuracy metrics, other task specific losses can be considered for hypersolvers. In the standard setting, numerical solvers act as adversaries preserving the ODE solution accuracy at the cost of expressivity. Taking this analogy further, we propose adversarial optimization in the form where is the solution at mesh point given by an adaptive step solver. When used either during hypersolver pretraining or as a regularization term for the main task, the above gives rise to emerging behaviors in the dynamics which exploit solver weaknesses. We observe, as briefly discussed in the Appendix, that direct adversarial training teaches to leverage stiffness (Shampine, 2018) of the differential equation to increase the hypersolver solution error.
Conclusion
Computational overheads represent a great obstacle for the utilization of continuous–depth models in large scale or real–time applications. This work develops the novel hypersolver framework, designed to alleviate performance limitations by leveraging the key model–solver interplay of continuous–depth architectures. Hypersolvers, neural networks trained to solve Neural ODEs accurately and with low overhead, improve solution accuracy at a negligible computational cost, ultimately improving pareto efficiency of traditional methods. Indeed, the synergistic combinations of Hypersolvers and Neural ODEs enjoy large speedups during inference steps of standard benchmarks of continuous–depth models, allowing in example accurate sampling from continuous normalizing flows (CNFs) in as little as 2 number of function evaluations (NFEs). Finally, we discuss how the hypesolver paradigm can be extended to enhance Neural ODE training through continual learning, pretraining or joint optimization of model and hypersolver.
Broader Impact
Major application areas for continuous deep learning architectures so far have been generative modeling (Grathwohl et al., 2018) and forecasting, particularly in the context of patient medical data (Jia and Benson, 2019). While these models have an intrinsic interpretability advantages over discrete counterparts, it is important that future iterations preserve these properties in the search for greater scalability. Early adoption of the hypersolver paradigm would speed up widespread utilization of Neural ODEs in these domains, ultimately leading to positive impact in healthcare applications.
Acknowledgment
We thank Patrick Kidger for helpful discussions. This work was supported by the Korea Agency for Infrastructure Technology Advancement (KAIA) grant, funded by the Ministry of Land, Infrastructure and Transport under Grant 19PIYR-B153277-01.
References
We can directly compute the local truncation error for the hypersolver as
A.2 Proof of Proposition 1
For the Lipschitz continuity of , it holds
Appendix B Further Discussion
We provide PyTorch (Paszke et al., 2017) code showcasing a general hypersolver template:
B.2 Adversarial training
Stiffness in differential equations is an important problem of practical relevance as it often requires development of specialized solution methods (Shampine and Gear, 1979; Cash, 2003). While challenging to fully characterize, stiffness occurs when adaptive–step solvers require a high number of solution steps to maintain the error below specified tolerances, in regions where the solution appears otherwise relatively smooth. Indeed, stiff ODEs are generally difficult to solve accurately for fixed–step solvers. Direct adversarial training allows to find and exploit common weaknesses of numerical methods, which in turn improves hypersolver resilience to a wider class of dynamics.
Appendix C Experimental Details
The experiments have been carried out on a machine equipped with a single NVIDIA Tesla V100 GPU and an eight–core Intel Xeon processor. In addition, we measure wall–clock speedups on a few additional hardware setups and found the results to be consistent.
C.1 Additional Experiments
To evaluate the effectiveness of the trajectory fitting method, we consider a Galërkin Neural ODE (Massaroli et al., 2020b) tasked to tracking of a periodic signal . The Neural ODE is optimized with an integral loss of the type in the integration domain . After the initial training of the model, we fit a three–layer HyperEuler of hidden dimensions using a trajectory fitting approach.
Fig. 8 shows that the pareto efficiency in terms of global truncation error is preserved when training with trajectory fitting. In the 10 - 25 NFE range, HyperEuler results more efficient than higher–order solvers such as midpoint and RK4.
C.2 Image Classification
We report a detailed discussion on the hyperparameter and architectural choices made for the image classification experiments. Further pareto efficiency experimental results, measured in NFEs instead of MACs, are provided in Fig. 9. We omit test accuracy loss NFE pareto fronts since hypersolvers avoid test accuracy losses altogether as shown in the main text.
On MNIST, we optimized Neural ODEs for epochs with batch size utilizing the Adam optimizer with learning rate and a cosine annealing scheduler down to at the end of training. On CIFAR10, we utilized a similar strategy, with epochs, batch size and the same optimizer.
The HyperEuler hypersolver has been trained utilizing fitting the residuals of the Dormand–Prince solver (dopri5) (Dormand and Prince, 1980) with absolute and relative tolerances set to . We use the AdamW (Loshchilov and Hutter, 2017) optimizer with and a cosine annealing schedule down to .
The hypersolver training is subdivided into two phases, proceeding as follows. First, we stabilize the optimization by pretrainining the hypersolver on the trajectories generated from a single batch for several iterations, usually . After this initial phase, the data batch is swapped every iterations. This allows the hypersolver to generalize by having access to trajectories generated from different batches of the training set.
We experimented with different numbers of iterations for hypersolver training. Convergence has been observed in as quickly as iterations, corresponding to less than epochs of the MNIST training dataset with batch size . In practice, iterations (or epochs) is sufficient to produce results comparable to the ones shown in Figure 3. A similar discussion applies to CIFAR10.
In the following, we report PyTorch code defining the Neural ODE and hypersolver architectures in full. The code snippets are followed by a text description for accessibility. In MNIST, the architecture takes the form
where the input–augmented layer (Massaroli et al., 2020b) Neural ODE is defined as a sequence of convolutional layers of channel dimensions and kernel size . The complete architecture is then composed of the above defined Neural ODE with a deconvolution layer, and a linear fully–connected layer to output the classification probabilities.
The HyperEuler architecture is simpler and is composed of only a two–layer CNN with parametric–ReLU (PReLU) (He et al., 2015) activation. The input layer channel dimension is whereas the input to , is only augmented to channels. This is because takes a concatenation of which yields channels.
For the CIFAR10 experiments, on the other hand, and the complete architectures are defined as
It should be noted that even though the Neural ODEs achieve comparable results as (Dupont et al., 2019; Massaroli et al., 2020b), the focus of these experiments has not been optimizing for task–performance. Indeed, we observed that HyperEuler obtains similar results to those shown in the main body of the paper and in Figures 3 and 9 across a variety of different . The setup for base solver generalization experiments has been the same as MNIST experiments, with the only major difference being a choice of HyperMidpoint and an evaluation across different base solvers.
To highlight the efficacy of hypersolvers, we utilize the following metrics
Absolute error of the numerical solution at different solution mesh points. These results provide qualitative proof of the higher solution accuracy of hypersolvers across different types of data samples.
Mean absolute percentage error (MAPE) of the terminal solution. Pareto efficiency of hypersolver numerical solutions.
Average test accuracy decrement. We measure the average (across samples) accuracy lost by a transition away from dopri5. The objective has been to show that outside of solution accuracy, hypersolvers offer pareto efficiency over other solvers in terms of task–specific metrics.
C.3 Continuous Normalizing Flows
We optimize continuous normalizing flows (CNF) (Chen et al., 2018) on density estimation tasks, closely following the setup of (Grathwohl et al., 2018). For a complete reference on normalizing flows we refer to (Kobyzev et al., 2019).
In particular, the training for the two–dimensional tasks is carried out for iterations with an Adam optimizer set to constant learning rate . The CNF is constructed with a three–layer MLP of hidden dimensions and the corresponding ODE is solved with dopri5 with absolute and relative tolerances set to for an accurate forward propagation of the log–density change (Chen et al., 2018). We consider several standard two–dimensional densities following (Grathwohl et al., 2018), namely pinwheel, rings, checkerboard and a modified, more challenging circles where the annuli are connected by three curves.
After this initial step, we train an Heun hypersolver for 30000 iterations of residual fitting on backward trajectories utilizing a similar strategy as discussed in the previous subsection. Namely, we leverage AdamW (Loshchilov and Hutter, 2017) with , weight decay and a two–stage training where the data–sample generating the residuals is switched after every iterations.