Wasserstein GAN

Martin Arjovsky, Soumith Chintala, Léon Bottou

Introduction

For this to make sense, we need the model density PθP_{\theta} to exist. This is not the case in the rather common situation where we are dealing with distributions supported by low dimensional manifolds. It is then unlikely that the model manifold and the true distribution’s support have a non-negligible intersection (see ), and this means that the KL distance is not defined (or simply infinite).

The typical remedy is to add a noise term to the model distribution. This is why virtually all generative models described in the classical machine learning literature include a noise component. In the simplest case, one assumes a Gaussian noise with relatively high bandwidth in order to cover all the examples. It is well known, for instance, that in the case of image generation models, this noise degrades the quality of the samples and makes them blurry. For example, we can see in the recent paper that the optimal standard deviation of the noise added to the model when maximizing likelihood is around 0.1 to each pixel in a generated image, when the pixels were already normalized to be in the range $$. This is a very high amount of noise, so much that when papers report the samples of their models, they don’t add the noise term on which they report likelihood numbers. In other words, the added noise term is clearly incorrect for the problem, but is needed to make the maximum likelihood approach work.

Variational Auto-Encoders (VAEs) and Generative Adversarial Networks (GANs) are well known examples of this approach. Because VAEs focus on the approximate likelihood of the examples, they share the limitation of the standard models and need to fiddle with additional noise terms. GANs offer much more flexibility in the definition of the objective function, including Jensen-Shannon , and all ff-divergences as well as some exotic combinations . On the other hand, training GANs is well known for being delicate and unstable, for reasons theoretically investigated in .

In Section 2, we provide a comprehensive theoretical analysis of how the Earth Mover (EM) distance behaves in comparison to popular probability distances and divergences used in the context of learning distributions.

In Section 3, we define a form of GAN called Wasserstein-GAN that minimizes a reasonable and efficient approximation of the EM distance, and we theoretically show that the corresponding optimization problem is sound.

In Section 4, we empirically show that WGANs cure the main training problems of GANs. In particular, training WGANs does not require maintaining a careful balance in training of the discriminator and the generator, and does not require a careful design of the network architecture either. The mode dropping phenomenon that is typical in GANs is also drastically reduced. One of the most compelling practical benefits of WGANs is the ability to continuously estimate the EM distance by training the discriminator to optimality. Plotting these learning curves is not only useful for debugging and hyperparameter searches, but also correlate remarkably well with the observed sample quality.

Different Distances

The Earth-Mover (EM) distance or Wasserstein-1

The following example illustrates how apparently simple sequences of probability distributions converge under the EM distance but do not converge under the other distances and divergences defined above.

Example 1 gives us a case where we can learn a probability distribution over a low dimensional manifold by doing gradient descent on the EM distance. This cannot be done with the other distances and divergences because the resulting loss function is not even continuous. Although this simple example features distributions with disjoint supports, the same conclusion holds when the supports have a non empty intersection contained in a set of measure zero. This happens to be the case when two low dimensional manifolds intersect in general position .

The following corollary tells us that learning by minimizing the EM distance makes sense (at least in theory) with neural networks.

All this shows that EM is a much more sensible cost function for our problem than at least the Jensen-Shannon divergence. The following theorem describes the relative strength of the topologies induced by these distances and divergences, with KL the strongest, followed by JS and TV, and EM the weakest.

The statements in (1) imply the statements in (2).

This highlights the fact that the KL, JS, and TV distances are not sensible cost functions when learning distributions supported by low dimensional manifolds. However the EM distance is sensible in that setup. This obviously leads us to the next section where we introduce a practical approximation of optimizing the EM distance.

Wasserstein GAN

Weight clipping is a clearly terrible way to enforce a Lipschitz constraint. If the clipping parameter is large, then it can take a long time for any weights to reach their limit, thereby making it harder to train the critic till optimality. If the clipping is small, this can easily lead to vanishing gradients when the number of layers is big, or batch normalization is not used (such as in RNNs). We experimented with simple variants (such as projecting the weights to a sphere) with little difference, and we stuck with weight clipping due to its simplicity and already good performance. However, we do leave the topic of enforcing Lipschitz constraints in a neural network setting for further investigation, and we actively encourage interested researchers to improve on this method.

