Introduction to Generative AI

Mini Project: Train a Variational Autoencoder on MNIST


By the end of this build you will have a model that takes 20 random numbers and turns them into a handwritten digit that nobody ever wrote. It trains in about five minutes on a laptop CPU and under a minute on a GPU.

That is the easy part. The interesting part is that a variational autoencoder has three characteristic ways of failing, all of which produce a loss curve that looks perfectly healthy. You can train one to completion, see a smooth descending line, and have a model that generates nothing but grey smudges. So this build includes the diagnostics that catch each failure, and you should run them even when everything appears fine.

Interpolating the latent line from a 3 to an 8z of a 30.250.500.75z ofan 8nullstill adigitIf the midpoint is a smear rather than a digit, the KL term was too weak and the space has holes.
A smooth walk between two encodings is the proof that the latent space is organised, not just memorised.

Setup

Bash
python -m venv .venvsource .venv/bin/activate          # Windows: .venv\Scripts\activatepip install torch torchvision matplotlib numpy

Nothing else is needed. MNIST downloads automatically on first run — 60,000 training images of 28×28 greyscale digits, about 10 MB.

Python
import torchimport torch.nn as nnimport torch.nn.functional as Ffrom torch.utils.data import DataLoaderfrom torchvision import datasets, transformsimport matplotlib.pyplot as pltdevice = ("cuda" if torch.cuda.is_available()          else "mps" if torch.backends.mps.is_available()          else "cpu")torch.manual_seed(0)print("device:", device)

Data

Python
tf = transforms.ToTensor()          # -> float tensor in [0, 1]train_ds = datasets.MNIST("./data", train=True,  download=True, transform=tf)test_ds  = datasets.MNIST("./data", train=False, download=True, transform=tf)train_loader = DataLoader(train_ds, batch_size=128, shuffle=True,  drop_last=True)test_loader  = DataLoader(test_ds,  batch_size=256, shuffle=False)

One detail decides the whole loss function: ToTensor() alone, with no normalisation. Pixels stay in [0,1][0, 1]. That range lets you treat each pixel as a Bernoulli probability and use binary cross-entropy for reconstruction, which is both principled and better-behaved than squared error on this data.

If you normalise to [−1,1][-1, 1] — a habit from classification work — binary cross-entropy will produce NaN the moment a target goes negative. The fix in that case is a Tanh output and MSE loss, but simply not normalising is easier.

The model

Python
class VAE(nn.Module):    def __init__(self, latent_dim=20):        super().__init__()        self.latent_dim = latent_dim        self.enc = nn.Sequential(            nn.Linear(784, 512), nn.ReLU(),            nn.Linear(512, 256), nn.ReLU(),        )        self.fc_mu     = nn.Linear(256, latent_dim)        self.fc_logvar = nn.Linear(256, latent_dim)        self.dec = nn.Sequential(            nn.Linear(latent_dim, 256), nn.ReLU(),            nn.Linear(256, 512),        nn.ReLU(),            nn.Linear(512, 784),        nn.Sigmoid(),        )    def encode(self, x):        h = self.enc(x)        return self.fc_mu(h), self.fc_logvar(h)    def reparameterise(self, mu, logvar):        if not self.training:            return mu                       # deterministic when evaluating        std = torch.exp(0.5 * logvar)        return mu + std * torch.randn_like(std)    def forward(self, x):        mu, logvar = self.encode(x.view(-1, 784))        z = self.reparameterise(mu, logvar)        return self.dec(z), mu, logvar

Three choices in there are load-bearing.

Two heads on one trunk. The encoder does not output a code. It outputs a mean and a log-variance — the centre and spread of a small Gaussian blob in latent space. The code fed to the decoder is a random draw from that blob, which forces the decoder to produce a sensible digit from any point in it, not just from one exact location. Blobs have volume, and volume is what fills the gaps that make random sampling work.

Log-variance, not standard deviation. A standard deviation must be positive; a log-variance can be any real number, so the network never has to be constrained and you never take the log of something that drifted to zero.

Sigmoid on the output. It matches the [0,1][0,1] pixel range and pairs correctly with binary cross-entropy.

The loss

Python
def vae_loss(recon, x, mu, logvar, beta=1.0):    x = x.view(-1, 784)    # SUM over pixels and batch -- not mean. See the note below.    recon_loss = F.binary_cross_entropy(recon, x, reduction="sum")    kld = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())    return recon_loss + beta * kld, recon_loss, kld

