Course Content
Introduction to Generative AI
3 sections · 9 lessons
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.
Setup
python -m venv .venvsource .venv/bin/activate # Windows: .venv\Scripts\activatepip install torch torchvision matplotlib numpyNothing else is needed. MNIST downloads automatically on first run — 60,000 training images of 28×28 greyscale digits, about 10 MB.
1import torch2import torch.nn as nn3import torch.nn.functional as F4from torch.utils.data import DataLoader5from torchvision import datasets, transforms6import matplotlib.pyplot as plt78device = ("cuda" if torch.cuda.is_available()9 else "mps" if torch.backends.mps.is_available()10 else "cpu")11torch.manual_seed(0)12print("device:", device)Data
1tf = transforms.ToTensor() # -> float tensor in [0, 1]23train_ds = datasets.MNIST("./data", train=True, download=True, transform=tf)4test_ds = datasets.MNIST("./data", train=False, download=True, transform=tf)56train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, drop_last=True)7test_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]. 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] — 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
1class VAE(nn.Module):2 def __init__(self, latent_dim=20):3 super().__init__()4 self.latent_dim = latent_dim56 self.enc = nn.Sequential(7 nn.Linear(784, 512), nn.ReLU(),8 nn.Linear(512, 256), nn.ReLU(),9 )10 self.fc_mu = nn.Linear(256, latent_dim)11 self.fc_logvar = nn.Linear(256, latent_dim)1213 self.dec = nn.Sequential(14 nn.Linear(latent_dim, 256), nn.ReLU(),15 nn.Linear(256, 512), nn.ReLU(),16 nn.Linear(512, 784), nn.Sigmoid(),17 )1819 def encode(self, x):20 h = self.enc(x)21 return self.fc_mu(h), self.fc_logvar(h)2223 def reparameterise(self, mu, logvar):24 if not self.training:25 return mu # deterministic when evaluating26 std = torch.exp(0.5 * logvar)27 return mu + std * torch.randn_like(std)2829 def forward(self, x):30 mu, logvar = self.encode(x.view(-1, 784))31 z = self.reparameterise(mu, logvar)32 return self.dec(z), mu, logvarThree 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] pixel range and pairs correctly with binary cross-entropy.
The loss
1def vae_loss(recon, x, mu, logvar, beta=1.0):2 x = x.view(-1, 784)3 # SUM over pixels and batch -- not mean. See the note below.4 recon_loss = F.binary_cross_entropy(recon, x, reduction="sum")5 kld = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())6 return recon_loss + beta * kld, recon_loss, kldreduction="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 and σ=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:
It is worth checking one value by hand so the code stops being magic. For a single dimension with μ=0.5 and logσ2=−0.3: σ2=e−0.3=0.741, so the bracket is 1−0.3−0.25−0.741=−0.291, giving DKL=0.146 nats for that dimension. A dimension sitting exactly at μ=0, σ2=1 contributes precisely zero — it has become the prior and carries no information.
Training with KL annealing
1model = VAE(latent_dim=20).to(device)2opt = torch.optim.Adam(model.parameters(), lr=1e-3)34EPOCHS, WARMUP = 30, 105history = []67for epoch in range(EPOCHS):8 beta = min(1.0, (epoch + 1) / WARMUP) # ramp the KL weight in9 model.train()10 tot = tot_r = tot_k = 0.01112 for x, _ in train_loader:13 x = x.to(device)14 recon, mu, logvar = model(x)15 loss, r, k = vae_loss(recon, x, mu, logvar, beta)1617 opt.zero_grad()18 loss.backward()19 torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)20 opt.step()2122 n = x.size(0)23 tot += loss.item() / n; tot_r += r.item() / n; tot_k += k.item() / n2425 nb = len(train_loader)26 history.append((tot / nb, tot_r / nb, tot_k / nb))27 print(f"epoch {epoch+1:2d} beta {beta:.2f} "28 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 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 β 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 β times KL):
| Epoch | Total | Reconstruction | KL | Reading |
|---|---|---|---|---|
| 1 | ~149 | ~146 | ~34 | β=0.1; the latent is used freely while the penalty is cheap |
| 5 | ~95 | ~80 | ~31 | Penalty ramping; KL starts to come down |
| 15 | ~102 | ~78 | ~24 | Full penalty; settling |
| 30 | ~99 | ~75 | ~23.5 | Healthy converged VAE |
Notice that the total rises between roughly epochs 5 and 10. That is not the model getting worse: β is still climbing, so the same KL costs more each epoch. Once β 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.
1@torch.no_grad()2def active_units(model, loader, threshold=0.1):3 model.eval()4 kl_sum = torch.zeros(model.latent_dim, device=device)5 n = 06 for x, _ in loader:7 x = x.to(device)8 mu, logvar = model.encode(x.view(-1, 784))9 per_dim = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp())10 kl_sum += per_dim.sum(0)11 n += x.size(0)12 kl = (kl_sum / n).cpu()13 order = kl.argsort(descending=True)14 print("per-dimension KL (nats):",15 " ".join(f"{v:.2f}" for v in kl[order]))16 print(f"active dimensions: {(kl > threshold).sum().item()} / {model.latent_dim}")17 return kl1819kl_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 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
1@torch.no_grad()2def show_reconstructions(model, n=8):3 model.eval()4 x, _ = next(iter(test_loader))5 x = x[:n].to(device)6 recon, _, _ = model(x)7 fig, axes = plt.subplots(2, n, figsize=(n * 1.2, 3))8 for i in range(n):9 axes[0, i].imshow(x[i].cpu().squeeze(), cmap="gray")10 axes[1, i].imshow(recon[i].view(28, 28).cpu(), cmap="gray")11 axes[0, i].axis("off"); axes[1, i].axis("off")12 axes[0, 0].set_title("original", loc="left")13 axes[1, 0].set_title("reconstruction", loc="left")14 plt.tight_layout(); plt.savefig("recon.png", dpi=120)1516show_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.
1@torch.no_grad()2def plot_latent(model2d, loader):3 model2d.eval()4 zs, ys = [], []5 for x, y in loader:6 mu, _ = model2d.encode(x.to(device).view(-1, 784))7 zs.append(mu.cpu()); ys.append(y)8 z = torch.cat(zs); y = torch.cat(ys)9 plt.figure(figsize=(7, 6))10 sc = plt.scatter(z[:, 0], z[:, 1], c=y, cmap="tab10", s=3, alpha=0.6)11 plt.colorbar(sc, label="digit"); plt.xlabel("z1"); plt.ylabel("z2")12 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] 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
1@torch.no_grad()2def sample(model, n=64):3 model.eval()4 z = torch.randn(n, model.latent_dim, device=device) # straight from N(0, I)5 imgs = model.dec(z).view(-1, 28, 28).cpu()6 fig, axes = plt.subplots(8, 8, figsize=(8, 8))7 for ax, im in zip(axes.flat, imgs):8 ax.imshow(im, cmap="gray"); ax.axis("off")9 plt.tight_layout(); plt.savefig("samples.png", dpi=120)1011sample(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
1@torch.no_grad()2def interpolate(model, x_a, x_b, steps=10):3 model.eval()4 mu_a, _ = model.encode(x_a.view(1, 784).to(device))5 mu_b, _ = model.encode(x_b.view(1, 784).to(device))6 ts = torch.linspace(0, 1, steps, device=device).view(-1, 1)7 z = (1 - ts) * mu_a + ts * mu_b # walk the straight line8 imgs = model.dec(z).view(-1, 28, 28).cpu()9 fig, axes = plt.subplots(1, steps, figsize=(steps, 1.4))10 for ax, im in zip(axes, imgs):11 ax.imshow(im, cmap="gray"); ax.axis("off")12 plt.tight_layout(); plt.savefig("interp.png", dpi=120)1314x, y = next(iter(test_loader))15a = x[(y == 3).nonzero()[0].item()]16b = x[(y == 8).nonzero()[0].item()]17interpolate(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
| Symptom | Most likely cause | Fix |
|---|---|---|
| All outputs are the same grey blur | Posterior collapse, usually from reduction="mean" | Use "sum"; lengthen the β warm-up |
| Final KL below ~1 | Same | Same |
| Loss becomes NaN | Log-variance exploded, or pixels outside [0,1] | logvar.clamp(-6, 2); remove any normalisation transform |
| Reconstruction good, random samples poor | Aggregate posterior does not match the prior | Train longer; raise β slightly; check active-unit count |
| Reconstruction poor even at epoch 30 | Latent too small, or β too high | Raise latent_dim to 32; cap β below 1 |
| Interpolations jump abruptly | Holes in the latent space | Raise β; train longer |
| Device mismatch error | Sampling z on CPU while the model is on GPU | Pass device=device to every tensor constructor |
Extensions worth doing
Sweep β. Train at β∈{0,0.5,1,4} and compare reconstruction quality, sample quality and active-unit count. At β=0 you have a plain autoencoder: excellent reconstruction, useless samples. At β=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.