The fact that the EM distance is continuous and differentiable a.e. means that we can (and should) train the critic till optimality. The argument is simple, the more we train the critic, the more reliable gradient of the Wasserstein we get, which is actually useful by the fact that Wasserstein is differentiable almost everywhere. For the JS, as the discriminator gets better the gradients get more reliable but the true gradient is 0 since the JS is locally saturated and we get vanishing gradients, as can be seen in Figure 1 of this paper and Theorem 2.4 of . In Figure 2 we show a proof of concept of this, where we train a GAN discriminator and a WGAN critic till optimality. The discriminator learns very quickly to distinguish between fake and real, and as expected provides no reliable gradient information. The critic, however, can’t saturate, and converges to a linear function that gives remarkably clean gradients everywhere. The fact that we constrain the weights limits the possible growth of the function to be at most linear in different parts of the space, forcing the optimal critic to have this behaviour.

Perhaps more importantly, the fact that we can train the critic till optimality makes it impossible to collapse modes when we do. This is due to the fact that mode collapse comes from the fact that the optimal generator for a fixed discriminator is a sum of deltas on the points the discriminator assigns the highest values, as observed by and highlighted in .

In the following section we display the practical benefits of our new algorithm, and we provide an in-depth comparison of its behaviour and that of traditional GANs.

Empirical Results

We run experiments on image generation using our Wasserstein-GAN algorithm and show that there are significant practical benefits to using it over the formulation used in standard GANs.

a meaningful loss metric that correlates with the generator’s convergence and sample quality

improved stability of the optimization process

We run experiments on image generation. The target distribution to learn is the LSUN-Bedrooms dataset – a collection of natural images of indoor bedrooms. Our baseline comparison is DCGAN , a GAN with a convolutional architecture trained with the standard GAN procedure using the −log⁡D-\log D trick . The generated samples are 3-channel images of 64x64 pixels in size. We use the hyper-parameters specified in Algorithm 1 for all of our experiments.

2 Meaningful loss metric

Because the WGAN algorithm attempts to train the critic ff (lines 2–8 in Algorithm 1) relatively well before each generator update (line 10 in Algorithm 1), the loss function at this point is an estimate of the EM distance, up to constant factors related to the way we constrain the Lipschitz constant of ff.

Our first experiment illustrates how this estimate correlates well with the quality of the generated samples. Besides the convolutional DCGAN architecture, we also ran experiments where we replace the generator or both the generator and the critic by 4-layer ReLU-MLP with 512 hidden units.

Figure 3 plots the evolution of the WGAN estimate (3) of the EM distance during WGAN training for all three architectures. The plots clearly show that these curves correlate well with the visual quality of the generated samples.

To our knowledge, this is the first time in GAN literature that such a property is shown, where the loss of the GAN shows properties of convergence. This property is extremely useful when doing research in adversarial networks as one does not need to stare at the generated samples to figure out failure modes and to gain information on which models are doing better over others.

However, we do not claim that this is a new method to quantitatively evaluate generative models yet. The constant scaling factor that depends on the critic’s architecture means it’s hard to compare models with different critics. Even more, in practice the fact that the critic doesn’t have infinite capacity makes it hard to know just how close to the EM distance our estimate really is. This being said, we have succesfully used the loss metric to validate our experiments repeatedly and without failure, and we see this as a huge improvement in training GANs which previously had no such facility.

In contrast, Figure 4 plots the evolution of the GAN estimate of the JS distance during GAN training. More precisely, during GAN training, the discriminator is trained to maximize

This quantity clearly correlates poorly the sample quality. Note also that the JS estimate usually stays constant or goes up instead of going down. In fact it often remains very close to log⁡2≈0.69\log 2\approx 0.69 which is the highest value taken by the JS distance. In other words, the JS distance saturates, the discriminator has zero loss, and the generated samples are in some cases meaningful (DCGAN generator, top right plot) and in other cases collapse to a single nonsensical image . This last phenomenon has been theoretically explained in and highlighted in .

When using the −log⁡D-\log D trick , the discriminator loss and the generator loss are different. Figure 8 in Appendix E reports the same plots for GAN training, but using the generator loss instead of the discriminator loss. This does not change the conclusions.

Finally, as a negative result, we report that WGAN training becomes unstable at times when one uses a momentum based optimizer such as Adam (with β1>0\beta_{1}>0) on the critic, or when one uses high learning rates. Since the loss for the critic is nonstationary, momentum based methods seemed to perform worse. We identified momentum as a potential cause because, as the loss blew up and samples got worse, the cosine between the Adam step and the gradient usually turned negative. The only places where this cosine was negative was in these situations of instability. We therefore switched to RMSProp which is known to perform well even on very nonstationary problems .

3 Improved stability