reduction="sum" is the single most common source of silent failure in a first VAE. With "mean", the reconstruction term is divided by every pixel of every image in the batch (784 × 128) while the KL term is still summed, so the KL outweighs it by around five orders of magnitude. The optimiser then does the obvious thing: it drives the KL to zero by making the encoder output μ=0\mu = 0 and σ=1\sigma = 1 for every input, ignoring the image entirely. The decoder learns to produce the average of all MNIST digits — a grey blur — and returns it for everything. The loss curve looks completely normal throughout.

The KL formula is the closed form for the divergence between the encoder's Gaussian and a standard normal:

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

It is worth checking one value by hand so the code stops being magic. For a single dimension with μ=0.5\mu = 0.5 and log⁡σ2=−0.3\log\sigma^2 = -0.3: σ2=e−0.3=0.741\sigma^2 = e^{-0.3} = 0.741, so the bracket is 1−0.3−0.25−0.741=−0.2911 - 0.3 - 0.25 - 0.741 = -0.291, giving DKL=0.146D_{KL} = 0.146 nats for that dimension. A dimension sitting exactly at μ=0\mu = 0, σ2=1\sigma^2 = 1 contributes precisely zero — it has become the prior and carries no information.

Training with KL annealing

Python
model = VAE(latent_dim=20).to(device)opt = torch.optim.Adam(model.parameters(), lr=1e-3)EPOCHS, WARMUP = 30, 10history = []for epoch in range(EPOCHS):    beta = min(1.0, (epoch + 1) / WARMUP)      # ramp the KL weight in    model.train()    tot = tot_r = tot_k = 0.0    for x, _ in train_loader:        x = x.to(device)        recon, mu, logvar = model(x)        loss, r, k = vae_loss(recon, x, mu, logvar, beta)        opt.zero_grad()        loss.backward()        torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)        opt.step()        n = x.size(0)        tot += loss.item() / n; tot_r += r.item() / n; tot_k += k.item() / n    nb = len(train_loader)    history.append((tot / nb, tot_r / nb, tot_k / nb))    print(f"epoch {epoch+1:2d}  beta {beta:.2f}  "          f"total {tot/nb:7.2f}  recon {tot_r/nb:7.2f}  kl {tot_k/nb:6.2f}")

Two things here are deliberate.

Beta annealing. Starting with β=0\beta = 0 makes the model a plain autoencoder for the first epoch, so it learns to use the latent space while there is no penalty for doing so. Ramping β\beta to 1 over ten epochs then imposes the structure gradually. Without this ramp, a VAE can take the easy route early — zero out the latent, minimise the KL, and never recover, because by the time it might benefit from using the latent it has already learned to ignore it.

Logging the two terms separately. The total is nearly useless for diagnosis. The ratio between reconstruction and KL is what tells you whether the model is healthy.

Expect roughly this shape (per-image, summed over pixels; the total is reconstruction plus β\beta times KL):

EpochTotalReconstructionKLReading
1~149~146~34β=0.1\beta = 0.1; the latent is used freely while the penalty is cheap
5~95~80~31Penalty ramping; KL starts to come down
15~102~78~24Full penalty; settling
30~99~75~23.5Healthy converged VAE

Notice that the total rises between roughly epochs 5 and 10. That is not the model getting worse: β\beta is still climbing, so the same KL costs more each epoch. Once β\beta reaches 1 the total falls again. This is one more reason to read the two terms rather than the sum.

A final KL near zero means the model collapsed. A final KL above about 60 on MNIST means the latent is being used as a lookup table and random samples will be poor.

Diagnostic one: how many latent dimensions are alive

This is the check almost nobody runs and it is the most informative one available.

Python
@torch.no_grad()def active_units(model, loader, threshold=0.1):    model.eval()    kl_sum = torch.zeros(model.latent_dim, device=device)    n = 0    for x, _ in loader:        x = x.to(device)        mu, logvar = model.encode(x.view(-1, 784))        per_dim = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp())        kl_sum += per_dim.sum(0)        n += x.size(0)    kl = (kl_sum / n).cpu()    order = kl.argsort(descending=True)    print("per-dimension KL (nats):",          " ".join(f"{v:.2f}" for v in kl[order]))    print(f"active dimensions: {(kl > threshold).sum().item()} / {model.latent_dim}")    return klkl_per_dim = active_units(model, test_loader)

