Introduction to Generative AI

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, zAz_A and zBz_B, and decode the midpoint (zA+zB)/2(z_A + z_B)/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 a VAE can generate and an autoencoder cannotEncoder reads the imageOutputs a mean and a log-variance, not a pointReparameterise: z equals mean plus sigma times noiseDecoder reconstructs from the sampled zLoss: reconstruction plus KL to a standard normal
The KL term is what makes the space between encoded points decodable, and that space is where new faces come from.

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)(40, -18), another near (−95,60)(-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 xx it produces a mean vector μ\mu and a spread vector σ\sigma, 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 xx 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)\mathcal{N}(0, I). Without this the encoder would cheat by shrinking σ\sigma 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 zz is drawn from a simple prior p(z)=N(0,I)p(z) = \mathcal{N}(0, I), and then the observation is drawn from pθ(x∣z)p_\theta(x \mid z), the decoder. The probability of an image under this model is

pθ(x)=∫pθ(x∣z) p(z) dzp_\theta(x) = \int p_\theta(x \mid z)\, p(z)\, dz

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 xx. 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)q_\phi(z \mid x) — the encoder — whose job is to guess which codes could have produced xx. Then, for any choice of qϕq_\phi:

log⁡pθ(x)=Eqϕ(z∣x)[log⁡pθ(x∣z)]−DKL(qϕ(z∣x) ∥ p(z))⏟ELBO  +  DKL(qϕ(z∣x) ∥ pθ(z∣x))⏟≥ 0\log p_\theta(x) = \underbrace{\mathbb{E}_{q_\phi(z \mid x)}\big[\log p_\theta(x \mid z)\big] - D_{KL}\big(q_\phi(z \mid x) \,\|\, p(z)\big)}_{\text{ELBO}} \;+\; \underbrace{D_{KL}\big(q_\phi(z \mid x) \,\|\, p_\theta(z \mid x)\big)}_{\ge\, 0}

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 log⁡pθ(x)\log p_\theta(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[log⁡pθ(x∣z)]\mathbb{E}_{q}[\log p_\theta(x \mid z)]. Encode xx, sample a zz 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][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))D_{KL}(q_\phi(z \mid 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:

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

Work it with real numbers. Take a 2-dimensional latent where the encoder outputs μ=[0.5, −1.2]\mu = [0.5,\, -1.2] and log⁡σ2=[−0.3, 0.4]\log \sigma^2 = [-0.3,\, 0.4].

Dimensionμj\mu_jlog⁡σj2\log \sigma_j^2σj2\sigma_j^21+log⁡σj2−μj2−σj21 + \log\sigma_j^2 - \mu_j^2 - \sigma_j^2
10.5-0.30.7411−0.3−0.25−0.741=−0.2911 - 0.3 - 0.25 - 0.741 = -0.291
2-1.20.41.4921+0.4−1.44−1.492=−1.5321 + 0.4 - 1.44 - 1.492 = -1.532

Sum is −1.823-1.823, so DKL=−0.5×(−1.823)=0.911D_{KL} = -0.5 \times (-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\mu = 0 and σ2=1\sigma^2 = 1 contributes exactly zero — it is already the prior.

Note that the network outputs log⁡σ2\log \sigma^2, not σ\sigma. This is deliberate: log⁡σ2\log \sigma^2 can take any real value, whereas σ\sigma 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)z \sim \mathcal{N}(\mu, \sigma^2). Sampling is not differentiable. You cannot ask "how would the loss change if μ\mu moved slightly?" when the step between μ\mu and zz 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 zz from a distribution parameterised by μ\mu and σ\sigma, draw a fixed standard normal ε\varepsilon and construct zz arithmetically:

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

The distribution of zz is identical. But now zz is a deterministic function of μ\mu, σ\sigma and an external constant, so ∂z/∂μ=1\partial z / \partial \mu = 1 and ∂z/∂σ=ε\partial z / \partial \sigma = \varepsilon. 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]\varepsilon = [0.42,\, -1.30]:

  • σ1=0.741=0.861\sigma_1 = \sqrt{0.741} = 0.861, so z1=0.5+0.861×0.42=0.862z_1 = 0.5 + 0.861 \times 0.42 = 0.862
  • σ2=1.492=1.221\sigma_2 = \sqrt{1.492} = 1.221, so z2=−1.2+1.221×(−1.30)=−2.788z_2 = -1.2 + 1.221 \times (-1.30) = -2.788

A different ε\varepsilon next epoch gives a different zz 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

Python
import torchimport torch.nn as nnimport torch.nn.functional as Fclass VAE(nn.Module):    def __init__(self, input_dim=784, hidden=512, latent_dim=20):        super().__init__()        self.encoder = nn.Sequential(            nn.Linear(input_dim, hidden), nn.ReLU(),            nn.Linear(hidden, 256),       nn.ReLU(),        )        # Two heads on the same trunk: one for mu, one for log-variance        self.fc_mu     = nn.Linear(256, latent_dim)        self.fc_logvar = nn.Linear(256, latent_dim)        self.decoder = nn.Sequential(            nn.Linear(latent_dim, 256), nn.ReLU(),            nn.Linear(256, hidden),     nn.ReLU(),            nn.Linear(hidden, input_dim),            nn.Sigmoid(),               # pixels in [0, 1]        )    def encode(self, x):        h = self.encoder(x)        return self.fc_mu(h), self.fc_logvar(h)    def reparameterise(self, mu, logvar):        if not self.training:            return mu                      # deterministic at eval time        std = torch.exp(0.5 * logvar)      # sigma from log-variance        eps = torch.randn_like(std)        return mu + std * eps    def forward(self, x):        mu, logvar = self.encode(x)        z = self.reparameterise(mu, logvar)        return self.decoder(z), mu, logvar
Python
def vae_loss(recon_x, x, mu, logvar, beta=1.0):    # SUM over pixels, not mean -- otherwise the two terms are on    # different scales and the KL silently dominates.    recon = F.binary_cross_entropy(recon_x, x, reduction='sum')    kld   = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())    return recon + beta * kld, recon, kld

The 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\mu = 0, \sigma = 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:

Python
model.eval()with torch.no_grad():    z = torch.randn(64, latent_dim)       # straight from N(0, I)    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.

PatternDiagnosisAction
KL falls to near 0 and stays therePosterior collapse — the latent carries no informationAnneal β\beta from 0; add free bits; weaken the decoder
KL climbs without boundReconstruction term dominating; latent is being used as a lookup tableIncrease β\beta; check the sum/mean bug
Both fall, samples still blurryNormal VAE behaviourExpected — see below
Reconstruction good, random samples badAggregate posterior does not match the priorLonger training, larger β\beta, or a richer prior
Loss becomes NaNlog⁡σ2\log\sigma^2 exploded, or the decoder output hit exactly 0 or 1Clamp logvar to [−6,2][-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 xx from any zz 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)q_\phi(z\mid x) = \mathcal{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 zz, 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\beta = 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

VariantChangeEffectUse when
β\beta-VAEWeight the KL term by β>1\beta > 1More disentangled dimensions; worse reconstructionYou want interpretable latent factors
Conditional VAEFeed a label yy to both encoder and decoderGenerate a chosen class on demandYou need controllable output
VQ-VAESnap each latent to the nearest entry in a learned codebookDiscrete codes, sharp output, no posterior collapseFeeding a transformer; high-fidelity images or audio
Hierarchical VAELatents at several resolutionsMuch sharper samples; heavier to trainQuality matters more than simplicity

The β\beta dial is the clearest way to feel the trade-off. At β=0\beta = 0 you have a plain autoencoder — perfect reconstruction, unusable for generation. At β=1\beta = 1 you have the proper ELBO. Push to β=4\beta = 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.