Variational autoencoders
Be able to derive the ELBO and train a VAE.
Prerequisites
- EAutoencodersrequired
- EInformation theory: entropy and KL divergencerequired
Intuition
An ordinary autoencoder encodes each input to a point in the latent space. Between the points there are holes where the decoder produces nonsense — which is why it cannot be sampled from.
A VAE encodes to a distribution instead: the encoder gives a mean and a spread, and a random sample is drawn from it.
Two things follow from that:
- Since each input covers a small region, the regions overlap and the space becomes coherent.
- A penalty term pulls the distributions towards a standard normal distribution, so that the whole space is filled out.
Then it can be sampled: draw , run the decoder, get something new.
The price is blurrier reconstructions than an ordinary autoencoder's. That is a direct consequence of the model being forced to make the space usable instead of just minimising the reconstruction error.
Derivation
The goal is to maximise , but the integral over all is intractable:
The variational approach: introduce an approximate posterior and rewrite. For any :
The last term is unknown but always non-negative. The ELBO is therefore a lower bound on , and maximising the ELBO pushes the floor up.
The ELBO has two parts with a clear interpretation:
| Term | Means |
|---|---|
| reconstruction — can the decoder recreate x? | |
| regularisation — does the encoded distribution resemble the prior? |
The reparameterisation trick. Sampling is not differentiable with respect to and . The solution is to move the randomness out:
Now is a differentiable function of and , and is a constant input. Without that trick the VAE cannot be trained with gradient descent at all.
The KL term in closed form for a Gaussian posterior and a standard normal prior:
Posterior collapse is the classic problem: the KL term goes to zero, the encoder is ignored, and the decoder learns to produce an average whatever is. It happens especially when the decoder is powerful (autoregressive, for instance).
| Remedy | How |
|---|---|
| KL warm-up | start with weight 0 on the KL term and increase it gradually |
| Free bits | allow a certain KL per dimension without a penalty |
| A weaker decoder | give it less capacity to manage without z |
| β-VAE with β < 1 | weight the KL term down |
β-VAE with β > 1 is the opposite setting: harder regularisation gives more disentangled latent dimensions, at the price of worse reconstruction.
Code
import torch, torch.nn as nn, torch.nn.functional as F
class VAE(nn.Module):
def __init__(self, in_dim=784, hidden=400, latent=20):
super().__init__()
self.enc = nn.Sequential(nn.Linear(in_dim, hidden), nn.ReLU())
self.mu = nn.Linear(hidden, latent)
self.log_var = nn.Linear(hidden, latent)
self.dec = nn.Sequential(nn.Linear(latent, hidden), nn.ReLU(),
nn.Linear(hidden, in_dim))
def encode(self, x):
h = self.enc(x)
return self.mu(h), self.log_var(h)
def reparameterise(self, mu, log_var):
std = torch.exp(0.5 * log_var)
return mu + std * torch.randn_like(std) # the randomness moved out → differentiable
def forward(self, x):
mu, log_var = self.encode(x)
z = self.reparameterise(mu, log_var)
return self.dec(z), mu, log_var
def elbo(logits, x, mu, log_var, beta=1.0, free_bits=0.0):
rec = F.binary_cross_entropy_with_logits(logits, x, reduction="sum") / len(x)
kl_per_dim = -0.5 * (1 + log_var - mu.pow(2) - log_var.exp())
if free_bits > 0:
kl_per_dim = torch.clamp(kl_per_dim, min=free_bits) # allow a certain KL for free
kl = kl_per_dim.sum(dim=1).mean()
return rec + beta * kl, {"reconstruction": float(rec), "kl": float(kl)}
# KL warm-up counteracts posterior collapse
def beta_schedule(step, warmup=5000, max_beta=1.0):
return min(max_beta, max_beta * step / warmup)
def train(model, dataloader, opt, total_steps=20000, free_bits=0.05):
step = 0
for x, _ in dataloader:
x = x.flatten(1)
logits, mu, log_var = model(x)
loss, parts = elbo(logits, x, mu, log_var,
beta=beta_schedule(step, 5000), free_bits=free_bits)
opt.zero_grad(); loss.backward(); opt.step()
step += 1
if step % 1000 == 0:
print(f" step {step}: rec {parts['reconstruction']:.1f} kl {parts['kl']:.3f}")
if step >= total_steps:
break
# Detect posterior collapse: the KL per dimension near zero
def collapsed_dimensions(model, dataloader, threshold=0.01):
model.eval()
kl_sum = None
n = 0
with torch.no_grad():
for x, _ in dataloader:
mu, log_var = model.encode(x.flatten(1))
kl = -0.5 * (1 + log_var - mu.pow(2) - log_var.exp())
kl_sum = kl.sum(0) if kl_sum is None else kl_sum + kl.sum(0)
n += len(x)
per_dim = kl_sum / n
return {"kl_per_dim": [round(float(v), 4) for v in per_dim],
"number_collapsed": int((per_dim < threshold).sum()),
"out_of": len(per_dim)}
# {'number_collapsed': 14, 'out_of': 20} ← only 6 dimensions are used
# Generate: sample from the prior
@torch.no_grad()
def sample(model, n=16, latent=20):
return torch.sigmoid(model.dec(torch.randn(n, latent)))
collapsed_dimensions is the diagnostic that is missing most often. A VAE can look as though it is training well while fourteen of twenty latent dimensions are unused — and then the latent representation is far narrower than you think.
Mastery means
- Derives the ELBO
- Explains the reparameterisation trick
- Recognises and remedies posterior collapse
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Auto-Encoding Variational Bayes — arXiv (open access; licence per article)
- Higgins m.fl. — beta-VAE (ICLR 2017) — open review, free to read
- PyTorch — tutorials (BSD-3) — BSD-3-Clause