Skip to content
AI-grafen
FAI engineeringGenerative models· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

Variational autoencoders

Be able to derive the ELBO and train a VAE.

Prerequisites

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:

  1. Since each input covers a small region, the regions overlap and the space becomes coherent.
  2. 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 z∼N(0,I)z \sim \mathcal{N}(0, I), 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 log⁡p(x)\log p(x), but the integral over all zz is intractable:

p(x)=∫p(x∣z) p(z) dzp(x) = \int p(x \mid z)\,p(z)\, dz

The variational approach: introduce an approximate posterior qϕ(z∣x)q_\phi(z\mid x) and rewrite. For any qq:

log⁡p(x)=Eq[log⁡pθ(x∣z)]−KL ⁣(qϕ(z∣x) ∥ p(z))⏟the ELBO+KL ⁣(qϕ(z∣x) ∥ p(z∣x))⏟≥0\log p(x) = \underbrace{\mathbb{E}_{q}\left[\log p_\theta(x\mid z)\right] - \mathrm{KL}\!\left(q_\phi(z\mid x)\,\|\,p(z)\right)}_{\text{the ELBO}} + \underbrace{\mathrm{KL}\!\left(q_\phi(z\mid x)\,\|\,p(z\mid x)\right)}_{\geq 0}

The last term is unknown but always non-negative. The ELBO is therefore a lower bound on log⁡p(x)\log p(x), and maximising the ELBO pushes the floor up.

The ELBO has two parts with a clear interpretation:

TermMeans
Eq[log⁡p(x∣z)]\mathbb{E}_q[\log p(x\mid z)]reconstruction — can the decoder recreate x?
−KL(q∥p)-\mathrm{KL}(q \| p)regularisation — does the encoded distribution resemble the prior?

The reparameterisation trick. Sampling z∼N(μ,σ2)z \sim \mathcal{N}(\mu, \sigma^2) is not differentiable with respect to μ\mu and σ\sigma. The solution is to move the randomness out:

z=μ+σ⊙ε,ε∼N(0,I)z = \mu + \sigma \odot \varepsilon, \qquad \varepsilon \sim \mathcal{N}(0, I)

Now zz is a differentiable function of μ\mu and σ\sigma, and ε\varepsilon 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:

KL=−12∑j(1+log⁡σj2−μj2−σj2)\mathrm{KL} = -\tfrac{1}{2}\sum_j \left(1 + \log\sigma_j^2 - \mu_j^2 - \sigma_j^2\right)

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 zz is. It happens especially when the decoder is powerful (autoregressive, for instance).

RemedyHow
KL warm-upstart with weight 0 on the KL term and increase it gradually
Free bitsallow a certain KL per dimension without a penalty
A weaker decodergive it less capacity to manage without z
β-VAE with β < 1weight 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

All the sources and licences