A dimension whose average KL is near zero has collapsed to the prior — the encoder puts every input at μ=0,σ=1\mu = 0, \sigma = 1 there, so it carries no information about the image and the decoder ignores it.

With this exact setup, expect most of the 20 dimensions to clear the 0.1-nat threshold: two test runs with different seeds gave 18 and 19. Read the printed values as well as the count. They are not equal: a handful of dimensions carry around 2 nats each, a tail carries half a nat or less, and at least one sits at zero. That spread is the model telling you which directions handwritten digits actually need. If all 20 are active with no near-zero tail, the latent may be too small and reconstruction capped — try 32 and see whether the extra dimensions stay idle. If only a few are active, something has gone wrong: check the sum-versus-mean bug first, then lengthen the annealing warm-up. (With reduction="mean", a test run ended with 0 of 20 active.)

Per-dimension KL is the closest thing a VAE has to an honest self-report. It tells you both whether the model is healthy and how large the latent space should have been.

Diagnostic two: reconstruction

Python
@torch.no_grad()def show_reconstructions(model, n=8):    model.eval()    x, _ = next(iter(test_loader))    x = x[:n].to(device)    recon, _, _ = model(x)    fig, axes = plt.subplots(2, n, figsize=(n * 1.2, 3))    for i in range(n):        axes[0, i].imshow(x[i].cpu().squeeze(), cmap="gray")        axes[1, i].imshow(recon[i].view(28, 28).cpu(), cmap="gray")        axes[0, i].axis("off"); axes[1, i].axis("off")    axes[0, 0].set_title("original", loc="left")    axes[1, 0].set_title("reconstruction", loc="left")    plt.tight_layout(); plt.savefig("recon.png", dpi=120)show_reconstructions(model)

Reconstructions should be clearly the same digit, slightly softer than the original. Softness is expected and is not a bug — it follows directly from the decoder having to satisfy a whole blob of codes at once, so it hedges towards the average of the plausible outputs.

If reconstructions are identical for every input, the model has collapsed. If they are the wrong digit, training has not finished.

Diagnostic three: the latent space

Train a second model with latent_dim=2 so the space can be plotted directly. Expect slightly worse reconstruction — two dimensions is genuinely restrictive — but the picture is worth it.

Python
@torch.no_grad()def plot_latent(model2d, loader):    model2d.eval()    zs, ys = [], []    for x, y in loader:        mu, _ = model2d.encode(x.to(device).view(-1, 784))        zs.append(mu.cpu()); ys.append(y)    z = torch.cat(zs); y = torch.cat(ys)    plt.figure(figsize=(7, 6))    sc = plt.scatter(z[:, 0], z[:, 1], c=y, cmap="tab10", s=3, alpha=0.6)    plt.colorbar(sc, label="digit"); plt.xlabel("z1"); plt.ylabel("z2")    plt.savefig("latent.png", dpi=120)

What you want to see: ten overlapping clouds, one per digit, roughly centred on the origin and spanning about [−3,3][-3, 3] in each direction. Visually similar digits sit adjacent — 4, 7 and 9 usually cluster together, as do 3, 5 and 8.

The overlap is the point, and it is what a plain autoencoder would not give you. Separated islands with empty space between them mean random sampling will land in the void and produce nothing. The KL term is what pushed the clusters together until they cover the region you are about to sample from.

Generating new digits

Python
@torch.no_grad()def sample(model, n=64):    model.eval()    z = torch.randn(n, model.latent_dim, device=device)   # straight from N(0, I)    imgs = model.dec(z).view(-1, 28, 28).cpu()    fig, axes = plt.subplots(8, 8, figsize=(8, 8))    for ax, im in zip(axes.flat, imgs):        ax.imshow(im, cmap="gray"); ax.axis("off")    plt.tight_layout(); plt.savefig("samples.png", dpi=120)sample(model)

Realistic expectations for a 30-epoch MNIST VAE: roughly 70 to 85 percent of the 64 samples are recognisable digits. The rest are ambiguous hybrids — something between a 4 and a 9, or a 3 with an extra loop. All of them are soft-edged.

