Course Content
Introduction to Generative AI
3 sections · 9 lessons
Variational Autoencoders (VAEs)
Build a plain autoencoder on a face dataset. An encoder network squeezes each 64×64 image down to 32 numbers; a decoder expands those 32 numbers back into an image. Train with mean squared error between input and output. After an hour the reconstructions are excellent — you can recognise individuals. You have compressed 12,288 numbers into 32 with very little loss. Impressive.
Now generate a new face. Draw 32 random numbers, feed them to the decoder, look at the result.
It is garbage. Not a slightly odd face — coloured smears with no facial structure at all. Try a hundred more random draws and every one is garbage.
Try something gentler. Take the codes of two real faces, zA and zB, and decode the midpoint (zA+zB)/2. If the code space were meaningful this should look like a blend of two people. It does not. It looks like a third kind of garbage.
Why the plain autoencoder cannot generate
Nothing in the training objective ever asked it to. The loss said: compress and reconstruct. The encoder was free to put codes wherever it found convenient, and what is convenient is to scatter them — spread the training images as far apart as possible in code space, because distance makes them easy to tell apart and easy to reconstruct exactly.
Train a version with a 2-dimensional code so you can plot it, and you see the damage directly. The codes form tight islands separated by wide voids. Perhaps one cluster sits near (40,−18), another near (−95,60). Between and around them, nothing.
Two consequences follow.
- Random sampling fails because you have no idea where the islands are. Draw from a standard normal and you land near the origin, which may be empty ocean. The decoder has never seen a code from there and produces whatever its untrained extrapolation happens to give.
- Interpolation fails because the straight line between two islands passes over the void. Every point along it is a code the decoder was never trained on.
An autoencoder learns a lookup table with a compact index. A generative model needs a space where every point means something — and "every point" is a requirement nobody put in the loss.
The fix, stated before it is derived
A variational autoencoder changes two things.
The encoder outputs a cloud, not a point. For an input x it produces a mean vector μ and a spread vector σ, defining a small Gaussian blob in latent space. During training, the code actually passed to the decoder is a random draw from that blob. So the decoder must reconstruct x from any point in the blob, not just from one exact location. Blobs have volume; points do not. Volume is what fills the voids.
A penalty pulls every blob towards the origin. Specifically, towards a standard normal N(0,I). Without this the encoder would cheat by shrinking σ towards zero, turning blobs back into points and recovering the old behaviour. With it, the blobs are pushed together until they overlap and collectively cover the region around the origin — which is exactly the region you sample from at generation time.
Those two changes are the whole idea. Everything below explains where they come from and how to implement them without breaking gradient descent.
Where the objective comes from
The model you want to fit says: a latent code z is drawn from a simple prior p(z)=N(0,I), and then the observation is drawn from pθ(x∣z), the decoder. The probability of an image under this model is
which is intractable. To evaluate it you would have to integrate over every possible 32-dimensional code, and the vast majority contribute essentially nothing for any particular x. Monte Carlo estimation fails for the same reason: random draws from the prior almost never land in the tiny region that could have produced this specific image.
The variational move is to introduce an approximate posterior qϕ(z∣x) — the encoder — whose job is to guess which codes could have produced x. Then, for any choice of qϕ:
The final term is a KL divergence, which is never negative, so dropping it gives a lower bound — the evidence lower bound, or ELBO. Maximise it and you push up logpθ(x) from below. The gap between the two is exactly how badly the encoder approximates the true posterior, so improving the encoder tightens the bound. One objective improves both networks.
Reading the two terms
Reconstruction: Eq[logpθ(x∣z)]. Encode x, sample a z from its blob, decode, and ask how probable the original was. With a Gaussian decoder this is squared error; with a Bernoulli decoder — appropriate for pixel values in [0,1] — it is binary cross-entropy. Sum over pixels; never average, or the term gets scaled down relative to the KL and the balance breaks.
Regularisation: DKL(qϕ(z∣x)∥p(z)). How far this input's blob has drifted from the standard normal. Because both are diagonal Gaussians there is a closed form — no sampling needed:
Work it with real numbers. Take a 2-dimensional latent where the encoder outputs μ=[0.5,−1.2] and logσ2=[−0.3,0.4].
| Dimension | μj | logσj2 | σj2 | 1+logσj2−μj2−σj2 |
|---|---|---|---|---|
| 1 | 0.5 | -0.3 | 0.741 | 1−0.3−0.25−0.741=−0.291 |
| 2 | -1.2 | 0.4 | 1.492 | 1+0.4−1.44−1.492=−1.532 |
Sum is −1.823, so DKL=−0.5×(−1.823)=0.911 nats. Sanity check the direction: dimension 2 contributes far more penalty than dimension 1, because its mean sits further from zero and its variance is 1.49 rather than 1. Both deviations cost. A blob with μ=0 and σ2=1 contributes exactly zero — it is already the prior.
Note that the network outputs logσ2, not σ. This is deliberate: logσ2 can take any real value, whereas σ must stay positive. Predicting the log removes the constraint and avoids ever taking the log of a number that drifted to zero.
The reparameterisation trick
There is a problem sitting in the middle of the ELBO, and it is fatal if unaddressed. The forward pass samples z∼N(μ,σ2). Sampling is not differentiable. You cannot ask "how would the loss change if μ moved slightly?" when the step between μ and z is a random draw. Gradients reach the decoder and stop dead. The encoder never learns.
The fix is to move the randomness out of the path. Instead of drawing z from a distribution parameterised by μ and σ, draw a fixed standard normal ε and construct z arithmetically:
The distribution of z is identical. But now z is a deterministic function of μ, σ and an external constant, so ∂z/∂μ=1 and ∂z/∂σ=ε. Gradients flow straight through. The randomness is still there; it just entered from the side, as an input rather than as an operation.
Continuing the numbers above, with a draw of ε=[0.42,−1.30]:
- σ1=0.741=0.861, so z1=0.5+0.861×0.42=0.862
- σ2=1.492=1.221, so z2=−1.2+1.221×(−1.30)=−2.788
A different ε next epoch gives a different z for the same image — which is precisely the point. The decoder is forced to handle the whole blob.
The reparameterisation trick is not a numerical convenience. Without it the encoder receives no gradient at all, and the VAE does not train.
Implementation
1import torch2import torch.nn as nn3import torch.nn.functional as F45class VAE(nn.Module):6 def __init__(self, input_dim=784, hidden=512, latent_dim=20):7 super().__init__()8 self.encoder = nn.Sequential(9 nn.Linear(input_dim, hidden), nn.ReLU(),10 nn.Linear(hidden, 256), nn.ReLU(),11 )12 # Two heads on the same trunk: one for mu, one for log-variance13 self.fc_mu = nn.Linear(256, latent_dim)14 self.fc_logvar = nn.Linear(256, latent_dim)1516 self.decoder = nn.Sequential(17 nn.Linear(latent_dim, 256), nn.ReLU(),18 nn.Linear(256, hidden), nn.ReLU(),19 nn.Linear(hidden, input_dim),20 nn.Sigmoid(), # pixels in [0, 1]21 )2223 def encode(self, x):24 h = self.encoder(x)25 return self.fc_mu(h), self.fc_logvar(h)2627 def reparameterise(self, mu, logvar):28 if not self.training:29 return mu # deterministic at eval time30 std = torch.exp(0.5 * logvar) # sigma from log-variance31 eps = torch.randn_like(std)32 return mu + std * eps3334 def forward(self, x):35 mu, logvar = self.encode(x)36 z = self.reparameterise(mu, logvar)37 return self.decoder(z), mu, logvar1def vae_loss(recon_x, x, mu, logvar, beta=1.0):2 # SUM over pixels, not mean -- otherwise the two terms are on3 # different scales and the KL silently dominates.4 recon = F.binary_cross_entropy(recon_x, x, reduction='sum')5 kld = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())6 return recon + beta * kld, recon, kldThe reduction='sum' detail causes more silent failures than any other line. With 'mean', the reconstruction term is divided by every pixel of every image in the batch (784 × 128 with a batch of 128) while the KL is still summed, so the KL overwhelms everything, the encoder outputs μ=0,σ=1 for every input, and the model reconstructs the dataset average for all inputs. The loss curve looks fine. The output is a grey smudge.
Sampling afterwards is trivial, because the whole point of the KL term was to make the prior the right thing to sample from:
1model.eval()2with torch.no_grad():3 z = torch.randn(64, latent_dim) # straight from N(0, I)4 images = model.decoder(z).view(-1, 1, 28, 28)Reading the training curves
Watch the two loss components separately. Their ratio tells you more than the total ever will.
| Pattern | Diagnosis | Action |
|---|---|---|
| KL falls to near 0 and stays there | Posterior collapse — the latent carries no information | Anneal β from 0; add free bits; weaken the decoder |
| KL climbs without bound | Reconstruction term dominating; latent is being used as a lookup table | Increase β; check the sum/mean bug |
| Both fall, samples still blurry | Normal VAE behaviour | Expected — see below |
| Reconstruction good, random samples bad | Aggregate posterior does not match the prior | Longer training, larger β, or a richer prior |
| Loss becomes NaN | logσ2 exploded, or the decoder output hit exactly 0 or 1 | Clamp logvar to [−6,2]; use a numerically stable BCE |
The three failures that matter
Blur is structural, not a bug
VAE samples are soft, and no amount of training fixes it. The cause is in the objective. The decoder must reconstruct x from any z in the blob, and when several plausible sharp outputs are consistent with one code, the loss-minimising answer is their average. Averaging several sharp edges at slightly different positions produces one soft edge.
Squared error makes this worse than it needs to be, because it treats a sharp edge placed two pixels off as a large error while treating a uniform smear as a moderate one. Options that help: a perceptual loss computed in the feature space of a pretrained network rather than pixel space; a discriminator on top of the reconstruction, which reintroduces sharpness; or discrete latents, which sidestep the averaging.
Posterior collapse
The most common training failure. The KL term is minimised — driven to zero — when qϕ(z∣x)=N(0,I) for every input, meaning the encoder ignores the input entirely. If the decoder is powerful enough to produce decent output without using z, that is a genuinely good solution to the loss and the optimiser will find it. You end up with an expensive way to generate the dataset mean.
Three fixes, in order of how often they are needed:
- KL annealing. Start with β=0 and ramp it to 1 over the first several epochs. The model learns to use the latent while the penalty is cheap, and by the time the penalty arrives, the latent is load-bearing.
- Free bits. Apply no penalty until each latent dimension's KL exceeds a small threshold, typically 0.05–0.5 nats:
kld = torch.clamp(kld_per_dim, min=lambda_).sum(). This gives every dimension a free allowance of information. - Weaken the decoder. If the decoder is autoregressive or otherwise strong enough to model the data alone, it will. Reduce its capacity, or drop out parts of its conditioning.
A latent space that is not organised
If interpolations jump abruptly rather than morphing smoothly, the latent has holes. Usually the KL weight is too low, or the latent dimension is far larger than the data needs, leaving many unused directions that the decoder never learned to interpret.
Variants worth knowing
| Variant | Change | Effect | Use when |
|---|---|---|---|
| β-VAE | Weight the KL term by β>1 | More disentangled dimensions; worse reconstruction | You want interpretable latent factors |
| Conditional VAE | Feed a label y to both encoder and decoder | Generate a chosen class on demand | You need controllable output |
| VQ-VAE | Snap each latent to the nearest entry in a learned codebook | Discrete codes, sharp output, no posterior collapse | Feeding a transformer; high-fidelity images or audio |
| Hierarchical VAE | Latents at several resolutions | Much sharper samples; heavier to train | Quality matters more than simplicity |
The β dial is the clearest way to feel the trade-off. At β=0 you have a plain autoencoder — perfect reconstruction, unusable for generation. At β=1 you have the proper ELBO. Push to β=4 and reconstructions degrade while individual latent dimensions start aligning with human-recognisable factors such as rotation or thickness. There is no free lunch here; you are trading fidelity for structure.
VQ-VAE deserves particular attention because it powers a great deal of modern work. Replacing the continuous Gaussian latent with a discrete codebook lookup removes the averaging that causes blur, and because the latent is now a grid of integer codes, a transformer can model the distribution over those codes directly. The VAE becomes a compressor and a different model handles the generation.
What this means when you build something
Pick the latent dimension by measuring, not guessing. Too small and reconstruction is capped no matter how long you train. Too large and dimensions collapse to the prior and go unused. The practical diagnostic: after training, compute the per-dimension KL and count how many exceed roughly 0.1 nats. If 12 of your 64 dimensions are active, the model told you it needs about 12 — and the other 52 are costing compute for nothing. For MNIST-scale data, 10–20 is typical; for 64×64 faces, 64–256.
Log the KL and reconstruction terms separately from the first run. The total loss hides every failure described above. A collapsed model and a healthy model can post similar totals while behaving completely differently.
Choose a VAE for what it is uniquely good at. If you want the sharpest possible images, this is not the tool. What a VAE gives you that sharper methods do not is a fast, well-behaved encoder — a single forward pass that maps any input to a meaningful position in a continuous space. That is what makes it the right choice for representation learning, anomaly detection by reconstruction error, semantic interpolation, and as the compression stage inside a larger system. The generative capability is real but modest; the latent space is the product.