One of the benefits of WGAN is that it allows us to train the critic till optimality. When the critic is trained to completion, it simply provides a loss to the generator that we can train as any other neural network. This tells us that we no longer need to balance generator and discriminator’s capacity properly. The better the critic, the higher quality the gradients we use to train the generator.

We observe that WGANs are much more robust than GANs when one varies the architectural choices for the generator. We illustrate this by running experiments on three generator architectures: (1) a convolutional DCGAN generator, (2) a convolutional DCGAN generator without batch normalization and with a constant number of filters, and (3) a 4-layer ReLU-MLP with 512 hidden units. The last two are known to perform very poorly with GANs. We keep the convolutional DCGAN architecture for the WGAN critic or the GAN discriminator.

Figures 7, 7, and 7 show samples generated for these three architectures using both the WGAN and GAN algorithms. We refer the reader to Appendix Appendix F for full sheets of generated samples. Samples were not cherry-picked.

In no experiment did we see evidence of mode collapse for the WGAN algorithm.

Related Work

as an integral probability metric associated with the function class F\mathcal{F}. It is easily verified that if for every f∈Ff\in\mathcal{F} we have −f∈F-f\in\mathcal{F} (such as all examples we’ll consider), then dFd_{\mathcal{F}} is nonnegative, satisfies the triangular inequality, and is symmetric. Thus, dFd_{\mathcal{F}} is a pseudometric over Prob(X)\text{Prob}(\mathcal{X}).

While IPMs might seem to share a similar formula, as we will see different classes of functions can yeald to radically different metrics.

Since the total variation distance displays the same regularity as the JS, it can be seen that EBGANs will suffer from the same problems of classical GANs regarding not being able to train the discriminator till optimality and thus limiting itself to very imperfect gradients.

The great aspect of MMD is that via the kernel trick there is no need to train a separate network to maximize equation (4) for the ball of a RKHS. However, this has the disadvantage that evaluating the MMD distance has computational cost that grows quadratically with the amount of samples used to estimate the expectations in (4). This last point makes MMD have limited scalability, and is sometimes inapplicable to many real life applications because of it. There are estimates with linear computational cost for the MMD which in a lot of cases makes MMD very useful, but they also have worse sample complexity.

That being said, these numbers can be a bit unfair to the MMD, in the sense that we are comparing empirical sample complexity of GANs with the theoretical sample complexity of MMDs, which tends to be worse. However, in the original GMMN paper they indeed used a minibatch of size 1000, much larger than the standard 32 or 64 (even when this incurred in quadratic computational cost). While estimates that have linear computational cost as a function of the number of samples exist , they have worse sample complexity, and to the best of our knowledge they haven’t been yet applied in a generative context such as in GMMNs.

On another great line of research, the recent work of has explored the use of Wasserstein distances in the context of learning for Restricted Boltzmann Machines for discrete spaces. The motivations at a first glance might seem quite different, since the manifold setting is restricted to continuous spaces and in finite discrete spaces the weak and strong topologies (the ones of W and JS respectively) coincide. However, in the end there is more in commmon than not about our motivations. We both want to compare distributions in a way that leverages the geometry of the underlying space, and Wasserstein allows us to do exactly that.

Finally, the work of shows new algorithms for calculating Wasserstein distances between different distributions. We believe this direction is quite important, and perhaps could lead to new ways of evaluating generative models.

Conclusion

We introduced an algorithm that we deemed WGAN, an alternative to traditional GAN training. In this new model, we showed that we can improve the stability of learning, get rid of problems like mode collapse, and provide meaningful learning curves useful for debugging and hyperparameter searches. Furthermore, we showed that the corresponding optimization problem is sound, and provided extensive theoretical work highlighting the deep connections to other distances between distributions.

Acknowledgments

We would like to thank Mohamed Ishmael Belghazi, Emily Denton, Ian Goodfellow, Ishaan Gulrajani, Alex Lamb, David Lopez-Paz, Eric Martin, Maxime Oquab, Aditya Ramesh, Ronan Riochet, Uri Shalit, Pablo Sprechmann, Arthur Szlam, Ruohan Wang, for helpful comments and advice.

References

Appendix A Why Wasserstein is indeed weak

Note that if f∈Cb(X)f\in C_{b}(\mathcal{X}), we can define ∥f∥∞=max⁡x∈X∣f(x)∣\|f\|_{\infty}=\max_{x\in\mathcal{X}}|f(x)|, since ff is bounded. With this norm, the space (Cb(X),∥⋅∥∞)(C_{b}(\mathcal{X}),\|\cdot\|_{\infty}) is a normed vector space. As for any normed vector space, we can define its dual

