Metropolis-Hastings Generative Adversarial Networks

Ryan Turner, Jane Hung, Eric Frank, Yunus Saatci, Jason Yosinski

Introduction

Traditionally, density estimation is done with a model that can compute the data likelihood. Generative adversarial networks (GANs) (Goodfellow et al., 2014) present a radically new way to do density estimation: They implicitly represent the density of the data via a classifier that distinguishes real from generated data.

GANs iterate between updating a discriminator DD and a generator GG, where GG generates new (synthetic) samples of data, and DD attempts to distinguish samples of GG from the real data. In the typical setup, DD is thrown away at the end of training, and only GG is kept for generating new synthetic data points. In this work, we propose the Metropolis-Hastings GAN (MH-GAN), a GAN that constructs a new generator G′G^{\prime} that “wraps” GG using the information contained in DD. This principle is illustrated in Figure 1.Code found at: github.com/uber-research/metropolis-hastings-gans

The MH-GAN uses Markov chain Monte Carlo (MCMC) methods to sample from the distribution implicitly defined by the discriminator DD learned for the generator GG. This is built upon the notion that the discriminator classifies between the generator GG and a data distribution:

where pG{p_{G}} is the (intractable) density of samples from the generator GG, and pD{p_{D}} is the data density implied by the discriminator DD with respect to GG. If GAN training reaches its global optimum, then this discriminator distribution pD{p_{D}} is equal to the data distribution and the generator distribution (pD=pdata=pG{p_{D}}={p_{\textrm{data}}}={p_{G}}) (Goodfellow et al., 2014). Furthermore, if the discriminator DD is optimal for a fixed imperfect generator, GG then the implied distribution still equals the data distribution (pD=pdata≠pG{p_{D}}={p_{\textrm{data}}}\neq{p_{G}}).

We use an MCMC independence sampler (Tierney, 1994) to sample from pD{p_{D}} by taking multiple samples from GG. Amazingly, using our algorithm, one can show that given a perfect discriminator DD and a decent (but imperfect) generator GG, one can obtain exact samples from the true data distribution pdata{p_{\textrm{data}}}. Standard MCMC implementations require (unnormalized) densities for the target pD{p_{D}} and the proposal pG{p_{G}}, which are both unavailable for GANs. However, the Metropolis-Hastings (MH) algorithm requires only the ratio:

which we can obtain using only evaluation of D(x)D({\boldsymbol{\mathbf{x}}}).

Sampling from an MH-GAN is more computationally expensive than a standard GAN, but the bigger and more relevant training compute cost remains unchanged. Thus, the MH-GAN is best suited for applications where sample quality is more important than compute speed at test time.

The outline of this paper is as follows: Section 2 reviews diverse areas of relevant prior work. In Sections 3.1 and 3.2 we explain the necessary background on MCMC methods and GANs. We explain our methodology of combining these two seemingly disparate areas in Section 4 where we derive the wrapped generator G′G^{\prime}. Results on real data (CIFAR-10 and CelebA) and extending common GAN models (DCGAN, WGAN, and progressive GAN) are shown in Section 5. Section 6 discusses implications and conclusions.

Related Work

A few other works combine GANs and MCMC in some way. Song et al. (2017) use a GAN-like procedure to train a RealNVP (Dinh et al., 2016) MCMC proposal for sampling an externally provided target p⋆{p^{\star}}. Whereas Song et al. (2017) use GANs to accelerate MCMC, we use MCMC to enhance the samples from a GAN. Similar to Song et al. (2017), Kempinska & Shawe-Taylor (2017) improve proposals in particle filters rather than MCMC. Song et al. (2017) was recently generalized by Neklyudov et al. (2018).

A concurrent work with similar aims from Azadi et al. (2018) proposes discriminator rejection sampling (DRS) for GANs, which performs rejection sampling on the outputs of GG by using the probabilities given by DD. While conceptually appealing at first, DRS suffers from two major shortcomings in practice. First, it is necessary to find an upper-bound on DD over all possible samples in order to obtain a valid proposal distribution for rejection sampling. Because this is not possible, one must instead rely on estimating this bound by drawing many pilot samples. Secondly, even if one were to find a good bound, the acceptance rate would become very low due to the high-dimensionality of the sampling space. This leads Azadi et al. (2018) to use an extra γ\gamma heuristic to shift the logit DD scores, making the model sample from a distribution different from pdata{p_{\textrm{data}}} even when DD is perfect. We use MCMC instead, which was invented precisely as a replacement for rejection sampling in higher dimensions. We further improve the robustness of MCMC via use of a calibrator on the discriminator to get more accurate probabilities for computing acceptance.

Background and Notation

In this section, we briefly review the notation and equations with MCMC and GANs.

