Training Generative Adversarial Networks by Solving Ordinary Differential Equations
Chongli Qin, Yan Wu, Jost Tobias Springenberg, Andrew Brock, Jeff Donahue, Timothy P. Lillicrap, Pushmeet Kohli
Introduction
The training of Generative Adversarial Networks (GANs) has seen significant advances over the past several years. Most recently, GAN based methods have, for example, demonstrated the ability to generate images with high fidelity and realism such as the work of Brock et al. and Karras et al. . Despite this remarkable progress, there remain many questions regarding the instability of training GANs and their convergence properties.
In this work, we attempt to extend the understanding of GANs by offering a different perspective. We study the continuous-time dynamics induced by gradient descent on the GAN objective for commonly used losses. We find that under mild assumptions, the dynamics should converge in the vicinity of a differential Nash equilibrium, and that the rate of convergence is independent of the rotational part of the dynamics if we can follow the dynamics exactly. We thus hypothesise that the instability in training GANs arises from discretisation of the continuous dynamics, and we should focus on accurate integration to impart stability.
Consistent with this hypothesis, we demonstrate that we can use standard methods for solving ordinary differential equations (ODEs) – such as Runge-Kutta – to solve for GAN parameters. In particular, we observe that more accurate time integration of the ODE yields better convergence as well as better performance overall; a result that is perhaps surprising given that the integrators have to use noisy gradient estimates. We find that the main ingredient we need for stable GAN training is to avoid large integration errors, and a simple regulariser on the generator gradients is sufficient to achieve this. This alleviates the need for hard constraints on the functional space of the discriminator (e.g. spectral normalisation ) and enables GAN training without advanced optimisation techniques.
Overall, the contributions of this paper are as follows:
We present a novel and practical view that frames GAN training as solving ODEs.
We design a regulariser on the gradients to improve numerical integration of the ODE.
We show that higher-order ODE solvers lead to better convergence for GANs. Surprisingly, our algorithm (ODE-GAN) can train GANs to competitive levels without any adaptive optimiser (e.g. Adam ) and explicit functional constraints (Spectral Normalisation).
Background and Notation
We study the GAN objective which is often described as a two-player min-max game with a Nash equilibrium at the saddle point, , of the objective function
Ordinary Differential Equations Induced by GANs
Given the losses from Eq. (2), the evolution of the parameters , following simultaneous gradient descent (GD), is given by the following updates at iteration
where For brevity, we set in our analysis. are optional scaling factors. Then and correspond to the learning rates in the discrete case. Previous work has analysed the dynamics in this discrete case (mostly focusing on the min-max game), see e.g., Mescheder et al. . In contrast, we consider arbitrary loss pairs under the continuous dynamics induced by gradient descent. That is we consider and to have explicit dependence on continuous time. This perspective has been taken for min-max games in Nagarajan and Kolter . With this dependence and . Then, as , Eq. (3) yields a dynamical system described by the ordinary differential equation
Perhaps surprisingly, if we view the problem of training GANs from this perspective we can make the following observation. Assuming we track the dynamical system exactly – and the gradient vector field is bounded – then in the vicinity of a differential Nash equilibriumWe refer to Ratliff et al. for an overview on local and differential Nash equilibria. , converges to this point at a rate independent of the frequency of rotation with respect to the vector field Namely, the magnitude of the imaginary eigenvalues of the Hessian does not affect the rate of convergence..
There are two direct consequences from this observation: First, changing the rate of rotation of the vector field does not change the rate of convergence. Second, if the velocity field has no subspace which is purely rotational, then it suggests that in principle GAN training can be reduced to the problem of accurate time integration. Thus if the dynamics are attracted towards a differential Nash they should converge (though reaching this regime depends on initial conditions).
2 Convergence of GAN Training under Continous Dynamics
We next show that, under mild assumptions, close to a differential Nash equilibrium, the continuous time dynamics from Eq (4) converge to the equilibrium point for GANs under commonly used losses. We further clarify that they locally converge even under a strong rotational field as in Fig. 1, a problem studied in recent GAN literature (see also e.g. Balduzzi et al. , Gemp and Mahadevan and Mescheder et al. ).
We here show that, in the vicinity of a differential Nash equilibrium, the dynamics converge unless the field has a purely rotational subspace. We prove convergence under this assumption for the continuous dynamics induced by GANs using the cross-entropy or the non-saturated loss. This analysis has strong connections to work by Nagarajan and Kolter , where local convergence was analysed in a more restricted setting for min-max games.
Let us consider the dynamics close to a local Nash equilibrium where by definition. We denote the increment as , then where . This Jacobian has the following form for GANs:
where denote the elements of .
Given a linearised vector field of the form shown in Eq. (6) where either and or and ; and is full rank. Following the dynamics will always converge to .
The proof requires three parts: We first show that of this form must be invertible in Lemma A.1, then we show, as a result, in Lemma A.2, that the real part of the eigenvalues for must be strictly greater than 0. The third part follows from the previous Lemmas: as the eigenvalues of , , satisfy the solution of the linear system converges at least at rate . Consequently, as , . ∎
2.2 Convergence Under a Strong Rotational Field
As an illustrative example of Lemma 3.1, and to verify the importance of accurate time integration of the dynamics, we provide a simple toy example with a strong rotational vector field. Consider a two-player game where the loss functions are given by:
When , it has the analytical solution
where and can be determined from initial conditions. Thus the dynamics will converge to the Nash as independent of the initial conditions. In Fig. 1 we compare the difference of using a first order numerical integrator (Euler’s) vs a second order integrator (two-stage Runge-Kutta, also known as Heun’s method) for solving the dynamical system numerically when . When we chose for 200 timesteps, Euler’s method diverges while RK2 converges.
ODE-GANs
In this section we outline our practical algorithm for applying standard ODE solvers to GAN training; we call the resulting model an ODE-GAN. With few exceptions (such as the exponential function), most ODEs cannot be solved in closed-form and instead rely on numerical solvers that integrate the ODEs at discrete steps. To accomplish this, an ODE solver approximates the ODE’s solution as a cumulative sum in small increments.
We denote the following to be an update using an ODE solver; which takes the velocity function , current parameter states and a step-size as input:
Note that is the function for the velocity field defined in Eq. (4). For an Euler integrator the method would simply compute , and is equivalent to simultaneous gradient descent with step-size . However, as we will outline below, higher-order methods such as the fourth-order Runge-Kutta method can also be used. After the update step is computed, we add a small amount of regularisation to further control the truncation error of the numerical integrator. A listing of the full procedure is given in Algorithm 1.
We consider several classical ODE solvers for the experiments in this paper, although any ODE solver may be used with our method.
Different Orders of Numerical Integration We experimented with a range of solvers with different orders. The order controls how the truncation error changes with respect to step size . Namely, the errors of first order methods reduce linearly with decreasing , while the errors of second order methods decrease quadratically. The ODE solvers considered in this paper are: first order Euler’s method, a second order Runge Kutta method – Heun’s method – (RK2), and fourth order Runge-Kutta method (RK4). For details on the explicit updates to the GAN parameters applied by each of the methods we refer to Appendix E. Computational costs for calculating grow with higher-order integrators; in our implementation, the most expensive solver considered (RK4) was less than slower (in wall-clock time) than standard GAN training.
Connections to Existing Methods for Stabilizing GANs We further observe (derivation in Appendix F) that existing methods such as Consensus Optimization , Extragradient and Symplectic Gradient Adjustment can be seen as approximating higher order integrators.
2 Practical Considerations for Stable integration of GANs
To guarantee stable integration, two issues need to be considered: exploding gradients and the noise from mini-batch gradient estimation.
whose gradient wrt. the discriminator parameters is well-defined in the GAN setting. Importantly, since this regulariser vanishes as we approach a differential Nash equilibrium, it does not change the parameters at the equilibrium. Standard implementation of this regulariser incurs extra cost on each Runge-Kutta step even with efficient double-backpropation . Empirically, we found that modifying the first gradient step is sufficient to control integration errors (see Algorithm 1 and Appendix D for details).
Experiments
We compare training with different integrators for a mixture of Gaussians (Fig. 2). We use a two layer MLP (25 units) with ReLU activations, latent dimension of 32 and batch size of 512. This low dimensional problem is solvable both with Euler’s method as well as Runge-Kutta with gradient regularisation – but without regularisation gradients grow large and integration is harmed, see Appendix H.3. In Fig. 2, we see that both the discriminator and generator converge to the Nash payoff value as shown by their corresponding losses; at the same time, all the modes were recovered by both players. As expected, the convergence rate using Euler’s method is slower than RK4.
2 CIFAR-10 and ImageNet
We use 50K samples to evaluate IS/FID. Unless otherwise specified, we use the DCGAN from for the CIFAR-10 dataset and ResNet-based GANs for ImageNet.
As shown in Fig. 3, we find that moving from a first order to a second order integrator can significantly improve training convergence. But when we go past second order we see diminishing returns. We further observe that higher order methods allow for much larger step sizes: Euler’s method becomes unstable with while Heun’s method (RK2) and RK4 do not. On the other hand, if we increase the regularisation weight, the performance gap between Euler and RK2 is reduced (results for higher regularisation tabulated in Table 3 in the appendix). This hints at an implicit effect of the regulariser on the truncation error of the integrator.
2.2 Effects of Gradient Regularisation
The regulariser controls the truncation error by penalising large gradient magnitudes (see Appendix D). We illustrate this with an embedded method (Fehlberg method, comparing errors between 3rd and 2nd order), which tracks the integration error over the course of training. We observe that larger leads to smaller error. For example, gives an average error of , yields with RK4 (see appendix Fig. 11). We depict the effects of using different values in Fig. 5.
2.3 Loss Profiles for the Discriminator and Generator
Our experiments reveal that using more accurate ODE solvers results in loss profiles that differ significantly to curves observed in standard GAN training, as shown in Fig. 5. Strikingly, we find that the discriminator loss and the generator loss stay very close to the values of a Nash equilibrium, which are for the discriminator and for the generator (shown by red lines in Fig. 5, see ). In contrast, the discriminator dominates the game when using the Adam optimiser, evidenced by a continuously decreasing discriminator loss, while the generator loss increases during training. This imbalance correlates with the well-known phenomenon of worsening FID and IS in late stages of training (we show this in Table 3, see also e.g. Arjovsky and Bottou ).
2.4 Comparison to Standard GAN training
This section compares using ODE solvers versus standard methods for GAN training. Our results challenge the widely-held view that adaptive optimisers are necessary for training GANs, as revealed in Fig. 7. Moreover, the often observed degrading performance towards the end of training disappears with improved integration. To our knowledge, this is the first time that competitive results for GAN training have been demonstrated for image generation without adaptive optimisers. We also compare ODE-GAN with SN-GAN in Fig. 7 (re-trained and tuned using our code for fair comparison) and find that ODE-GAN can improve significantly upon SN-GAN for both IS and FID. Comparisons to more baselines are listed in Table 1. For the DCGAN architecture, ODE-GAN (RK4) achieves 17.66 FID as well as 7.97 IS, note that these best scores are remarkably close to scores we observe at the end of training.
ODE solvers and adaptive optimisers can also be combined. We considered this via the following approach (listed as ODE-GAN(RK4+Adam) in tables): we use the adaptive learning rates computed by Adam to scale the gradients used in the ODE solver. Table 1 shows that this combination reaches similar best IS/FID, but then deteriorates. This observation suggests that Adam can efficiently accelerate training, but the convergence properties may be lost due to the modified gradients – analysing these interactions further is an interesting avenue for future work. For a more detailed comparison of RK4 and Adam, see Table 3 in the appendix.
In Tables 1 and 2, we present results to test ODE-GANs at a larger scale. For CIFAR-10, we experiment with the ResNet architecture from Gulrajani et al. and report the results for baselines using the same model architecture. ODE-GAN achieves 11.85 in FID and 8.61 in IS. Consistent with the behavior we see in DCGAN, ODE-GAN (RK4) results in stable performance throughout training. Additionally, we trained a conditional model on ImageNet 128 128 with the ResNet used in SNGAN without Spectral Normalisation (for further details see Appendix G and H.1). ODE-GAN achieves 26.16 in FID and 38.71 in IS. We have also trained a larger ResNet on ImageNet (see Appendix G), where we can obtain 22.29 for FID and 46.17 for IS using ODE-GAN (see Table 2). Consistent with all previous experiments, we find that the performance is stable and does not degrade over the course of training: see Fig. 8 and Tables 1 and 2.
Discussion and Relation to Existing Work
Our work explores higher-order approximations of the continuous dynamics induced by GAN training. We show that improved convergence and stability can be achieved by faithfully following the vector field of the adversarial game – without the necessity for more involved techniques to stabilise training, such as those considered in Salimans et al. , Balduzzi et al. , Gulrajani et al. , Miyato et al. . Our empirical results thus support the hypothesis that, at least locally, the GAN game is not inherently unstable. Rather, the discretisation of GANs’ continuous dynamics, yielding inaccurate time integration, causes instability. On the empirical side, we demonstrated for training on CIFAR-10 and ImageNet that both Adam and spectral normalisation , two of the most popular techniques, may harm convergence, and that they are not necessary when higher-order ODE solvers are available.
The dynamical systems perspective has been employed for analysing GANs in previous works . They mainly consider simultaneous gradient descent to analyse the discretised dynamics. In contrast, we study the link between GANs and their underlying continuous time dynamics which prompts us to use higher-order integrators in our experiments. Others made related connections: for example, using a second order ODE integrator was also considered in a simple 1-D case for GANs in Gemp and Mahadevan , and Nagarajan and Kolter also analysed the continuous dynamics in a more restrictive setting – in a min-max game around the optimal solution. We hope that our paper can encourage more work in the direction of this connection , and adds to the valuable body of work on analysing GAN training convergence .
Lastly, it is worth noting that viewing traditionally discrete systems through the lens of continuous dynamics has recently attracted attention in other parts of machine learning. For example, the Neural ODE interprets the layered processing of residual neural networks as Euler integration of a continuous system. Similarly, we hope our work can contribute towards establishing a bridge for utilising tools from dynamical systems for generative modelling.
Broader Impact
This work offers a perspective on training generative adversarial networks through the lens of solving an ordinary differential equation. As such it helps us connect an important part of current studies in machine learning (generative modelling) to an old and well studied field of research (integration of dynamical systems).
Making this connection more rigorous over time could help us understand how to better model natural phenomena, see e.g. Zoufal et al. and Casert et al. for recent steps in this direction. Further, tools developed for the analysis of dynamical systems could potentially help reveal in what form exploitable patterns exist in the models we are developing – or their dynamics – and as a result contribute to the goal of learning robust and fair representations .
The techniques proposed in this paper make training of GAN models more stable. This may result in making it easier for non-experts to train such models for beneficial applications like creating realistic images or audio for assistive technologies (e.g. for the speech-impaired, or technology for restoration of historic text sources). On the other hand, the technique could also be used to train models used for nefarious applications, such as forging images and videos (often colloquially referred to as “DeepFakes”). There are some research projects to find ways to mitigate this issue, one example is the DeepFakes Detection Challenge .
Acknowledgments and Disclosure of Funding
We would especially like to thank David Balduzzi for insightful initial discussions and Ian Gemp for careful reading of our paper and feedback on the work.
References
Supplementary
We present details on the proofs from the main paper in Sections A-C. We include analysis on the effects of regularisation on the truncation error (Section D). Update rules for the ODE solvers considered in the main paper are presented in Section E.The connections between our method and Consensus optimisation, SGA and extragradient are reported in Section F. Further details of experiments/additional experimental results are in Section G-H. Image samples are shown in Section I.
Appendix A Real Part of the Eigenvalues are Positive
Given is full rank, if either and or and , is invertible.
Let’s assume that , then we note . Thus by definition and are invertible as there are no zero eigenvalues. The Schur complement decomposition of is given by
Each matrix in this is invertible, thus is invertible. ∎
if either and or vice versa, and is full rank, the positive part of the eigenvalues of is strictly positive.
We assume and . The proof for and is analogous.
To see that the real part of the eigenvalues are positive, we first note that the matrix satisfies the following:
where and are the dimensions of and respectively. From this, we can derive the following property about the eigenvalue with corresponding eigenvector .
Thus we can multiply the top equation by and the bottom equation by to retrieve the following:
To see that this is strictly above zero, we note that we can split the eigenvector via the following:
Thus now can be rewritten as the following:
Since , for this to be zero we note that the following must hold
If this is the form of an eigenvector of then we get the following:
where . One condition which is needed is that . Another condition needed for this to be an eigenvector is that is an eigenvalue of . However this would mean that the eigenvalue is real. If this eigenvalue was thus at , the matrix would not be invertible. Thus concluding the proof. ∎
Appendix B Off-Diagonal Elements are Opposites at the Nash
Here we show that the off diagonal elements of the Hessian with respect to the Wasserstein, cross-entropy and non-saturating loss are opposites. For the cross-entropy and Wasserstein losses, this property hold for zero-sum games by definition. Thus we only show this for the non-saturating loss.
For the non-saturating loss, the objectives for the discriminator and the generator is given by the following:
We transform this with where is now drawn from the probability distribution . We can rewrite this loss as the following:
For the non-saturating loss this becomes the following:
At the global Nash, we know that and , thus this is identically zero.
In this section, we show that when we use piecewise-linear activation functions such as ReLU or LeakyReLUs in the case of the cross-entropy loss or the non saturating loss (as they are the same loss for the discriminator), the Hessian wrt. the discriminator’s parameters will be semi-positive definite. Here we also make the assumption that we are never at the point where the piece-wise function switches state (as the curvature is not defined at these points). Thus we note
For the cross-entropy loss and non-saturating losses, the discriminator network often outputs the logit, , for a sigmoid function . The Hessian with respect to the parameters is given by:
Appendix D Gradient Norm Regularisation and Effects on Integration Error
Here we provide intuition for why gradient regularisation after applying the might be needed. We start by considering how we can approximate the truncation error of Euler’s method with stepsize .
Note here . The truncation error is approximated by comparing this update to that when we half the step size and take two update steps:
Here we see that the truncation error is linear with respect to the magnitude of the gradient. Thus we need the magnitude of the gradient to be bounded in order for the truncation error to be small.
Appendix E Numerical Integration Update Steps
Note corresponds to the vector field element corresponding to , similarly for .
Appendix F Existing Methods which Approximate Second-order ODE Solvers
Here we show that previous methods such as consensus optimisation , SGA or Crossing-the-Curl , and extragradient methods approximates second-order ODE solversNote that methods such as consensus optimisation or SGA/Crossing-the-Curl are methods which address the rotational aspects of the gradient vector field..
To see this let’s consider an “second-order" ODE solver of the following form:
where and scales of each term.
Extra-gradient: Note that when and and , this is the extra-gradient method by definition .
Consensus Optimisation: Note that and , we get Heun’s method (RK2). For consensus optimisation, we only need . If we perform Taylor Expansion on Eq (30):
With this, we see that this approximates consensus optimisation .
SGA/Crossing the Curl: To see we can approximate one part of SGA /Crossing-the-Curl , the update is now given by:
All algorithms are known to improve GAN’s convergence, and we hypothesise that these effects are also related to improving numerical integration.
Appendix G Experimental Setup
We use two GAN architectures for image generation of CIFAR-10: the DCGAN modified by Miyato et al. and the ResNet from Gulrajani et al. , with additional parameters from , but removed spectral-normalisation. For conditional ImageNet generation we use a similar ResNet architecture from Miyato et al. with conditional batch-normalisation and projection but with spectral-normalisation removed. We also consider another ResNet where the size of the hidden layers is the previous ResNet, we also increase the latent size from 128 to 256, we denote this as ResNet (large).
For Euler, RK2 and RK4 integration, we first set for 500 steps. Then we go to till 400k steps and then we decrease the learning rate by half. For Euler integration we found that will be unstable, so we use . We train using batch size 64. The regularisation weight used is . For the Adam optimiser, we use for the generator and for the discriminator, with .
G.2 Hyperparameters for CIFAR-10 (ResNet)
For RK4 first we set for the first 500 steps. Then we go to till 400k steps and then we decrease the learning rate by . We train with batch size 64. The regularisation weight used is . For the Adam optimiser, we use for both the generator and the discriminator, with .
G.3 Hyperparameters for ImageNet (ResNet)
The same hyperparameters are used for ResNet and ResNet (large). For RK4 first we set for the first 15k steps, then we go to . We train with batch size 256. The regularisation weight used is .
Appendix H Additional Experiments
Here we show experiments on ImageNet; ablation studies on using different orders of numerical integrators; effects of gradient regularisation; as well as experiments on combining the RK4 ODE solver with the Adam optimiser.
Fig. 8 shows that ODE-GAN can significantly improve upon SNGAN with respect to both Inception Score (IS) and Fréchet Inception Distance (FID). As in we use conditional batch-normalisation and projection . Similarly to what we have observed in CIFAR-10, the performance degrades over the course of training when we use SNGAN whereas ODE-GAN continues improving. We also note that as we increase the architectural size (i.e. increasing the size of hidden layers and latent size to 256) the IS and FID we obtain using SNGAN gets worse. Whereas, for ODE-GAN we see an improvement in both IS and FID, see Table 2 and Fig. 8. We want to note that our algorithm seems to be more prone to landing in NaNs during training for conditional models, something we would like to further understand in future work.
H.2 Ablation Studies
H.3 Supplementary: Effects of Regularisation
Gradient regularisation allows us to control the magnitude of the gradient. We hypothesise that this helps us control for integration errors. Holding the step size constant, we observe that decreasing the regularisation weight lead to increased gradient norms and integration errors (Fig. 11), causing divergence. This is shown explicitly in Fig. 9, where we show that, with low regularisation weight , the losses for the discriminator and the generator start oscillating heavily around the point where the gradient norm rapidly increases.
We find that our regularisation (Grad Reg) outperforms spectral normalisation (SN) as measured by FID and IS (Fig. 11). Meanwhile, Fig. 11 depicts the integration error (from Fehlberg’s method) over the course of training. As is visible, heavier regularisation leads to smaller integration errors.