Stabilizing Training of Generative Adversarial Networks through Regularization
Kevin Roth, Aurelien Lucchi, Sebastian Nowozin, Thomas Hofmann
Introduction
The standard way of representing a specific is through a family of statistics or discriminants , typically realized by a neural network . In GANs, we use these discriminators in a logistic classification loss as follows
where is the log-logistic function (for reference, in ).
We consider three different challenges for learning the model distribution:
(B) density misspecification: the model distribution and true distribution both have a density function with respect to the same base measure but there exists no parameter for which these densities are sufficiently similar. Here, the principle of parameter estimation via divergence minimization is provably sound in that it achieves a well-defined limit . It therefore provides a solid foundation for statistical inference that is robust with regard to model misspecifications.
In what follows, we will take Eq. (3) as the starting point and derive an approximation via a regularizer that is simple to implement as an integral operator penalizing the squared gradient norm. As opposed to a naïve norm penalization, each -divergence has its own characteristic weighting function over the input space, which depends on the discriminator output. We demonstrate the effectiveness of our approach on a simple Gaussian mixture as well as on several benchmark image datasets commonly used for generative models. In both cases, our proposed regularization yields stable GAN training and produces samples of higher visual quality. We also perform pairwise tests of regularized vs. unregularized GANs using a novel cross-testing protocol.
In summary, we make the following contributions:
We systematically derive a novel, efficiently computable regularization method for -GAN.
We show how this addresses the dimensional misspecification challenge.
We empirically demonstrate stable GAN training across a broad set of models.
Background
Integral Probability Metrics (IPM).
where the supremum is taken over functions which have a bounded Lipschitz constant.
As shown in , the Wasserstein metric implies a different notion of convergence compared to the JS divergence used in the original GAN. Essentially, the Wasserstein metric is said to be weak as it requires the use of a weaker topology, thus making it easier for a sequence of distribution to converge. The use of a weaker topology is achieved by restricting the function class to the set of bounded Lipschitz functions. This yields a hard constraint on the function class that is empirically hard to satisfy. In , this constraint is implemented via weight clipping, which is acknowledged to be a "terrible way" to enforce the Lipschitz constraint. As will be shown later, our regularization penalty can be seen as a soft constraint on the Lipschitz constant of the function class which is easy to implement in practice. Recently, has also proposed a similar regularization; while their proposal was motivated for Wasserstein GANs and does not extend to -divergences it is interesting to observe that both their and our regularization work on the gradient.
Training with Noise.
Regularization for Mode Dropping.
Other regularization techniques address the problem of mode dropping and are complementary to our approach. This includes the work of which incorporates a supervised training signal as a regularizer on top of the discriminator target. To implement supervision the authors use an additional auto-encoder as well as a two-step training procedure which might be computationally expensive. A similar approach was proposed by that stabilizes GANs by unrolling the optimization of the discriminator. The main drawback of this approach is that the computational cost scales with the number of unrolling steps. In general, it is not clear to what extent these methods not only stabilize GAN training, but also address the conceptual challenges listed in Section 1.
Noise-Induced Regularization
From now onwards, we consider the general -GAN objective defined as
2 Convolved Discriminants
With symmetric noise, , we can write Eq. (8) equivalently as
3 Analytic Approximations
In general, it may be difficult to analytically compute or – equivalently – . However, for small we can use a Taylor approximation of around (cf. ):
where denotes the Hessian, whose trace is known as the Laplace operator. The properties of white noise result in the approximation
and thereby lead directly to an approximation of (see Eq. (3)) via plus a correction, i.e.
We can interpret Eq. (13) as follows: the Laplacian measures how much the scalar fields and differ at each point from their local average. It is thereby an infinitesimal proxy for the (exact) convolution.
The Laplace operator is a sum of terms, where is the dimensionality of the ambient data space. As such it does not suffer from the quadratic blow-up involved in computing the Hessian. If we realize the discriminator via a deep network, however, then we need to be able to compute the Laplacian of composed functions. For concreteness, let us assume that , and look
So at the intermediate layer, we would need to effectively operate with a full Hessian, which is computationally demanding, as has already been observed in .
4 Efficient Gradient-Based Regularization
The relevance of this becomes clear, if we apply the chain rule to , assuming that is twice differentiable
as now we get a convenient cancellation of the Laplacians at
We can (heuristically) turn this into a regularizer by taking the leading terms,
Regularizing GANs
We have shown that training with noise is equivalent to regularizing the discriminator. Inspired by the above analysis, we propose the following class of -GAN regularizers:
Regularizing the discriminator provides an efficient way to convolve the distributions and is thereby sufficient to address the dimensional misspecification challenges outlined in the introduction. This leaves open the possibility to use the regularizer also in the objective of the generator. On the one hand, optimizing the generator through the regularized objective may provide useful gradient signal and therefore accelerate training. On the other hand, it destabilizes training close to convergence (if not dealt with properly), since the generator is incentiviced to put probability mass where the discriminator has large gradients. In the case of JS-GANs, we recommend to pair up the regularized objective of the discriminator with the “alternative” or “non-saturating” objective for the generator, proposed in , which is known to provide strong gradients out of the box (see Algorithm 1).
2 Annealing
The regularizer variance lends itself nicely to annealing. Our experimental results indicate that a reasonable annealing scheme consists in regularizing with a large initial early in training and then (exponentially) decaying to a small non-zero value. We leave to future work the question of how to determine an optimal annealing schedule.
Experiments
To demonstrate the stabilizing effect of the regularizer, we train a simple GAN architecture on a 2D submanifold mixture of seven Gaussians arranged in a circle and embedded in 3D space (further details and an illustration of the mixture distribution are provided in the Appendix). We emphasize that this mixture is degenerate with respect to the base measure defined in ambient space as it does not have fully dimensional support, thus precisely representing one of the failure scenarios commonly described in the literature . The results are shown in Fig. 1 for both standard unregularized GAN training as well as our regularized variant.
While the unregularized GAN collapses in literally every run after around 50k iterations, due to the fact that the discriminator concentrates on ever smaller differences between generated and true data (the stakes are getting higher as training progresses), the regularized variant can be trained essentially indefinitely (well beyond 200k iterations) without collapse for various degrees of noise variance, with and without annealing. The stabilizing effect of the regularizer is even more pronounced when the GANs are trained with five discriminator updates per generator update step, as shown in Fig. 6.
2 Stability across various architectures
To demonstrate the stability of the regularized training procedure and to showcase the excellent quality of the samples generated from it, we trained various network architectures on the CelebA , CIFAR-10 and LSUN bedrooms datasets. In addition to the deep convolutional GAN (DCGAN) of , we trained several common architectures that are known to be hard to train , therefore allowing us to establish a comparison to the concurrently proposed gradient-penalty regularizer for Wasserstein GANs . Among these architectures are a DCGAN without any normalization in either the discriminator or the generator, a DCGAN with tanh activations and a deep residual network (ResNet) GAN . We used the open-source implementation of for our experiments on CelebA and LSUN, with one notable exception: we use batch normalization also for the discriminator (as our regularizer does not depend on the optimal transport plan or more precisely the gradient penalty being imposed along it).
All networks were trained using the Adam optimizer with learning rate and hyperparameters recommended by . We trained all datasets using batches of size 64, for a total of 200K generator iterations in the case of LSUN and 100k iterations on CelebA. The results of these experiments are shown in Figs. 3 & 2. Further implementation details can be found in the Appendix.
3 Training time
We empirically found regularization to increase the overall training time by a marginal factor of roughly (due to the additional backpropagation through the computational graph of the discriminator gradients). More importantly, however, (regularized) -GANs are known to converge (or at least generate good looking samples) faster than their WGAN relatives .
4 Regularization vs. explicitly adding noise
We compare our regularizer against the common practitioner’s approach to explicitly adding noise to images during training. In order to compare both approaches (analytic regularizer vs. explicit noise), we fix a common batch size (64 in our case) and subsequently train with different noise-to-signal ratios (NSR): we take (NSR) samples (both from the dataset and generated ones) to each of which a number of NSR noise vectors is added and feed them to the discriminator (so that overall both models are trained on the same batch size). We experimented with NSR 1, 2, 4, 8 and show the best performing ratio (further ratios in the Appendix). Explicitly adding noise in high-dimensional ambient spaces introduces additional sampling variance which is not present in the regularized variant. The results, shown in Fig. 4, confirm that the regularizer stabilizes across a broad range of noise levels and manages to produce images of considerably higher quality than the unregularized variants.
5 Cross-testing protocol
We propose the following pairwise cross-testing protocol to assess the relative quality of two GAN models: unregularized GAN (Model 1) vs. regularized GAN (Model 2). We first report the confusion matrix (classification of 10k samples from the test set against 10k generated samples) for each model separately. We then classify 10k samples generated by Model 1 with the discriminator of Model 2 and vice versa. For both models, we report the fraction of false positives (FP) (Type I error) and false negatives (FN) (Type II error). The discriminator with the lower FP (and/or lower FN) rate defines the better model, in the sense that it is able to more accurately classify out-of-data samples, which indicates better generalization properties. We obtained the following results on CIFAR-10:
Regularized GAN () True condition Positive Negative Predicted Positive Negative
Unregularized GAN True condition Positive Negative Predicted Positive Negative
For both models, the discriminator is able to recognize his own generator’s samples (low FP in the confusion matrix). The regularized GAN also manages to perfectly classify the unregularized GAN’s samples as fake (cross-testing FP 0.0) whereas the unregularized GAN classifies the samples of the regularized GAN as real (cross-testing FP 1.0). In other words, the regularized model is able to fool the unregularized one, whereas the regularized variant cannot be fooled.
Conclusion
We introduced a regularization scheme to train deep generative models based on generative adversarial networks (GANs). While dimensional misspecifications or non-overlapping support between the data and model distributions can cause severe failure modes for GANs, we showed that this can be addressed by adding a penalty on the weighted gradient-norm of the discriminator. Our main result is a simple yet effective modification of the standard training algorithm for GANs, turning them into reliable building blocks for deep learning that can essentially be trained indefinitely without collapse. Our experiments demonstrate that our regularizer improves stability, prevents GANs from overfitting and therefore leads to better generalization properties (cf cross-testing protocol). Further research on the optimization of GANs as well as their convergence and generalization can readily be built upon our theoretical results.
Acknowledgements
We would like to thank Devon Hjelm for pointing out that the regularizer works well with ResNets. KR is thankful to Yannic Kilcher, Lars Mescheder, Paulina Grnarova and the dalab team for insightful discussions. Big thanks also to Ishaan Gulrajani and Taehoon Kim for their open-source GAN implementations. This work was supported by Microsoft Research through its PhD Scholarship Programme.
References
APPENDIX
The Jensen-Shannon GAN is typically encountered in one of two equivalent parametrizations: the commonly used “original” GAN parametrization ,
and the Fenchel-dual -GAN parametrization , where ,
Depending on whether we train the JS-GAN through its Fenchel-dual parametrization or with the original GAN objective we have to either use the regularizer in the general -GAN form in Eq. (19), or in the specific Jensen-Shannon parametrization given in Eq. (20).
We now show how to derive the specific Jensen-Shannon regularizer, which is basically a repetition of the derivation of the general -GAN regularizer presented in the main text. Using the same terminology and following the same line of thought as in section 3.3, i.e. assuming the noise variance is small, we can Taylor approximate the statistics around ,
Expanding the noise convolved version of the objective in Eq. (21) and making use of the zero-mean and uncorrelatedness properties of the noise distribution, as well as applying the chain rule, we obtain
Following the same arguments as in section 3.4, one can again show that the Laplacian terms cancel at and we arrive at
allowing us to read off the corresponding JS regularizer,
In order to obtain the regularizer in the logit-parameterization, with , we make use of the following property of the sigmoid
Let us finally also show that these two parametrizations are indeed equivalent at the optimum:
2 Further Considerations Regarding the Regularizer
We can also justify the approximation in Eq. (18) more rigorously. Starting from the Taylor approximation of at , we get pointwise
3 2D Submanifold Mixture of Gaussians in 3D Space
The experimental setup is inspired by the two-dimensional mixture of Gaussians in . The dataset is constructed as follows. We sample from a mixture of seven Gaussians with standard deviation and means equally spaced around the unit circle. This 2D mixture is then embedded in 3D space , rotated by around the axis and translated by . As illustrated in Fig. 5, this yields a mixture distribution that lives in a tilted 2D submanifold embedded in 3D space. It is important to emphasize that the mixture distribution is by design degenerate with respect to the base measure in 3D as it does not have full dimensional support. This precisely represents the dimensional misspecification scenario for GANs that we aim to address with our regularizer.
Architecture.
The architecture corresponds to the one used in with one notable exception. We use 2 dimensional latent vectors , sampled from a multivariate normal prior, (whereas uses 256 dimensional ), as we found lower dimensional latent variables greatly improve the performance of the unregularized GAN against which we compare. We did all experiments also for latent vectors of dimension 15: the obtained results are in accordance with those presented in the main text and below.
Both networks are optimized using Adam with a learning rate of and standard hyper-parameters. We trained on batches of size 512. Ten batches were generated to produce one image of the mixture at given time steps. The generator and discriminator network parameters were updated alternatively.
4 Image Datasets and Network Architectures
We trained on CelebA , CIFAR-10 and LSUN . All datasets were trained on minibatches of size 64. The respective GAN architectures are listed below.
For the CIFAR-10 experiments, we used the DCGAN (Deep Convolutional GAN) architecture of implemented in Tensorflow by https://github.com/carpedm20/DCGAN-tensorflow.
For the LSUN experiments, we used the DCGAN architecture of implemented in Tensorflow by https://github.com/igul222/improved_wgan_training.
In both cases, the discriminator uses batch normalization in all but the first and last layer (except where explicitly stated otherwise). The generator uses batch normalization in all layers except the last one. The latent code is sampled from a 100-dimensional uniform distribution over $$ (carpedm20) resp. 128-dimensional normal-(0,1) distribution (Gulrajani).
Both networks were trained using the Adam optimizer (with hyper-parameters recommended by the DCGAN authors) for various different learning rates in the range . The recommended learning rate was found to perform best.
ResNet GAN.
For the CelebA and LSUN experiments, we used a deep residual network ResNet implemented in Tensorflow by https://github.com/igul222/improved_wgan_training (implementation details can be found in ).
The ResNets use pre-activation residual blocks with two 3 x 3 convolutional layers each and ReLU nonlinearities. The generator has one linear layer, followed by four residual blocks, one deconvolutional layer and tanh output activations. The discriminator has one convolutional layer, followed by four residual blocks and a linear output layer (the discriminator logits are then fed into the sigmoid GAN loss).
We use batch normalization in the generator and discriminator (except explicitly stated otherwise). Both networks were optimized using Adam with learning rate and standard hyperparameters. For further architectural details, please refer to and the excellent open-source implementation referenced above.