auto-encoding-variational-bayes

drawing.png

This is the Variational Auto-Encoder paper by Kingma & Welling.

For a simple pytorch implementation, check out my github repo, where I auto-encoded CryptoPunks.

Problem statement

There is a random variable XX that is generated by:

  • sampling a latent variable z∼pθ∗(z)z\sim p_{\theta^*}(z)
  • sampling x∼pθ∗(x∣z)x\sim p_{\theta^*}(x\vert z)

where the true parameters θ∗\theta^* of the distribution and the latent variable zz are hidden.

Reminder on the evidence lower bound

Suppose that we want to estimate the distribution pθ∗p_{\theta^*} by estimating parameters θ\theta. θ\theta is both used to generate zz and then xx (by sampling pθ(x∣z)p_\theta({x\vert z})).

Let's say we found some appropriate θ\theta, it's intractable to estimate pθ(x)p_\theta(x) for a given xx because we need to marginalize over all zz's:

pθ(x)=∫pθ(x∣z)pθ(z)dzp_\theta(x)=\int p_\theta(x\vert z)p_\theta(z)dz

However, we also know that, given a certain zz we have: pθ(x,z)=pθ(x∣z)p(z)p_\theta(x,z) = p_\theta(x\vert z)p(z) and pθ(x,z)=pθ(z∣x)p(x)p_\theta(x, z) = p_\theta(z\vert x)p(x).

Combining the two, for any zz, we have:

pθ(x)=pθ(z)pθ(x∣z)pθ(z∣x)p_\theta(x)=\frac{p_\theta(z)p_\theta(x\vert z)}{p_\theta(z\vert x)}

If we can find a good approximation of pθ(z∣x)p_\theta(z\vert x), we can compute pθ(x)p_\theta(x). However:

  • we might not have a closed form solution
  • pθ(z∣x)=pθ(x∣z)pθ(z)/pθ(x)p_\theta(z\vert x) = p_\theta(x\vert z)p_\theta(z)/p_\theta(x) and we already established that p(x)p(x) is intractable.

Therefore, we'll use another distribution family qϕ(z∣x)q_\phi(z\vert x) to approximate pθ(z∣x)p_\theta(z\vert x).

Now, to estimate the best model pθp_\theta we wish to maximize its likelihood, which is the probability density function: pθ(x)p_\theta(x). For convenience, we usually use the log-likelihood. The optimization problem is equivalent since log⁡\log is a monotonously increasing function.

ln⁡pθ(x)=ln⁡∫pθ(x,z)dz=Ez∼pθ(x,z)[1]\ln p_\theta(x) = \ln \int p_\theta(x, z)dz=\mathbb{E}_{z\sim p_{\theta}(x,z)}[1]

Recall that this is intractable, because of the integral over zz. We can estimate this using importance sampling:

ln⁡pθ(x)=ln⁡Eqϕ(z∣x)pθ(x,z)qϕ(z∣x)≥Eqϕ(z∣x)ln⁡pθ(x,z)qϕ(z∣x) (Jensen’s inequality)\ln p_\theta(x) = \ln \mathbb{E}_{q_\phi(z\vert x)}\frac{p_\theta(x, z)}{q_\phi(z\vert x)} \geq \mathbb{E}_{q_\phi(z\vert x)}\ln\frac{p_\theta(x, z)}{q_\phi(z\vert x)}\text{ (Jensen's inequality)}

Now, this term: Eqϕ(z∣x)ln⁡pθ(x,z)qϕ(z∣x)\mathbb{E}_{q_\phi(z\vert x)}\ln\frac{p_\theta(x, z)}{q_\phi(z\vert x)} is what we call the evidence lower bound L(θ,ϕ)\mathcal{L}(\theta, \phi). But it's not over. Let's put it on the left hand side and combine it with the log-likelihood term ln⁡pθ(x)\ln p_\theta(x):

ln⁡pθ(x)−Eqϕ(z∣x)ln⁡pθ(x,z)qϕ(z∣x)≥0\ln p_\theta(x) - \mathbb{E}_{q_\phi(z\vert x)}\ln\frac{p_\theta(x, z)}{q_\phi(z\vert x)} \geq 0

where

ln⁡pθ(x)−Eqϕ(z∣x)ln⁡pθ(x,z)qϕ(z∣x)=−Eqϕ(z∣x)[ln⁡pθ(x,z)qϕ(z∣x)−ln⁡pθ(x)]=−Eqϕ(z∣x)[ln⁡pθ(x,z)/pθ(x)qϕ(z∣x)]=−Eqϕ(z∣x)[ln⁡pθ(z∣x)qϕ(z∣x)]=DKL(qϕ(z∣x)∥pθ(z∣x))\begin{aligned}\ln p_\theta(x) - \mathbb{E}_{q_\phi(z\vert x)}\ln\frac{p_\theta(x, z)}{q_\phi(z\vert x)} & = -\mathbb{E}_{q_\phi(z\vert x)}[\ln\frac{p_\theta(x, z)}{q_\phi(z\vert x)} - \ln p_\theta(x)] \\& = -\mathbb{E}_{q_\phi(z\vert x)}[\ln\frac{p_\theta(x, z)/p_\theta(x)}{q_\phi(z\vert x)}] \\& = -\mathbb{E}_{q_\phi(z\vert x)}[\ln\frac{p_\theta(z\vert x)}{q_\phi(z\vert x)}] \\& = D_{KL}(q_\phi(z\vert x)\Vert p_\theta(z\vert x)) \\\end{aligned}

So we basically have:

ln⁡pθ(x)⏟log likelihood−L(θ,ϕ)⏟ELBO=DKL(qϕ(z∣x)∥pθ(z∣x))⏟KL divergence≥0\underbrace{\ln p_\theta(x)}_{\text{log likelihood}} - \underbrace{\mathcal{L}(\theta, \phi)}_{\text{ELBO}} = \underbrace{D_{KL}(q_\phi(z\vert x)\Vert p_\theta(z\vert x))}_{\text{KL divergence}} \geq 0

Therefore, by maximizing the ELBO we are maximizing the log-likelihood. The point of the paper is to find a way to differentiate and optimize the ELBO with low variance.

To recap, here's our encoder: qϕ(z∣x)≈pθ(z∣x)q_\phi(z\vert x)\approx p_\theta(z\vert x); here's our decoder pθ(x∣z)p_\theta(x\vert z). We'll learn ϕ\phi and θ\theta jointly.

We want the algo to work in the case of:

  • intractibility: can't compute pθ(z)=∫pθ(z)pθ(x∣z)dzp_\theta(z) = \int p_\theta(z)p_\theta(x\vert z)dz or pθ(z∣x)=pθ(x∣z)pθ(z)pθ(x)p_\theta(z\vert x)=\frac{p_\theta(x\vert z)p_\theta(z)}{p_\theta(x)} (EM can't be used). pθ(x)p_\theta(x) intractable because of large dimensionality (e.g. images). pθ(x∣z)p_\theta(x\vert z) intractable due to large number of hidden variables (in case of neural net with nonlinear hidden layer for example).
  • large dataset: monte carlo EM would be too slow (expensive sampling loop per datapoint)

Some relevant applications:

  • the parameters θ\theta can be of interest if we're analyzing some natural process or want to generate artificial data (by sampling p(x∣θ)p(x\vert \theta)). We want efficient max likelihood (maximize probability of seeing the data given the model) or max à priori estimation (maximize probability of the model given the data; requires a model prior) of parameters θ\theta.
  • representing data (e.g. generating image embeddings): posterior inference of zz given xx
  • marginal inference of xx for tasks where a prior over xx is required like image denoising, inpainting, superresolution

Stochastic Gradient Variational Bayes

The ELBO L(pθ,qϕ)L(p_\theta, q_\phi) can also be written as L(pθ,qϕ)=−DKL(qϕ(z∣x)∥pθ(z))+Eqϕ(z∣x)(log⁡pθ(x∣z))L(p_\theta, q_\phi)=-D_{KL}(q_\phi(z\vert x)\lVert p_\theta(z)) + \mathbb{E_{q_\phi(z\vert x)}}(\log p_\theta(x\vert z))

The first term is the KL divergence which can be seen as a regularization term with respect to the prior pθ(z)p_\theta(z). The second term is the reconstruction loss: given zz we want to reconstruct xx.

Naively, we can use a monte carlo gradient estimator for either term: ∇ϕEqϕ(z)[f(z)]=∇ϕ∫zq(z)f(z)dz=∫z∇ϕq(z)f(z)dz=∫zq(z)1q(z)∇ϕq(z)⏟∇ϕlog⁡qϕ(z)f(z)dz=E[f(z)∇ϕlog⁡qϕ(z)]\nabla_\phi \mathbb{E}_{q_\phi(z)}[f(z)]=\nabla_\phi \int_z q(z)f(z)dz=\int_z \nabla_\phi q(z) f(z) dz =\int_z q(z)\underbrace{\frac{1}{q(z)}\nabla_\phi q(z)}_{\nabla_\phi \log q_\phi(z)} f(z) dz = \mathbb{E}[f(z)\nabla_\phi \log q_\phi(z)]

where f(z)=log⁡pθ(z)qϕ(z∣x)f(z) = \log\frac{p_\theta(z)}{q_\phi(z\vert x)} or f(z)=log⁡pθ(x∣z)f(z)=\log p_\theta(x\vert z)

Thus: ∇ϕEqϕ(z)[f(z)]≈1L∑f(z)∇ϕlog⁡qϕ(z)\nabla_\phi \mathbb{E}_{q_\phi(z)}[f(z)] \approx \frac{1}{L}\sum f(z)\nabla_\phi \log q_\phi(z)

However, this estimator has very high variance.

We can re-parameterize qϕ(z∣x)q_\phi(z\vert x) as z~=gϕ(ϵ,x),ϵ∼p(ϵ)\tilde z = g_\phi(\epsilon , x), \epsilon\sim p(\epsilon).

For instance, z=N(μ,σ)z=\mathcal{N}(\mu, \sigma): z~=μ+σϵ,ϵ∼N(0,1)\tilde z=\mu + \sigma \epsilon, \epsilon\sim\mathcal{N}(0, 1)

Eqϕ(z∣x)[f(z)]=Ep(ϵ)[f(gϕ(ϵ,x))]≈1L∑f(gϕ(ϵ,x))\mathbb{E}_{q_\phi(z\vert x)}[f(z)]=\mathbb{E}_{p(\epsilon)}[f(g_\phi(\epsilon, x))]\approx \frac{1}{L}\sum f(g_\phi(\epsilon, x))

KL divergence can be integrated analytically (e.g. when prior pθ(z)p_\theta(z) and posterior qϕ(z∣x)q_\phi(z\vert x) are gaussian, see appendix B in the paper). KL divergence can be interpreted as regularization to encourage q(x∣z)q(x\vert z) to be close to prior. Only the expected reconstruction error requires estimation Eqϕ(z∣x)[log⁡pθ(x∣z)]\mathbb{E}_{q_\phi(z\vert x)}[\log p_\theta(x\vert z)]

The second reason for the re-parameterization trick is that it makes the random sampling step differentiable. Instead of sampling z∼N(μ,σ2)z\sim \mathcal{N}(\mu, \sigma^2), we compute zz as a deterministic transformation of ϵ∼N(0,1)\epsilon \sim \mathcal{N}(0, 1) and can compute gradients w.r.t. μ\mu and σ\sigma.

Example: Variational Auto-Encoder

  • we set a prior over the latent pθ(z)=N(z;0,I)p_\theta(z)=\mathcal{N}(z; 0, I).
  • For discrete data, encoder pθ(x∣z)p_\theta(x\vert z) is a multivariate Bernoulli distribution:
log⁡p(x∣z)=∑xilog⁡yi+(1−xi)log⁡(1−yi)\log p(x\vert z) = \sum x_i \log y_i + (1-x_i)\log(1-y_i)

(product of each pixel independently, p(xi∣z)=yip(x_i\vert z) = y_i if xi=1x_i=1 and 1−yi1-y_i otherwise)

where yi=sigmoid(W2+tanh⁡(W1z+b1)+b2)y_i=\text{sigmoid}(W_2 + \tanh(W_1 z +b_1) + b_2)

  • For real value data, encoder pθ(x∣z)p_\theta(x\vert z) is a multivariate gaussian:
log⁡p(x∣z)=log⁡N(x;μ,σ2I)\log p(x\vert z) = \log \mathcal{N}(x; \mu, \sigma^2 I)

where μ=W1h+b1,log⁡σ2=W2h+b2,h=tanh⁡(W0z+b0)\mu=W_1 h + b_1, \log \sigma^2 = W_2 h + b_2, h =\tanh(W_0 z + b_0).

  • The true posterior pθ(z∣x)p_\theta(z\vert x) is intractable. We approximate it with decoder qϕ(z∣x)q_\phi(z\vert x). For qϕ(z∣x)q_\phi(z\vert x), same formula as a multivariate gaussian (with different parameters) and zz and xx are swapped

KL divergence can be computed and differentiated analytically. For reconstruction loss, we sample z∼μ+σϵz\sim \mu + \sigma \epsilon where ϵ∼N(0,I)\epsilon\sim\mathcal{N}(0,I).

log⁡likelihood=L(θ,ϕ;x(i))≈12∑j=1J(1+log⁡((σj(i))2)−(μj(i))2−(σj(i))2)+1L∑l=1Llog⁡pθ(x(i)∣z(i,l))\log\text{likelihood} = \mathcal{L}(\theta, \phi; x^{(i)}) \approx \frac{1}{2}\sum_{j=1}^J(1+\log((\sigma_j^{(i)})^2) - (\mu_j^{(i)})^2 - (\sigma_j^{(i)})^2) + \frac{1}{L}\sum_{l=1}^L\log p_\theta(x^{(i)}\vert z^{(i, l)})

where z(i,l)=μ(i)+σ(i)⊙ϵ(l)z^{(i, l)} = \mu^{(i)}+\sigma^{(i)} \odot \epsilon^{(l)}

and ϵ(l)∼N(0,I)\epsilon^{(l)} \sim \mathcal{N}(0, I)

The second term 1L∑l=1Llog⁡pθ(x(i)∣z(i,l))\frac{1}{L}\sum_{l=1}^L\log p_\theta(x^{(i)}\vert z^{(i, l)}) is equivalent to the mean-squared error, for real-valued data.

AEVB algorithm:

aevb.png

Posterior collapse

Signal from input xx to posterior parameters is either too weak or too noisy and as a result, decoder starts ignoring zz samples drawn from posterior qϕ(z∣x)q_\phi(z\vert x), i.e. qϕ(z∣x)≈qϕ(z)=N(cst a,cstb)q_\phi(z\vert x)\approx q_\phi (z) = \mathcal{N}(\text{cst} a, \text{cst} b) (the parameters of the distribution collapse to constants). In practice, this produces generic outputs x^\hat{x} that are crude representations of all seens xx's.