From optimal transport to generative modeling: the VEGAN cookbook
Olivier Bousquet, Sylvain Gelly, Ilya Tolstikhin, Carl-Johann Simon-Gabriel, Bernhard Schoelkopf
Introduction
The field of representation learning was initially driven by supervised approaches, with impressive results using large labelled datasets. Unsupervised generative modeling, in contrast, used to be a domain governed by probabilistic approaches focusing on low-dimensional data. Recent years have seen a convergence of those two approaches. In the new field that formed at the intersection, variational autoencoders (VAEs) form one well-established approach, theoretically elegant yet with the drawback that they tend to generate blurry images. In contrast, generative adversarial networks (GANs) turned out to be more impressive in terms of the visual quality of images sampled from the model, but have been reported harder to train. There has been a flurry of activity in assaying numerous configurations of GANs as well as combinations of VAEs and GANs. A unifying theory relating GANs to VAEs in a principled way is yet to be discovered. This forms a major motivation for the studies underlying the present paper.
Following , we approach generative modeling from the optimal transport point of view. The optimal transport (OT) cost is a way to measure a distance between probability distributions and provides a much weaker topology than many others, including -divergences associated with the original GAN algorithms. This is particularly important in applications, where data is usually supported on low dimensional manifolds in the input space . As a result, stronger notions of distances (such as -divergences, which capture the density ratio between distributions) often max out, providing no useful gradients for training. In contrast, the optimal transport behave nicer and may lead to a more stable training .
In this work we aim at minimizing the optimal transport cost between the true (but unknown) data distribution and a latent variable model . We do so via the primal form and investigate theoretical properties of this optimization problem. Our main contributions are listed below; cf. also Figure 1.
We derive an equivalent formulation for the primal form of , which makes the role of latent space and probabilistic encoders explicit (Theorem 1).
Unlike in VAE, we arrive at an optimization problem where the are constrained. We relax the constraints by penalization, arriving at the penalized optimal transport (POT) objective (13) which can be minimized with stochastic gradient descent by sampling from and .
We show that for squared Euclidean cost (Section 4.1), POT coincides with the objective of adversarial auto-encoders (AAE) . We believe this provides the first theoretical justification for AAE, showing that they approximately minimize the 2-Wasserstein distance . We also compare POT to VAE, adversarial variational Bayes (AVB) , and other methods based on the marginal log-likelihood. In particular, we show that all these methods necessarily suffer from blurry outputs, unlike POT or AAE.
When is the Euclidean distance (Section 4.2), POT and WGAN both minimize the 1-Wasserstein distance . They approach this problem from primal/dual forms respectively, which leads to a different behaviour of the resulting algorithms.
Finally, following a somewhat well-established tradition of using acronyms based on the “GAN” suffix, and because we give a recipe for blending VAE and GAN (using POT), we propose to call this work a VEGAN cookbook.
The authors of address computing the OT cost in large scale using stochastic gradient descent (SGD) and sampling. They approach this task either through the dual formulation, or via a regularized version of the primal. They do not discuss any implications for generative modeling. Our approach is based on the primal form of OT, we arrive at regularizers which are very different, and our main focus is on generative modeling. The Wasserstein GAN minimizes the 1-Wasserstein distance for generative modeling. The authors approach this task from the dual form. Unfortunately, their algorithm cannot be readily applied to any other OT cost, because the famous Kantorovich duality holds only for . In contrast, our algorithm POT approaches the same problem from the primal form and can be applied for any cost function .
The present paper is structured as follows. In Section 2 we introduce our notation and discuss existing generative modeling techniques, including GANs, VAEs, and AAE. Section 3 contains our main results, including a theoretical analysis of the primal form of OT and the novel POT objective. Section 4 discusses the implications of our new results. Finally, we discuss future work in Section 5. Proofs may be found in Appendix B.
Notations and preliminaries
We use calligraphic letters (i.e. ) for sets, capital letters (i.e. ) for random variables, and lower case letters (i.e. ) for their values. We denote probability distributions with capital letters (i.e. ) and densities with lower case letters (i.e. ). By we denote the Dirac distribution putting mass 1 on , and denotes the support of .
where is any measurable cost function and is a set of all joint distributions of with marginals and respectively. A particularly interesting case is when is a metric space and for . In this case , the -th root of , is called the -Wasserstein distance. Finally, the Kantorovich-Rubinstein theorem establishes a duality for the 1-Wasserstein distance, which holds under mild assumptions on and :
where is the class of all bounded 1-Lipschitz functions on . Note that the same symbol is used for and , but only is a number and thus the above refers to the Wasserstein distance.
Even though GANs and VAEs are quite different—both in terms of the conceptual frameworks and empirical performance—they share important features: (a) both can be trained by sampling from the model without knowing an analytical form of its density and (b) both can be scaled up with SGD. As a result, it becomes possible to use highly flexible implicit models defined by a two-step procedure, where first a code is sampled from a fixed distribution on a latent space and then is mapped to the image with a (possibly random) transformation . This results in latent variable models defined on with density of the form
assuming all involved densities are properly defined. These models are indeed easy to sample and, provided can be differentiated analytically with respect to its parameters, can be trained with SGD. The field is growing rapidly and numerous variations of VAEs and GANs are available in the literature. Next we introduce and compare several of them.
The original generative adversarial network (GAN) approach minimizes
Variational auto-encoders (VAE) utilize models of the form (3) and minimize
Minimizing the primal of optimal transport
We have argued that minimizing the optimal transport cost between the true data distribution and the model is a reasonable goal for generative modeling. We now will explain how this can be done in the primal formulation of the OT problem (1) by reparametrizing the space of couplings (Section 3.1) and relaxing the marginal constraint (Section 3.2), leading to a formulation involving expectations over and that can thus be solved using SGD and sampling.
We will consider certain sets of joint probability distributions of three random variables . The reader may wish to think of as true images, as images sampled from the model, and as latent codes. We denote by a joint distribution of a variable pair , where is first sampled from and next from . Note that defined in (3) and used throughout this work is the marginal distribution of when .
In the optimal transport problem, we consider joint distributions which are called couplings between values of and . Because of the marginal constraint, we can write and we can consider as a non-deterministic mapping from to . In this section we will show how to factor this mapping through , i.e., decompose it into an encoding distribution and the generating distribution .
In order to give a more intuitive explanation of this decomposition, consider the case where all probability distributions have densities with respect to the Lebesgue measure. In this case our results show that some elements of have densities of the form
where the density of the conditional distribution satisfies for all . Equation (8) allows to express the search space (couplings) of the OT problem in terms of the probabilistic encoders . Unlike VAE, AVB, and other methods based on the marginal log-likelihood, where the are not constrained, the encoders of the OT problem need to match the aggregated posterior to the prior .
As in Section 2, denotes the set of all joint distributions of with marginals , and likewise for . The set of all joint distributions of such that , , and (Y\mathchoice{\mathrel{\hbox to0.0pt{\displaystyle\perp\hss}\mkern 2.0mu{\displaystyle\perp}}}{\mathrel{\hbox to0.0pt{\textstyle\perp\hss}\mkern 2.0mu{\textstyle\perp}}}{\mathrel{\hbox to0.0pt{\scriptstyle\perp\hss}\mkern 2.0mu{\scriptstyle\perp}}}{\mathrel{\hbox to0.0pt{\scriptscriptstyle\perp\hss}\mkern 2.0mu{\scriptscriptstyle\perp}}}X)|Z will be denoted by . Finally, we denote by and the sets of marginals on and (respectively) induced by distributions in . Note that , , and depend on the choice of conditional distributions , while does not. In fact, it is easy to check that . From the definitions it is clear that and we get the following upper bound:
If are Dirac measures (i.e., ), the two sets are actually coincide, thus justifying the reparametrization (8) and the illustration in Figure 1(b), as demonstrated in the following theorem:
If for all , where , we have
where is the marginal distribution of when and .
The r.h.s. of (10) is the optimal transport between and with the cost function c_{g}(x,z):=c\bigl{(}x,G(z)\bigr{)} defined on . If corresponds to a deterministic mapping , the resulting model is the push-forward of through . Theorem 1 states that in this case the two OT problems are equivalent, .
The conditional distributions in (11) are constrained to ensure that when is sampled from and then from , the resulting marginal distribution of (which was denoted and called aggregated posterior in Section 2.1) coincides with . A very similar result also holds for the case when are not necessarily Dirac. Nevertheless, for simplicity we will focus on the Dirac case for now and shortly summarize the more general case in the end of this section.
2 Relaxing the constraints
Minimizing boils down to a min-min optimization problem. Unfortunately, the constraint on makes the variational problem even harder. We thus propose to replace the constrained optimization problem (11) with its relaxed unconstrained version in a standard way. Namely, use any convex penalty , such that if and only if , and for any , consider the following relaxed unconstrained version of :
It is well known that under mild conditions adding a penalty as in (12) is equivalent to adding a constraint of the form for some . As increases, the corresponding decreases, and as , the solutions of (12) reach the feasible region where . This shows that for all and the gap reduces with increasing .
The above discussion assumed Dirac measures . If this is not the case, we can still upper bound the 2-Wasserstein distance , corresponding to , in a very similar way to Theorem 1. The case of Gaussian decoders , which will be particularly useful when discussing the relation to VAEs, is summarized in the following remark:
Implications: relations to AAE, VAEs, and GANs
Thus far we showed that the primal form of the OT problem can be relaxed to allow for efficient minimization by SGD. We now discuss the main implications of these results.
Consider the squared Euclidean cost function , for which is the squared 2-Wasserstein distance . The goal of this section is to compare the minimization of to other generative modeling approaches. Let us focus our attention on generative distributions typically used in VAE, AVB, and AAE, i.e., Gaussians P_{G}(Y|Z)=\mathcal{N}\bigl{(}Y;G(Z),\sigma^{2}\!\cdot\!I_{d}\bigr{)}. In order to verify the differentiability of all three methods require and have problems handling the case of deterministic decoders (). To emphasize the role of the variance we will denote the resulting latent variable model .
The analysis of Section 3 shows that the value of is upper bounded by of the form (16) and the two coincide when . Next we summarize properties of solutions minimizing these two values and :
Let and assume , P_{G}(Y|Z)=\mathcal{N}\bigl{(}Y;G(Z),\sigma^{2}\!\cdot\!I\bigr{)} with any function . If then the functions and minimizing and respectively are different: depends on , while does not. The function is also a minimizer of .
In contrast, Proposition 2 shows that for the same Gaussian models with any given we can minimize and the solution will be indeed the one resulting in the smallest 2-Wasserstein distance between and the noiseless implicit model , used in practice.
Relation to AAE
The authors of tried to establish a connection between AAE and log-likelihood maximization. They argued that AAE is “a crude approximation” to AVB. Our results suggest that AAE is in fact attempting to minimize the 2-Wasserstein distance between and , which may explain its good empirical performance reported in .
Blurriness of VAE and AVB
We next add to the discussion regarding the blurriness commonly attributed to VAE samples. Our argument shows that VAE, AVB, and other methods based on the marginal log-likelihood necessarily lead to an averaging in the input space if are Gaussian.
This overlap necessarily happens in VAEs, which use Gaussian encoders supported on the entire . When probabilistic encoders are allowed to be flexible enough, as in AVB, for any fixed the optimal will try to invert the decoder (see Appendix A) and take the form
This approximation becomes exact in the nonparametric limit of . When is Gaussian we have for all and , showing that for all . This will again lead to the overlap of encoders if . In contrast, the optimal encoders of AAE and POT do not necessarily overlap, as they are not inverting the decoders.
2 The 1-Wasserstein distance: relation to WGAN
This means we can now approach the problem of optimizing in two distinct ways, taking gradient steps either in the primal or in the dual forms. Denote by the optimal encoder in the primal and the optimal witness function in the dual. By the envelope theorem, gradients of with respect to can be computed by taking a gradient of the criteria evaluated at the optimal points or .
Despite the theoretical equivalence of both approaches, practical considerations lead to different behaviours and to potentially poor approximations of the real gradients. For example, in the dual formulation, one usually restricts the witness functions to be smooth, while in the primal formulation, the constraint on is only approximately enforced. We will study the effect of these approximations.
There exists a constant such that for any , one can construct distributions , and pick witness functions and that are -optimal , , but which give (at some point ) gradients whose direction is at least -wrong: , , and .
Imperfect posterior in the primal (i.e., for POT)
In the primal formulation, when the constraint is violated, that is the aggregated posterior is not matching , there can be two kinds of negative effects: (i) the gradient of the criterion is only computed on a (possibly small) subset of the latent space reached by ; (ii) several input points could be mapped by to the same latent code , thus giving gradients that encourage to be the average/median of several inputs (hence encouraging a blurriness). See Section 4.1 for the details and Figure 1 for an illustration.
Conclusion
This work proposes a way to fit generative models by minimizing any optimal transport cost. It also establishes novel links between different popular unsupervised probabilistic modeling techniques. Whilst our contribution is on the theoretical side, it is reassuring to note that the empirical results of show the strong performance of our method for the special case of the 2-Wasserstein distance. Experiments with other cost functions are beyond the scope of the present work and left for the future studies.
Acknowledgments
The authors are thankful to Mateo Rojas-Carulla and Fei Sha for stimulating discussions. CJSG is supported by a Google European Doctoral Fellowship in Causal Inference.
References
Appendix A Further details on VAEs and GANs
For models of the form (3) and any conditional distribution it can be easily verified that
Here the conditional distribution is induced by a joint distribution , which is in turn specified by the 2-step latent variable procedure: (a) sample from , (b) sample from . Note that the first term on the r.h.s. of (15) is always non-positive, while the l.h.s. does not depend on . This shows that if conditional distributions are not restricted then
where the infimum is achieved for . However, for any restricted class of conditional distributions we only have
where the inequality accounts for the fact that might be not flexible enough to match for all values of .
Relation between AAE, AVB, and VAE
For any distributions and :
Appendix B Proofs
We start by introducing an important lemma relating the two sets over which and are optimized.
with identity if are Dirac distributions for all We conjecture that this is also a necessary condition. The necessity is not used in the remainder of the paper..
Inequality (9) and the first identity in (10) obviously follows from Lemma 6. The tower rule of expectation, and the conditional independence property of implies
Let and assume the conditional distributions have mean values and marginal variances for all , where . Take . Then
Proof is similar to the one of Theorem 1. See Section B.2. ∎
First inequality follows from (9). For the identity we proceed similarly to the proof of Theorem 1 and write
Together with (17) and the fact that this concludes the proof.
B.3 Proof of Proposition 2
Corollary 7 shows that does not depend on the variance . When the distribution turns into Dirac. In this case we combine Theorem 1 and Corollary 7 to conclude that also minimizes . Next we prove that generally depends on .
We will need the following simple result, which is basically saying that the variance of a sum of two independent random variables is a sum of the variances:
Under conditions of Proposition 2, assume . Then
Next we prove the remaining implication of Proposition 2. Namely, that when the function minimizing depends on . The proof is based on the following example: , , , and . Note that by setting for any we ensure that is the Gaussian distribution, because a convolution of two Gaussians is also Gaussian. In particular if we take Lemma 8 implies that is the standard normal Gaussian . In other words, we obtain the global minimum and clearly depends on .
B.4 Proof of Proposition 3
We just give a sketch of the proof: consider discrete distributions supported on two points , and supported on and let , (). Given an optimal , one can modify locally it around without changing its Lipschitz constant such that the obtained is an -approximation of whose gradients at and point in directions arbitrarily different from those of . For smooth functions, by moving and away from the segment but close to each other , the gradients of will point in directions roughly opposite but the constraint on the Hessian will force the gradients of at and to be very close. Finally, putting on the segment , one can get an whose gradients at and are exactly opposite, while taking , we can swap the direction at one of the points while changing the criterion by less than .