Reinterpreting Importance-Weighted Autoencoders

Chris Cremer, Quaid Morris, David Duvenaud

Background

The importance-weighted autoencoder (IWAE; Burda et al. (2016)) is a variational inference strategy capable of producing arbitrarily tight evidence lower bounds. IWAE maximizes the following multi-sample evidence lower bound (ELBO):

which is a tighter lower bound than the ELBO maximized by the variational autoencoder (VAE; Kingma & Welling (2014)):

In this section, we derive the implicit distribution that arises from importance sampling from a distribution pp using qq as a proposal distribution. Given a batch of samples z2...zkz_{2}...z_{k} from q(z∣x)q(z|x), the following is the unnormalized importance-weighted distribution:

Here are some properties of the approximate IWAE posterior:

For a more detailed derivation, see the appendix. Note that we are abusing the VAE lower bound notation because this implies an expectation over an unnormalized distribution. Consequently, we replace the expectation with an equivalent integral.

See section 5.2 for a proof that qEWq_{EW} is a normalized distribution. Using qEWq_{EW} in the VAE ELBO, LVAE[qEW]\mathcal{L}_{VAE}[q_{EW}], results in an upper bound of LIWAE[q]\mathcal{L}_{IWAE}[q]. See section 5.3 for the proof, which is a special case of the proof in Naesseth et al. (2017). The procedure to sample from qEW(z∣x)q_{EW}(z|x) is shown in Algorithm 1. It is equivalent to sampling-importance-resampling (SIR).

3 Visualizing the nonparameteric approximate posterior

Resampling for prediction

During training, we sample the qq distribution and implicitly weight them with the IWAE ELBO. After training, we need to explicitly reweight samples from qq.

In figure 2, we demonstrate the need to sample from qEWq_{EW} rather than q(z∣x)q(z|x) for reconstructing MNIST digits. We trained the model to maximize the IWAE ELBO with K=50 and 2 latent dimensions, similar to Appendix C in Burda et al. (2016). When we sample from q(z∣x)q(z|x) and reconstruct the samples, we see a number of anomalies. However, if we perform the sampling-resampling step (Alg. 1), then the reconstructions are much more accurate. The intuition here is that we trained the model with qEWq_{EW} with K=50K=50 then sampled from q(z∣x)q(z|x) (qEWq_{EW} with K=1K=1), which are very different distributions, as seen in Fig. 1.

Discussion

We’d like to thank an anonymous ICLR reviewer for providing insightful future directions for this work. We’d like to thank Yuri Burda, Christian Naesseth, and Scott Linderman for bringing our attention to oversights in the paper. We’d also like to thank Christian Naesseth for the derivation in section 5.3 and for providing many helpful comments.

References

Appendix

(8): Change of notation z=z1z=z_{1}. (10): ziz_{i} has the same expectation as z1z_{1} so we can replace kk with the sum of kk terms.

(17): Change of notation z=z1z=z_{1}. (19): ziz_{i} has the same expectation as z1z_{1} so we can replace kk with the sum of kk terms. (20): Linearity of expectation.

Let p^(x∣z1:k)=1k(p(x,z)q(z∣x)+∑j=2kp(x,zj)q(zj∣x))\hat{p}(x|z_{1:k})=\frac{1}{k}\left(\frac{p(x,z)}{q(z|x)}+\sum_{j=2}^{k}\frac{p(x,z_{j})}{q(z_{j}|x)}\right)

(28): Given that f(A)=−AlogAf(A)=-AlogA is concave for A>0A>0, and f(E[x])≥E[f(x)]f(E[x])\geq E[f(x)], then f(E[x])=−E[x]logE[x]≥E[−xlogx]f(E[x])=-E[x]logE[x]\geq E[-xlogx]. (30): Change of notation z=z1z=z_{1}. (34): ziz_{i} has the same expectation as z1z_{1} so we can replace kk with the sum of kk terms.

The previous section showed that LIWAE(q)≤LVAE(qEW)\mathcal{L}_{IWAE}(q)\leq\mathcal{L}_{VAE}(q_{EW}). That is, the IWAE ELBO with the base qq is a lower bound to the VAE ELBO with the importance weighted qEWq_{EW}. Due to Jensen's inequality and as shown in Burda et al. (2016), we know that the IWAE ELBO is an upper bound of the VAE ELBO: LIWAE(q)≥LVAE(q){L}_{IWAE}(q)\geq{L}_{VAE}(q). Furthermore, the log marginal likelihood can be factorized into: log(p(x))=LVAE(q)+KL(q∣∣p)log(p(x))={L}_{VAE}(q)+KL(q||p), and rearranged to: KL(q∣∣p)=log(p(x))−LVAE(q)KL(q||p)=log(p(x))-{L}_{VAE}(q).

Following the observations above and substituting qEWq_{EW} for qq:

Thus, KL(qEW∣∣p)≤KL(q∣∣p)KL(q_{EW}||p)\leq KL(q||p), meaning qEWq_{EW} is closer to the true posterior than qq in terms of KL divergence.

5 In the limit of the number of samples

Another perspective is in the limit of k=∞{k=\infty}. Recall that the marginal likelihood can be approximated by importance sampling: