A plain autoencoder compresses data but its latent space is disorganised: decoding a random point gives garbage. Kingma and Welling's Variational Autoencoder (2013) — with the parallel work of Rezende, Mohamed and Wierstra — made the latent space probabilistic and smooth, so we can sample new data by decoding random latent vectors. VAEs introduced ideas central to modern generative modelling, and their autoencoder component lives on inside latent diffusion models.
The generative model#
Assume each data point $\mathbf{x}$ is generated from a latent variable $\mathbf{z}$:
where the decoder $p_\theta(\mathbf{x} \mid \mathbf{z})$ is a neural network (e.g. outputting pixel means). The marginal likelihood
is intractable: we cannot integrate over all $\mathbf{z}$, and the true posterior $p_\theta(\mathbf{z} \mid \mathbf{x})$ is intractable too.
Amortised variational inference#
Introduce an encoder $q_\phi(\mathbf{z} \mid \mathbf{x})$ — a neural network outputting the mean $\boldsymbol{\mu}$ and (log-)variance $\boldsymbol{\sigma}^2$ of a Gaussian — to approximate the posterior. "Amortised" means one network infers latents for all data points rather than optimising a separate distribution for each.
The Evidence Lower Bound (ELBO)#
For any $q$, the log-likelihood decomposes as
The last KL term is non-negative, so the ELBO is a lower bound on $\log p_\theta(\mathbf{x})$. We maximise it jointly over encoder and decoder. Its two terms have clear meanings:
- Reconstruction term: decoded samples should reproduce the input.
- KL regulariser: each encoding distribution should stay close to the prior $\mathcal{N}(\mathbf{0}, \mathbf{I})$ — packing the latent space densely and smoothly so random prior samples decode to sensible data.
For Gaussian $q$ and standard normal prior, the KL has a closed form:
The reparameterisation trick#
We need gradients of $\mathbb{E}_{q_\phi}[\cdot]$ with respect to $\phi$, but sampling is not differentiable. Rewrite the sample as a deterministic function of $\phi$ and independent noise:
Now gradients flow through $\boldsymbol{\mu}$ and $\boldsymbol{\sigma}$; the randomness sits in $\boldsymbol{\epsilon}$. This low-variance gradient estimator is what made VAEs trainable with ordinary backpropagation.
Implementation#
import torch
import torch.nn as nn
import torch.nn.functional as F
class VAE(nn.Module):
def __init__(self, x_dim=784, h=400, z_dim=20):
super().__init__()
self.enc = nn.Sequential(nn.Linear(x_dim, h), nn.ReLU())
self.mu, self.logvar = nn.Linear(h, z_dim), nn.Linear(h, z_dim)
self.dec = nn.Sequential(nn.Linear(z_dim, h), nn.ReLU(), nn.Linear(h, x_dim))
def forward(self, x):
hdn = self.enc(x)
mu, logvar = self.mu(hdn), self.logvar(hdn)
z = mu + torch.exp(0.5 * logvar) * torch.randn_like(mu) # reparameterisation
return self.dec(z), mu, logvar
def vae_loss(logits, x, mu, logvar, beta=1.0):
recon = F.binary_cross_entropy_with_logits(logits, x, reduction="sum")
kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
return (recon + beta * kl) / x.size(0)
model = VAE()
x = torch.rand(64, 784) # e.g. flattened MNIST images in [0, 1]
logits, mu, logvar = model(x)
print(vae_loss(logits, x, mu, logvar).item())
with torch.no_grad(): # generation: decode samples from the prior
samples = torch.sigmoid(model.dec(torch.randn(16, 20)))After training on MNIST, decoding a grid of latent points shows digits morphing smoothly into one another — evidence of a continuous, structured latent space.
Known issues#
- Blurry samples: with a Gaussian (MSE-like) or per-pixel likelihood, the decoder averages over uncertainty, producing blurry images compared with GANs and diffusion models.
- Posterior collapse: with powerful decoders (e.g. autoregressive), the model may ignore $\mathbf{z}$, making the KL term zero. Remedies include KL annealing (gradually increasing its weight) and "free bits".
- The prior hole problem: regions of the prior rarely covered by any encoding decode poorly.
Variants#
- β-VAE (Higgins et al., 2017): weight the KL by $\beta > 1$ to encourage disentangled latents where individual dimensions capture separate factors (e.g. rotation, scale) — at some cost to reconstruction.
- Conditional VAE (CVAE): condition encoder and decoder on a label or other input.
- VQ-VAE (van den Oord et al., 2017): a discrete latent codebook with vector quantisation; paired with an autoregressive prior it generated high-quality images and audio, and discrete codebooks underlie many image and audio tokenisers used by modern multimodal models.
- Hierarchical VAEs (NVAE, VDVAE): many latent layers; much sharper samples.
- The autoencoder in Stable Diffusion is a VAE-style model (with perceptual and adversarial losses) that compresses images into a latent space where diffusion operates.