GANs
Be able to train a small GAN and explain the instability.
Prerequisites
- EGenerative models — an overviewrequired
Intuition
A GAN is two networks in a game:
- The generator G takes noise and produces an image. Its goal: fool the discriminator.
- The discriminator D is given an image and decides whether it is genuine or generated. Its goal: expose the generator.
They are trained alternately. If it works, G becomes so good that D cannot tell them apart — but the equilibrium is hard to reach.
Three classic problems:
- Mode collapse: G finds a handful of images that fool D and produces only those. The variety dies.
- Instability: if D gets too good, G gets no useful gradients and stops learning.
- No likelihood: there is no measure that says «how good is the model?» — so it is hard to know when to stop training.
Formal
The objective function (Goodfellow et al. 2014):
At the optimal this minimises the Jensen–Shannon divergence between the data distribution and the generator's distribution. The problem: when the distributions barely overlap (which they do not at the start) the JS divergence is constant and the gradient zero — which is why the «non-saturating» objective is used in practice.
Why WGAN helps: the Wasserstein distance gives useful gradients even without overlap, and the critic's value correlates with image quality — a measure to follow, which an ordinary GAN lacks. It requires a Lipschitz constraint (a gradient penalty).
Where GANs are still used: super-resolution, image-to-image, and anywhere one forward pass is required (real time). Diffusion won on quality and stability but needs many steps; GAN sampling is a single call.
Evaluation: FID (which compares statistics in a pretrained network's feature space) is the standard, but does not catch mode collapse well — complement it with precision/recall for generative models and human judgement.
Code
import torch, torch.nn as nn
G = nn.Sequential(nn.Linear(64, 256), nn.ReLU(), nn.Linear(256, 784), nn.Tanh())
D = nn.Sequential(nn.Linear(784, 256), nn.LeakyReLU(0.2), nn.Linear(256, 1))
optG = torch.optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999))
optD = torch.optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999))
bce = nn.BCEWithLogitsLoss()
for xb in loader:
# 1. The discriminator: genuine → 1, fake → 0
z = torch.randn(xb.size(0), 64)
fake = G(z).detach()
lossD = bce(D(xb), torch.ones(len(xb), 1)) + bce(D(fake), torch.zeros(len(xb), 1))
lossD.backward(); optD.step(); optD.zero_grad()
# 2. The generator: non-saturating — maximise log D(G(z))
z = torch.randn(xb.size(0), 64)
lossG = bce(D(G(z)), torch.ones(len(xb), 1))
lossG.backward(); optG.step(); optG.zero_grad()
# Diagnostics: if lossD → 0, D has won and G gets no gradients.
# If the generated images become more and more alike: mode collapse.
Mastery means
- Explains the generator, the discriminator and the minimax objective
- Recognises mode collapse and instability
- Knows why GANs have no likelihood
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Generative Adversarial Networks — arXiv (open access; licence per article)
- arXiv — Wasserstein GAN — arXiv (open access; licence per article)