Course Content
Introduction to Generative AI
3 sections · 9 lessons
Generative Adversarial Networks (GANs)
You have a network that turns 100 random numbers into a 64×64 image. Now write the loss function that makes it produce realistic faces.
Stop and actually try. This is harder than it looks, and the difficulty is the entire reason GANs exist.
Squared error against what? The generator produced a face from noise. There is no ground-truth image it was supposed to match — that is the whole point of generation. Pairing each output with a random real face and minimising pixel error just drives every output towards the dataset average, a grey oval.
Hand-written realism metrics? You could reward high-frequency detail to punish blur. The generator will find that adding fine-grained noise scores brilliantly, and you get static. You could add a symmetry term, a skin-tone histogram term, an edge-density term. Every one of them gets gamed, because the generator optimises exactly what you wrote and nothing else, and none of your terms says "face".
Ask a human? Correct, and useless — gradient descent needs millions of evaluations per training run.
The problem is that "looks real" is not a formula anyone can write. So do not write it. Train a second network to be the loss function. That network's job is to distinguish real images from generated ones — a plain binary classification problem, which neural networks are excellent at. Whatever it uses to tell them apart is, by construction, a difference between your output and reality. Its gradient points at that difference. Follow it.
A GAN's real innovation is not two networks competing. It is that the loss function is learned rather than specified, so it adapts to whatever flaw the generator has right now.
The game
Two networks, opposite objectives.
- Generator G: takes noise z∼N(0,I), outputs an image G(z). Wants the discriminator to call its output real.
- Discriminator D: takes an image, outputs a probability that it came from the real dataset. Wants to be right.
Written as one objective:
Read it one term at a time. The first term is large when D assigns high probability to real images. The second is large when D assigns low probability to generated ones. D maximises both. G appears only in the second term and pushes it the other way — it wants D(G(z)) near 1.
There is a genuinely elegant result underneath. Hold G fixed and the optimal discriminator is
Substitute that back in and the generator's objective becomes the Jensen–Shannon divergence between the real and generated distributions, minimised exactly when pg=pdata. So at the theoretical optimum the generator has matched the data distribution and the discriminator, unable to do better than guessing, outputs 0.5 everywhere.
That is the theory. In practice you never reach it, for reasons that occupy the rest of this lesson.
Why the original loss stalls, and the one-line fix
Early in training the generator is terrible and the discriminator spots it instantly, so D(G(z))≈0.01. Look at what the generator's gradient does there.
Its loss is log(1−D(G(z))), and
At D=0.01 that is −1.01. The generator is failing as badly as it possibly can and receives a gradient of magnitude one. The function log(1−D) is nearly flat near D=0 — it has saturated. The worse the generator does, the less it learns.
The fix is to have the generator maximise logD(G(z)) instead of minimising log(1−D(G(z))). Same direction of preference, completely different gradient:
At D=0.01 that is −100. A hundred times stronger, precisely when the generator most needs a strong signal.
| D(G(z)) | Generator status | Gradient, saturating form | Gradient, non-saturating form |
|---|---|---|---|
| 0.01 | Easily caught | 1.01 | 100 |
| 0.1 | Poor | 1.11 | 10 |
| 0.5 | Fooling D half the time | 2.0 | 2.0 |
| 0.9 | Winning | 10 | 1.11 |
The two forms agree at 0.5 and behave oppositely at the extremes. The non-saturating version gives the loudest signal when the generator is worst, which is what you want; the original gives it when the generator is already winning, which is useless. Every practical GAN uses the non-saturating loss. It is one line of code and the difference between training and not training.
Loss variants
| Loss | Discriminator output | Property | Cost |
|---|---|---|---|
| Non-saturating (standard) | Probability via sigmoid | Fixes the early-training stall | Still unstable mid-training |
| Least-squares (LSGAN) | Unbounded score | Penalises samples far from the boundary even when correctly classified; smoother gradients | Marginal quality gain |
| Wasserstein (WGAN) | Unbounded "critic" score | Loss correlates with sample quality, so the curve is finally informative | Requires a Lipschitz constraint |
| WGAN-GP | Unbounded score | Enforces the constraint with a gradient penalty rather than weight clipping | Extra backward pass per step |
| Hinge | Unbounded score | Simple, stable, standard in large image GANs | None notable |
The Wasserstein family is worth understanding for one reason beyond stability. With the standard loss, the discriminator's loss tells you nothing about sample quality — it measures how the contest is going, not how good the images are, so a falling loss can accompany worsening output. The Wasserstein critic's value approximates an actual distance between the distributions, so it goes down as samples improve. That turns a blind training run into a monitored one.
Architecture
The generator maps a low-dimensional vector to a full image, expanding spatial size while contracting channels. The discriminator does the reverse. To keep the code short, this version makes 32×32 colour images; one more stride-2 layer in each network gives the 64×64 of the opening example.
1import torch.nn as nn23class Generator(nn.Module):4 def __init__(self, z_dim=100, ch=64):5 super().__init__()6 self.net = nn.Sequential(7 # z: (B, 100, 1, 1) -> 4x48 nn.ConvTranspose2d(z_dim, ch*8, 4, 1, 0, bias=False),9 nn.BatchNorm2d(ch*8), nn.ReLU(True),10 # 4x4 -> 8x811 nn.ConvTranspose2d(ch*8, ch*4, 4, 2, 1, bias=False),12 nn.BatchNorm2d(ch*4), nn.ReLU(True),13 # 8x8 -> 16x1614 nn.ConvTranspose2d(ch*4, ch*2, 4, 2, 1, bias=False),15 nn.BatchNorm2d(ch*2), nn.ReLU(True),16 # 16x16 -> 32x3217 nn.ConvTranspose2d(ch*2, 3, 4, 2, 1, bias=False),18 nn.Tanh(), # outputs in [-1, 1]19 )20 def forward(self, z):21 return self.net(z)2223class Discriminator(nn.Module):24 def __init__(self, ch=64):25 super().__init__()26 self.net = nn.Sequential(27 nn.Conv2d(3, ch, 4, 2, 1, bias=False),28 nn.LeakyReLU(0.2, inplace=True), # no norm on the first layer29 nn.Conv2d(ch, ch*2, 4, 2, 1, bias=False),30 nn.BatchNorm2d(ch*2), nn.LeakyReLU(0.2, inplace=True),31 nn.Conv2d(ch*2, ch*4, 4, 2, 1, bias=False),32 nn.BatchNorm2d(ch*4), nn.LeakyReLU(0.2, inplace=True),33 nn.Conv2d(ch*4, 1, 4, 1, 0, bias=False), # single score34 )35 def forward(self, x):36 return self.net(x).view(-1)Several choices in there are not stylistic:
Tanhon the generator output, images scaled to [−1,1]. The ranges must match or the discriminator separates real from fake on brightness alone and learns nothing useful.LeakyReLUin the discriminator. Plain ReLU zeroes the gradient for every negative activation, and the generator's only learning signal passes back through this network. Dead units in the discriminator mean no signal for the generator.- No normalisation on the discriminator's first layer. It would erase the input statistics that distinguish real from generated data.
- Kernel 4, stride 2, padding 1. This combination halves or doubles resolution exactly, avoiding the checkerboard artefacts that appear when kernel size is not divisible by stride.
The training loop
1import torch, torch.nn.functional as F23opt_d = torch.optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999))4opt_g = torch.optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999))56for real in loader:7 b = real.size(0)89 # ---- discriminator ----10 z = torch.randn(b, z_dim, 1, 1, device=dev)11 fake = G(z)12 d_real = D(real)13 d_fake = D(fake.detach()) # detach: no generator update here14 loss_d = (F.binary_cross_entropy_with_logits(d_real, torch.full_like(d_real, 0.9))15 + F.binary_cross_entropy_with_logits(d_fake, torch.zeros_like(d_fake)))16 opt_d.zero_grad(); loss_d.backward(); opt_d.step()1718 # ---- generator (non-saturating) ----19 d_fake = D(fake) # re-score, gradients flow to G20 loss_g = F.binary_cross_entropy_with_logits(d_fake, torch.ones_like(d_fake))21 opt_g.zero_grad(); loss_g.backward(); opt_g.step()Two lines carry outsized weight. fake.detach() stops the discriminator's update from also modifying the generator — omit it and you are training the generator to help the discriminator catch it, which is exactly backwards. And the target of 0.9 instead of 1.0 is one-sided label smoothing: it stops the discriminator becoming absolutely certain, which keeps its gradients from vanishing. The betas of (0.5,0.999) rather than the usual (0.9,0.999) reduce momentum, because in a two-player game the landscape shifts under you and stale momentum overshoots.
The two failures
Mode collapse
Train on a dataset of ten digit classes and find the generator producing only 3s and 8s. Every sample is sharp. Every sample is convincing. Coverage is 20%.
The cause is directly visible in the objective. The generator is scored on whether each individual sample fools the discriminator. Nothing in V(D,G) asks whether the samples, taken together, resemble the dataset. If the generator finds one output the discriminator cannot reject, producing that output every time is an optimal strategy. It has won the game as written.
What follows is a chase: the discriminator eventually learns "too many 3s here", so the generator jumps to producing only 7s, then only 1s, cycling forever without ever covering everything at once.
| Remedy | Mechanism |
|---|---|
| Minibatch discrimination | Let D see relationships within a batch, so it can reject a batch that lacks variety |
| Unrolled GAN | G optimises against D's future response, removing the payoff for a short-lived jump |
| WGAN-GP | The Wasserstein distance penalises missing mass, unlike Jensen–Shannon |
| Two discriminators, or one with more capacity | Harder to fool with a single trick |
| Experience replay of past fakes | Stops the generator revisiting a mode D has already forgotten |
Instability
There is no loss being minimised. There is an equilibrium between two networks that both keep moving, and equilibria can be orbited rather than reached. The characteristic symptom is a run that improves for 40 epochs and then, over two epochs, degenerates into noise.
The usual proximate cause is an imbalance. If the discriminator becomes far stronger than the generator, it rejects everything with near-total confidence, D(G(z))→0, and even the non-saturating gradient becomes a signal pointing everywhere at once. If the generator gets far ahead, the discriminator stops being an informative critic and the generator drifts.
Practical levers, roughly in the order to try them: spectral normalisation on the discriminator, which is the single most reliable stabiliser available; a lower discriminator learning rate; different update ratios such as two discriminator steps per generator step; gradient penalty; and keeping an exponential moving average of the generator's weights for sampling, which smooths out the oscillation even when training itself is bumpy.
Save checkpoints often and evaluate them. GAN training does not converge to a best model at the end — the best model is usually somewhere in the middle of the run.
Evaluating output
You cannot use the loss. A GAN's loss measures the state of a contest, not the quality of images. Falling generator loss is compatible with worsening samples.
| Metric | Measures | Direction | Limitation |
|---|---|---|---|
| Inception Score | Confident classification plus class variety | Higher better | Blind to mode collapse within a class; ignores the real data entirely |
| Fréchet Inception Distance | Distance between real and generated feature distributions | Lower better | Needs 10k+ samples; assumes Gaussian features |
| Precision / Recall for generative models | Separates realism from coverage | Both higher | More work to compute |
| Human preference study | What you actually care about | — | Slow and expensive |
FID is the default. Its virtue is catching mode collapse — a generator producing beautiful images of one class has generated features clustered far too tightly, and the distance to the real distribution rises even as individual samples look flawless. Report the sample count alongside the number; FID computed on 1,000 samples is not comparable to FID on 50,000.
GAN or VAE
Both map noise to data in a single forward pass. The difference is what they optimise, and every practical distinction follows from that.
| GAN | VAE | |
|---|---|---|
| Objective | Fool a learned critic | Maximise a lower bound on likelihood |
| Sample sharpness | High — nothing averages | Blurry — averaging is built in |
| Mode coverage | Unreliable | Good; the likelihood term punishes ignoring data |
| Training | Can fail outright | Essentially always converges |
| Encoder for real inputs | None by default | Yes, free |
| Can score how likely an input is | No | Approximately |
| Loss curve is informative | No, unless Wasserstein | Yes |
Sharpness versus coverage, and stability versus quality. Those are the trades.
What this means when you build something
Before starting a GAN project, price in the failure mode that has no equivalent elsewhere: a full training run can end with nothing usable. Not a mediocre model — a collapsed one. If your schedule cannot absorb that, use an architecture with a monotone loss.
If you do proceed, three decisions do most of the work. Use a known-good recipe rather than designing from scratch; the published architectures encode years of hard-won detail about normalisation placement and kernel sizes. Add spectral normalisation to the discriminator from the start — it costs almost nothing and prevents the most common divergence. Compute FID every few epochs on a fixed noise batch and keep the best checkpoint, because the final checkpoint is frequently not the best one.
Where GANs remain genuinely the right answer is real-time generation. One forward pass, milliseconds, on modest hardware. For interactive editing, live style transfer, super-resolution in a video pipeline, or anything running on a device, that single-pass property is decisive — and it is why the adversarial idea survives inside modern systems even where the pure GAN has been displaced, most often as a component that sharpens or accelerates a model trained some other way.