and give it the dual norm ∥ϕ∥=sup⁡f∈Cb(X),∥f∥∞≤1∣ϕ(f)∣\|\phi\|=\sup_{f\in C_{b}(\mathcal{X}),\|f\|_{\infty}\leq 1}|\phi(f)|.

With this definitions, (Cb(X)∗,∥⋅∥)(C_{b}(\mathcal{X})^{*},\|\cdot\|) is another normed space. Now let μ\mu be a signed measure over X\mathcal{X}, and let us define the total variation distance

is a distance in Prob(X)\text{Prob}(\mathcal{X}) (called the total variation distance).

Now, all dual spaces (such as Cb(X)∗C_{b}(\mathcal{X})^{*} and thus Prob(X)\text{Prob}(\mathcal{X})) have a strong topology (induced by the norm), and a weak* topology. As the name suggests, the weak* topology is much weaker than the strong topology. In the case of Prob(X)\text{Prob}(\mathcal{X}), the strong topology is given by the total variation distance, and the weak* topology is given by the Wasserstein distance (among others) .

Appendix B Assumption definitions

Appendix C Proofs of things

By the definition of the Wasserstein distance, we have

If gg is continuous in θ\theta, then gθ(z)→θ→θ′gθ′(z)g_{\theta}(z)\to_{\theta\to\theta^{\prime}}g_{\theta^{\prime}}(z), so ∥gθ−gθ′∥→0\|g_{\theta}-g_{\theta^{\prime}}\|\to 0 pointwise as functions of zz. Since X\mathcal{X} is compact, the distance of any two elements in it has to be uniformly bounded by some constant MM, and therefore ∥gθ(z)−gθ′(z)∥≤M\|g_{\theta}(z)-g_{\theta^{\prime}}(z)\|\leq M for all θ\theta and zz uniformly. By the bounded convergence theorem, we therefore have

Now let gg be locally Lipschitz. Then, for a given pair (θ,z)(\theta,z) there is a constant L(θ,z)L(\theta,z) and an open set UU such that (θ,z)∈U(\theta,z)\in U, such that for every (θ′,z′)∈U(\theta^{\prime},z^{\prime})\in U we have

By taking expectations and z′=zz^{\prime}=z we

The counterexample for item 3 of the Theorem is indeed Example 1. ∎

We begin with the case of smooth nonlinearities. Since gg is C1C^{1} as a function of (θ,z)(\theta,z) then for any fixed (θ,z)(\theta,z) we have L(θ,Z)≤∥∇θ,xgθ(z)∥+ϵL(\theta,Z)\leq\|\nabla_{\theta,x}g_{\theta}(z)\|+\epsilon is an acceptable local Lipschitz constant for all ϵ>0\epsilon>0. Therefore, it suffices to prove

If HH is the number of layers we know that ∇zgθ(z)=∏k=1HWkDk\nabla_{z}g_{\theta}(z)=\prod_{k=1}^{H}W_{k}D_{k} where WkW_{k} are the weight matrices and DkD_{k} is are the diagonal Jacobians of the nonlinearities. Let fi:jf_{i:j} be the application of layers ii to jj inclusively (e.g. gθ=f1:Hg_{\theta}=f_{1:H}). Then, ∇Wkgθ(z)=((∏i=k+1HWiDi)Dk)f1:k−1(z)\nabla_{W_{k}}g_{\theta}(z)=\left(\left(\prod_{i=k+1}^{H}W_{i}D_{i}\right)D_{k}\right)f_{1:k-1}(z). We recall that if LL is the Lipschitz constant of the nonlinearity, then ∥Di∥≤L\|D_{i}\|\leq L and ∥f1:k−1(z)∥≤∥z∥Lk−1∏i=1k−1Wi\|f_{1:k-1}(z)\|\leq\|z\|L^{k-1}\prod_{i=1}^{k-1}W_{i}. Putting this together,

If C1(θ)=LH(∏i=1H∥Wi∥)C_{1}(\theta)=L^{H}\left(\prod_{i=1}^{H}\|W_{i}\|\right) and C2(θ)=∑k=1HLH(∏i=1k−1∥Wi∥)(∏i=k+1H∥Wi∥)C_{2}(\theta)=\sum_{k=1}^{H}L^{H}\left(\prod_{i=1}^{k-1}\|W_{i}\|\right)\left(\prod_{i=k+1}^{H}\|W_{i}\|\right) then