MCMC methods attempt to draw a chain of samples x1:K∈XK{\boldsymbol{\mathbf{x}}}_{1:K}\in\mathcal{X}^{K} that marginally come from a target distribution p⋆{p^{\star}}. We refer to the initial distribution as p0{p_{0}} and the proposal for the independence sampler as x′∼q(x′∣xk)=q(x′){\boldsymbol{\mathbf{x}}}^{\prime}\sim q({\boldsymbol{\mathbf{x}}}^{\prime}|{\boldsymbol{\mathbf{x}}}_{k})=q({\boldsymbol{\mathbf{x}}}^{\prime}). The proposal x′∈X{\boldsymbol{\mathbf{x}}}^{\prime}\in\mathcal{X} is accepted with probability

If x′{\boldsymbol{\mathbf{x}}}^{\prime} is accepted, xk+1=x′{\boldsymbol{\mathbf{x}}}_{k+1}={\boldsymbol{\mathbf{x}}}^{\prime}, otherwise xk+1=xk{\boldsymbol{\mathbf{x}}}_{k+1}={\boldsymbol{\mathbf{x}}}_{k}. Note that when estimating the distribution p⋆{p^{\star}}, one must include the duplicates that are a result of rejections in x′{\boldsymbol{\mathbf{x}}}^{\prime}.

Many evaluation metrics assume perfectly iid samples. Although MCMC methods are typically used to produce correlated samples, we can produce iid samples by using one chain per sample: Each chain samples x0∼p0{\boldsymbol{\mathbf{x}}}_{0}\sim{p_{0}} and then does KK MH iterations to get xK{\boldsymbol{\mathbf{x}}}_{K} as the output of the chain, which is the output of G′G^{\prime}. Using multiple chains is also better for GPU parallelization.

Detailed balance

The detailed balance condition implies that if xk∼p⋆{\boldsymbol{\mathbf{x}}}_{k}\sim{p^{\star}} exactly then xk+1∼p⋆{\boldsymbol{\mathbf{x}}}_{k+1}\sim{p^{\star}} exactly as well. Even if xk{\boldsymbol{\mathbf{x}}}_{k} is not exactly distributed according to p⋆{p^{\star}}, the Kullback-Leibler (KL) divergence between the implied density it is drawn from and p⋆{p^{\star}} always decreases as kk increases (Murray & Salakhutdinov, 2008). We use detailed balance to motivate our approach to MH-GAN initialization.

2 GANs

This implies a (intractable) distribution on the data x∼pG{\boldsymbol{\mathbf{x}}}\sim{p_{G}}. We refer to the unknown true distribution on the data x{\boldsymbol{\mathbf{x}}} as pdata{p_{\textrm{data}}}. The discriminator D∈X→D\in\mathcal{X}\rightarrow is a soft classifier predicting if a data point is real as opposed to being sampled from pG{p_{G}}.

If DD converges optimally for a fixed GG, then D=pdata/(pdata+pG)D={p_{\textrm{data}}}/({p_{\textrm{data}}}+{p_{G}}), and if both DD and GG converge then pG=pdata{p_{G}}={p_{\textrm{data}}} (Goodfellow et al., 2014). GAN training forms a game between DD and GG. In practice DD is often better at estimating the density ratio than G is at generating high-fidelity samples (Shibuya, 2017). This motivates wrapping an imperfect GG to obtain an improved G′G^{\prime} by using the density ratio information contained in DD.

Methods

In this section we show how to sample from the distribution pD{p_{D}} implied by the discriminator DD. We apply (2) and (3) for a target of p⋆=pD{p^{\star}}={p_{D}} and proposal q=pGq={p_{G}}:

The ratio pD/pG{p_{D}}/{p_{G}} is computed entirely from the discriminator scores DD. If DD is perfect, pD=pdata{p_{D}}={p_{\textrm{data}}}, so the sampler will marginally sample from pdata{p_{\textrm{data}}}. The use of (6) is further illustrated in Algorithm 1.

A toy one-dimensional example with just such a perfect discriminator is shown in Figure 2. In this example the MH-GAN is able to correctly reconstruct a missing mode in the generating distribution from the tail of a faulty generator.

The probabilities for DD must not merely provide a good AUC score, but must also be well calibrated. In other words, if one were to warp the probabilities of the perfect discriminator in (1) it may still suffice for standard GAN training, but it will not work in the MCMC procedure defined in (6), as it will result in erroneous density ratios.

We can demonstrate the miscalibration of DD using the statistic of Dawid (1997) on held out samples x1:N{\boldsymbol{\mathbf{x}}}_{1:N} and real/fake labels y1:N∈{0,1}Ny_{1:N}\in\{0,1\}^{N}. If DD is well calibrated, i.e., yy is indistinguishable from a y∼Bern(D(x))y\sim\textrm{Bern}(D({\boldsymbol{\mathbf{x}}})), then