That is a correct result, not a failed one. Sharp samples require machinery this model does not have, and the useful outcome here is that samples drawn from a plain standard normal decode into digits at all. That is the property the entire variational construction exists to produce.

Interpolation: the payoff

Python
@torch.no_grad()def interpolate(model, x_a, x_b, steps=10):    model.eval()    mu_a, _ = model.encode(x_a.view(1, 784).to(device))    mu_b, _ = model.encode(x_b.view(1, 784).to(device))    ts = torch.linspace(0, 1, steps, device=device).view(-1, 1)    z = (1 - ts) * mu_a + ts * mu_b            # walk the straight line    imgs = model.dec(z).view(-1, 28, 28).cpu()    fig, axes = plt.subplots(1, steps, figsize=(steps, 1.4))    for ax, im in zip(axes, imgs):        ax.imshow(im, cmap="gray"); ax.axis("off")    plt.tight_layout(); plt.savefig("interp.png", dpi=120)x, y = next(iter(test_loader))a = x[(y == 3).nonzero()[0].item()]b = x[(y == 8).nonzero()[0].item()]interpolate(model, a, b)

Every intermediate frame should be a plausible digit-like shape, morphing continuously from 3 to 8. This is the clearest possible demonstration that the latent space is continuous rather than a set of memorised points — run the same code on a plain autoencoder trained without the KL term and the middle frames are noise.

Troubleshooting

SymptomMost likely causeFix
All outputs are the same grey blurPosterior collapse, usually from reduction="mean"Use "sum"; lengthen the β\beta warm-up
Final KL below ~1SameSame
Loss becomes NaNLog-variance exploded, or pixels outside [0,1][0,1]logvar.clamp(-6, 2); remove any normalisation transform
Reconstruction good, random samples poorAggregate posterior does not match the priorTrain longer; raise β\beta slightly; check active-unit count
Reconstruction poor even at epoch 30Latent too small, or β\beta too highRaise latent_dim to 32; cap β\beta below 1
Interpolations jump abruptlyHoles in the latent spaceRaise β\beta; train longer
Device mismatch errorSampling z on CPU while the model is on GPUPass device=device to every tensor constructor

Extensions worth doing

Sweep β\beta. Train at β∈{0,0.5,1,4}\beta \in \{0, 0.5, 1, 4\} and compare reconstruction quality, sample quality and active-unit count. At β=0\beta = 0 you have a plain autoencoder: excellent reconstruction, useless samples. At β=4\beta = 4 reconstructions degrade while individual latent dimensions start to align with visible factors such as stroke thickness or slant. Feeling that trade-off directly is worth more than reading about it.

Make it conditional. Concatenate a one-hot digit label to both the encoder input and the decoder input. Now you can ask for a specific digit rather than taking what you are given — dec(cat([z, onehot(7)])) — and the latent is freed from having to encode identity, so it can spend its capacity on style.

Swap in convolutions. Replace the linear encoder with strided Conv2d layers and the decoder with ConvTranspose2d. Expect noticeably cleaner strokes, because convolution encodes the assumption that nearby pixels are related — an assumption the fully-connected version has to learn from scratch.

Detect anomalies. Feed the trained model images from a different dataset, such as Fashion-MNIST, and compare reconstruction error against MNIST test digits. The error distributions separate cleanly. That is out-of-distribution detection, and you got it free from a model you trained only to reconstruct.

What this build is really teaching

Three habits transfer to every generative model you will train after this one.

Never trust a single loss number. This model has two objectives pulling against each other, and their sum hides which one is winning. Every silent failure here — collapse, over-regularisation, an unused latent — is invisible in the total and obvious in the split. Log components separately from the very first run, before you need them.

Measure what the model uses, not just what it produces. The per-dimension KL told you how much of the 20-dimensional latent the model really uses, and how unevenly. That number is not in the loss, is not in the samples, and directly determines whether your architecture is oversized or undersized. The equivalent measurement exists in every architecture and is almost always worth finding.

Verify the property you actually wanted. You did not want low reconstruction error; a plain autoencoder wins on that and cannot generate anything. You wanted a continuous latent space where every point decodes to something sensible. Only the interpolation and sampling checks test that property. The loss curve never mentions it — and the general version of that observation is the one worth carrying: the metric you optimise is rarely the property you need, so test the property directly.