Let ϵ>0\epsilon>0 fixed, and An={fn>1+ϵ}A_{n}=\{f_{n}>1+\epsilon\}. Then,

This is a long known fact that WW metrizes the weak* topology of (C(X),∥⋅∥∞)(C(\mathcal{X}),\|\cdot\|_{\infty}) on Prob(X)\text{Prob}(\mathcal{X}), and by definition this is the topology of convergence in distribution. A proof of this can be found (for example) in .

This is a straightforward application of Pinsker’s inequality

This is trivial by recalling the fact that δ\delta and WW give the strong and weak* topologies on the dual of (C(X),∥⋅∥∞)(C(\mathcal{X}),\|\cdot\|_{\infty}) when restricted to Prob(X)\text{Prob}(\mathcal{X}).

Since X\mathcal{X} is compact, we know by the Kantorovich-Rubenstein duality that there is an f∈Ff\in\mathcal{F} that attains the value

for any f∈X∗(θ)f\in X^{*}(\theta) when both terms are well-defined.

Let f∈X∗(θ)f\in X^{*}(\theta), which we knows exists since X∗(θ)X^{*}(\theta) is non-empty for all θ\theta. Then, we get

under the condition that the first and last terms are well-defined. The rest of the proof will be dedicated to show that

when the right hand side is defined. For the reader who is not interested in such technicalities, he or she can skip the rest of the proof.

Since f∈Ff\in\mathcal{F}, we know that it is 1-Lipschitz. Furthermore, gθ(z)g_{\theta}(z) is locally Lipschitz as a function of (θ,z)(\theta,z). Therefore, f(gθ(z))f(g_{\theta}(z)) is locally Lipschitz on (θ,z)(\theta,z) with constants L(θ,z)L(\theta,z) (the same ones as gg). By Radamacher’s Theorem, f(gθ(z))f(g_{\theta}(z)) has to be differentiable almost everywhere for (θ,z)(\theta,z) jointly. Rewriting this, the set A=\{(\theta,z):\text{f\circ gis not differentiable}\} has measure 0. By Fubini’s Theorem, this implies that for almost every θ\theta the section Aθ={z:(θ,z)∈A}A_{\theta}=\{z:(\theta,z)\in A\} has measure 0. Let’s now fix a θ0\theta_{0} such that the measure of Aθ0A_{\theta_{0}} is null (such as when the right hand side of equation (5) is well defined). For this θ0\theta_{0} we have ∇θf(gθ(z))∣θ0\nabla_{\theta}f(g_{\theta}(z))|_{\theta_{0}} is well-defined for almost any zz, and since p(z)p(z) has a density, it is defined p(z)p(z)-a.e. By assumption 1 we know that

By differentiability, the term inside the integral converges p(z)p(z)-a.e. to 0 as θ→θ0\theta\to\theta_{0}. Furthermore,

Appendix D Energy-based GANs optimize total variation

In this appendix we show that under an optimal discriminator, energy-based GANs (EBGANs) optimize the total variation distance between the real and generated distributions.

Energy-based GANs are trained in a similar fashion to GANs, only under a different loss function. They have a discriminator DD who tries to minimize

for some m>0m>0 and [x]+=max⁡(0,x)[x]^{+}=\max(0,x) and a generator network gθg_{\theta} that’s trained to minimize

First, we prove that there exists an optimal discriminator. Let D:X→[0,+∞)D:\mathcal{X}\rightarrow[0,+\infty) be a measurable function, then D′(x):=min⁡(D(x),m)D^{\prime}(x):=\min(D(x),m) is also a measurable function, and LD(D′,gθ)≤LD(D,gθ)L_{D}(D^{\prime},g_{\theta})\leq L_{D}(D,g_{\theta}). Therefore, a function D∗:X→[0,+∞)D^{*}:\mathcal{X}\rightarrow[0,+\infty) is optimal if and only if D∗′{D^{*}}^{\prime} is. Furthermore, it is optimal if and only if LD(D∗,gθ)≤LD(D,gθ)L_{D}(D^{*},g_{\theta})\leq L_{D}(D,g_{\theta}) for all D:X→[0,m]D:\mathcal{X}\rightarrow[0,m]. We are then interested to see if there’s an optimal discriminator for the problem min⁡0≤D(x)≤mLD(D,gθ)\min_{0\leq D(x)\leq m}L_{D}(D,g_{\theta}).

Note now that if 0≤D(x)≤m0\leq D(x)\leq m we have

Furthermore, if ff is bounded between -1 and 1, we get

Appendix E Generator’s cost during normal GAN training

Appendix F Sheets of samples