That is, we expect the ZZ diagnostic to be a Gaussian in large NN for any well-calibrated classifier. This means that for large values of ZZ, such as when ∣Z∣>2|Z|>2, we reject the hypothesis that DD is well-calibrated.

Correcting Calibration

Initialization

We also avoid the burn-in issues that usually plague MCMC methods. Recall that via the detailed balance property (Gilks et al., 1996, Ch. 1), if the marginal distribution of a Markov chain state x∈X{\boldsymbol{\mathbf{x}}}\in\mathcal{X} at time step kk matches the target pD{p_{D}} (xk∼pD{\boldsymbol{\mathbf{x}}}_{k}\sim{p_{D}}), then the marginal at time step k+1k+1 will also follow pD{p_{D}} (xk+1∼pD{\boldsymbol{\mathbf{x}}}_{k+1}\sim{p_{D}}). In most MCMC applications it is not possible to get an initial sample from the target distribution (x0∼pD{\boldsymbol{\mathbf{x}}}_{0}\sim{p_{D}}).

However, for MH-GAN, we have access to real data from the target distribution. By initializing the chain at a sample of real data (the correct distribution), we apply the detailed balance property and avoid burn-in. If no generated sample is accepted by the end of the chain, we restart sampling from a synthetic sample to ensure the initial real sample is never output. To make restarts rare, we set KK large (often 640).

Using a restart after an MCMC chain of only rejects has a theoretical potential for bias. However, MCMC in practice often uses chain diagnostics as a stopping criterion, which suffers the same bias potential (Cowles et al., 1999). Alternatively, we could never restart and always report the state after KK samples, which will occasionally include the initial real sample. This might be a better approach in certain statistical problems, where we care more about eliminating any potential source of bias, than in image generation.

Perfect Discriminator

The assumption of a perfect DD may be weakened for two reasons: (A) Because we recalibrate the discriminator, the actual probabilities can be incorrect as long as the decision boundary between real and fake is correct. (B) Because the discriminator is only ever evaluated at samples from GG or the initial real sample x0{\boldsymbol{\mathbf{x}}}_{0}, DD only needs to be accurate on the manifold of samples from the generator pG{p_{G}} and the real data pdata{p_{\textrm{data}}}.

Results

We first show an illustrative synthetic mixture model example followed by real data with images.

We consider the 5×55\times 5 grid of two-dimensional Gaussians used in Azadi et al. (2018), which has become a popular toy example in the GAN literature (Dumoulin et al., 2016). The means are arranged on the grid μ∈{−2,−1,0,1,2}\mu\in\{{-2},{-1},0,1,2\} and use a standard deviation of σ=0.05\sigma=0.05.

Visual results

In Figure 3, we show the original data along with samples generated by the GAN. We also show samples enhanced via the MH-GAN (with calibration) and with DRS. The standard GAN creates spurious links along the grid lines between modes and misses some modes along the bottom row. DRS is able to reduce some of the spurious links but not fill in the missing modes. The MH-GAN further reduces the spurious links and recovers these under-estimated modes.

Quantitative results

These results are made more quantitative in Figure 4, where we follow some of the metrics for the example from Azadi et al. (2018). We consider the standard deviations within each mode in Figure 4(a) and the rate of “high quality” samples in Figure 4(b). A sample is assigned to a mode if its L2L_{2} distance is within four standard deviations (≤4σ=0.2\leq 4\sigma=0.2) of its mean. Samples within four standard deviations of any mixture component are considered “high quality”. The within standard deviation plot (Figure 4(a)) shows a slight improvement for MH-GAN, and the high quality sample rate (Figure 4(b)) approaches 100% faster for the MH-GAN than the GAN or DRS.

To test the spread of the distribution, we inspect the categorical distribution of the closest mode. Far away (non-high quality) samples are assigned to a 26th unassigned category. This categorical distribution should be uniform over the 25 real modes for a perfect generator. To assess generator quality, we look at the Jensen-Shannon divergence (JSD) between the sample mode distribution and a uniform distribution. This is a much more stringent test of appropriate spread of probability mass than checking if a single sample is produced near a mode (as in Azadi et al. (2018)).

In Figure 4(c), we see that the MH-GAN improves the JSD over DRS by 5×5\times on average, meaning it achieves a much more balanced spread across modes. DRS fails to make gains after epoch 30. Using the principled approach of the MH-GAN along with calibrated probabilities ensures a correct spread of probability mass.

2 Real Data

For real data experiments we considered the CelebA (Liu et al., 2015) and CIFAR-10 (Torralba et al., 2008) data sets modeled using the DCGAN (Radford et al., 2015) and WGAN (Arjovsky et al., 2017; Gulrajani et al., 2017). To evaluate the generator G′G^{\prime}, we plot Inception scores (Salimans et al., 2016) per epoch in Figure 5(a) after k=640k=640 MCMC iterations. Figure 5(b) shows Inception score per MCMC iteration: most gains are made in the first k=100k=100 iterations, but gains continue to k=400k=400. This shows that the MH-GAN allows a tunable trade-off between sample quality and computation cost.

In Table 1, we summarize performance (Inception score) across all experiments, running MCMC to k=640k=640 iterations in all cases. Behavior is qualitatively similar to that in Figure 5(a). While DRS improves on a direct GAN, MH-GAN improves Inception score more in every case. Calibration helps in every case; and we found a slight advantage for isotonic regression over other calibration methods. Results are computed at epoch 60, and as in Figure 5(a), error bars and p-values are computed using a paired t-test across Inception score batches. All results are significantly better than the baseline GAN at p<0.05p<0.05.

In Figure 5(c), we show what G′G^{\prime} does to the distribution on discriminator scores. MCMC shifts the distribution of the fakes to match the distribution on true images. We also observed that the MH acceptance rate is primarily determined by the overlap of the distributions on DD scores between real and fake samples. If the AUC of DD is less than 0.90 we see acceptance rates over 20%; but when the AUC of DD is 0.95, acceptance rates drop to 10%.

Calibration results

Figure 6 shows the results per epoch for both CIFAR-10 and CelebA. It shows that the raw discriminator is highly miscalibrated, but can be fixed with any of the calibration methods. The ZZ statistic for the raw discriminator DD (DCGAN on CIFAR-10) varies from −77.57-77.57 to 48.9848.98 in the first 60 epochs; even after Bonferroni correction at N ⁣ ⁣= ⁣ ⁣60N\!\!=\!\!60, we expect ∣Z∣<3.35|Z|<3.35 with 95% confidence for a calibrated classifier. The calibrated discriminator varies from −2.91-2.91 to 3.603.60, showing almost perfect calibration. Accordingly, it is unsurprising that the calibrated discriminator significantly boosts performance in the MH-GAN.

Visual results

We show example images from the CIFAR-10 and CelebA setups in the Appendix A (Figures 11–12). The selectors (such as MH-GAN) result in a wider spread of probability mass across background colors. For CIFAR-10, it enhances modes with animal-like outlines and vehicles.

3 Progressive GAN

To further illustrate the power of the MH-GAN approach we consider the progressive GAN (PGAN) (Karras et al., 2017), which recently produced shockingly realistic images. We applied the MH-GAN to a PGAN using the same setup as with DCGAN, at k=800k=800. We used the pre-trained network of Karras et al. (2017) on CelebA-HQ (1024×\times1024). Large batches of samples are in Appendix A (Figures 7–10).

In Table 2, we use the PGAN as our base GAN and generate random samples from the base, as well as from the addition of DRS and MH-GAN selectors. The different selectors (DRS and MH-GAN) are run on the same batches of images, so the same images may appear for both generators. Although the PGAN sometimes produces near photorealistic images, it also produces many flawed nightmare like images. To assess image quality, five human labelers manually labeled images as warped or acceptable. Table 2 shows that MH-GAN selects significantly fewer warped images.

Both DRS and MH-GAN show an ability to select just the realistic images. The MH-GAN samples are nearly perfect, while DRS still has many flawed samples.

Conclusions

We have shown how to incorporate the knowledge in the discriminator DD into an improved generator G′G^{\prime}. Our method is based on the premise that DD is better at density ratio estimation than GG is at sampling data, which may be a harder task. The principled MCMC setup selects among samples from GG to correct biases in GG. This is the only method in the literature which has the property that given a perfect DD one can recover GG such that pG=pdata{p_{G}}={p_{\textrm{data}}}.

We have shown the raw discriminators in GANs and DRS are poorly calibrated. To our knowledge, this is the first work to evaluate the discriminator in this way and to rigorously show the poor calibration of the discriminator. Because the MH-GAN algorithm may be used to wrap any other GAN, there are countless possible use cases.

Acknowledgements

We thank Rosanne Liu and Zoubin Ghahramani for useful discussions and comments.

References

Appendix A Supplementary Material

In this section we present some of the samples from the various GAN setups in full page figures below.

We also note that the GAN approach to density estimation is complementary to the earlier density ratio estimation approach (Sugiyama et al., 2012). In density ratio estimation, the generator GG is fixed, and the density is found by combining Bayes’ rule and the learned classifier DD. In GANs, the key is learning GG well; while in density ratio estimation, the key is learning DD well. The MH-GAN has flavors of both in that it uses both GG and DD to build G′G^{\prime}.