primer.ml.generative.gans

GANs: a forger against a detective

Run: python -m primer.ml.generative.gans

This lesson builds on the small networks and gradients of primer.ml.neural_net, the binary cross-entropy of primer.ml.losses and the Adam optimizer of primer.ml.optimizers; primer.notation explains every symbol from zero.

Level 1: The practitioner's guide

In one sentence. A generative adversarial network (GAN) trains a generator to turn random noise into samples by pitting it against a discriminator that learns to tell its samples from real ones, which yields sharp output in a single pass at the price of a training game that is hard to keep stable.

When you need it. Two tells. You need a generator that runs in one pass, because it sits in a real-time loop (a speech vocoder, an interactive tool, a game) and a many-step diffusion sampler is too slow. Or you have a decoder that rebuilds images or audio and its output is blurry: a reconstruction loss rewards averages, and an adversarial loss is the standard cure. You don't need to train a GAN from scratch to generate new images or video today: since 2021 diffusion and flow models are the default there, because they train stably and cover the data, which is exactly what this lesson shows GANs struggling to do. The number that shows the naive approach failing: on this lesson's toy data (eight small clouds on a ring), a GAN trained with the same learning rate for both players covers 3 of the 8 clouds after 2,500 rounds, and only 16% of its samples land on any cloud. Change nothing but the discriminator's learning rate to three times the generator's and it covers all 8 with 81% on target.

Your options. From the least commitment to the most:

Option What it does What it guarantees What it costs Where it lives
A diffusion or flow model instead Generates by removing noise over many passes Stable training, full coverage; matches BigGAN-deep's image quality with 25 passes and better coverage Tens of network passes per sample A hosted API or a diffusion library
A pretrained GAN generator One pass from noise to sample, with an editable latent space Real-time generation in its domain (faces, one class of object) You are limited to domains someone trained; less variety than diffusion Released checkpoints (StyleGAN family)
Train a GAN with the standard stabilisers Non-saturating loss, a faster discriminator, an R1 penalty or spectral normalization A one-pass generator for your own narrow domain Two networks, coverage monitoring, learning-rate tuning; a run can still collapse Your training loop
Wasserstein critic (WGAN-GP) Replaces the verdict with a distance that still points home when real and fake don't overlap A gradient wherever the generator is, and a loss that tracks quality Several critic steps per generator step, plus a gradient penalty Your training loop
Adversarial loss inside a decoder A discriminator judges reconstructions against originals Sharp detail where a reconstruction loss alone gives blur One more network to train and keep stable The training of VAEs, tokenizers and vocoders
Adversarial distillation of a diffusion model A student learns to match a diffusion teacher in 1 to 4 steps, with a discriminator keeping it sharp Real-time sampling from a foundation model A teacher, a distillation run, some loss of variety The fast sampling path of a diffusion system

How to choose. Start from the latency you need and the variety you can't lose.

  • A new image, audio or video generator with no hard latency limit: use diffusion or flow matching, and put adversarial losses only in the decoder and in a distilled fast path.
  • One-pass generation is non-negotiable and the domain is narrow: a GAN generator, trained with every stabiliser in the table, or a distilled diffusion model if a teacher exists for your domain.
  • A blurry decoder: add a discriminator that compares reconstructions with originals, and expect the same instabilities as any GAN.
  • A latent space to edit (age a face, change a pose): a style-based generator, whose latent space was designed for disentangled control.
  • Whatever you pick, measure coverage, not just per-sample quality. Mode collapse produces beautiful samples that all look the same, and a metric that scores one sample at a time cannot see it.

What it costs. Sampling is the GAN's strength: one generator pass per sample. HiFi-GAN generates 22.05 kHz speech 167.9 times faster than real time on one V100 GPU, and its small version runs 13.4 times faster than real time on a CPU; a diffusion model of the 2021 generation needed 25 network passes per image to match BigGAN-deep. Training is where the cost lies: two networks, one loss surface that moves every time the other player steps, and a set of dials whose settings decide the outcome. On the ring, 2,500 rounds with the discriminator learning at 0.003 against the generator's 0.001 covers all eight clouds; at 0.001 it covers three and at 0.01 it covers six. Each stabiliser has a price. An R1 gradient penalty needs the gradient of the discriminator's own gradient, an extra backward pass on every discriminator step; spectral normalization costs a couple of matrix-vector products per layer per step; a Wasserstein critic is trained for several steps per generator step. Quality and variety trade against each other on every dial: the ten-times-faster discriminator gives 86% on-target samples against the balanced run's 81%, and leaves two clouds empty for good. BigGAN exposed the same trade as a knob, its truncation trick, which trades sample variety for fidelity.

What breaks.

  • Mode collapse. The generator covers a few kinds of data and hops between them as the discriminator catches up: on the ring, 3 clouds of 8, alternating between odd and even ones every 250 steps. Track coverage (FID or a per-class count) and rebalance the learning rates.
  • A silent gradient. With the original generator loss, a confident discriminator hands the generator almost nothing: after an 800-step head start its verdict on fakes is 0.003 and the gradient 0.09, against 15 with the non-saturating loss. Use the non-saturating loss; every modern GAN does.
  • Oscillation. Plain gradient steps rotate around the equilibrium rather than descending into it: in the two-number Dirac GAN, simultaneous steps drift from 1 away to 2.6 away in 300 steps. Losses that swing without trending are the symptom; an R1 penalty (0.08 away after 50 steps) is the fix.
  • A discriminator that wins too fast. Every region the generator has not reached becomes a wall of confident "fake", and the generator polishes what it has: six clouds from step 500 to the end, at ten times the generator's rate. A faster discriminator helps up to a point, then hurts.
  • No overlap, no direction. When real and generated data don't overlap, the original objective is flat (Jensen-Shannon stuck at 0.693 whether the generator is 1 or 10 away). The Wasserstein critic's distance still slopes towards the data.
  • A seed that lies. The same settings can find all eight clouds with one seed and three with another. Judge a recipe over several seeds.

In the wild. StyleGAN (Karras, Laine and Aila) is the reference one-pass image generator: a style-based generator that disentangles high-level attributes from stochastic detail, trained on the FFHQ face dataset the paper introduced, and trained with the R1 penalty. BigGAN scaled class-conditional GANs to ImageNet with spectral normalization, reaching an Inception Score of 166.5 and an FID of 7.4 at 128 × 128. FID itself, the standard score for generated images, came from the two time-scale paper (Heusel et al.), along with the proof that separate learning rates converge. Dhariwal and Nichol's Diffusion Models Beat GANs marked the handover on image quality. Adversarial losses now live inside other systems: the autoencoders of latent diffusion and VQGAN's tokenizer are trained with a discriminator, HiFi-GAN and the neural audio codecs use adversaries to keep waveforms clean, and Adversarial Diffusion Distillation turns a diffusion model into a one-to-four-step sampler by pairing score distillation with an adversarial loss. Every paper is linked at the end of the lesson.

Go deeper. Level 2 builds both players in NumPy, derives the best possible discriminator and the game's equilibrium, shows why the original generator loss goes silent, then reproduces each failure on the ring (the Dirac GAN's spiral, mode collapse, the over-strong discriminator) and runs each fix. If you only needed to choose, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

A forger wants to print banknotes that pass as real. A detective wants to catch every fake. At first the forger is hopeless: smudged ink, the wrong colour, and the detective spots every note. But every time the detective rejects a note, the forger learns what gave it away and fixes that. Every time the forger improves, the detective has to look more closely. Neither is ever told what a real banknote should look like. They only push against each other, and both keep getting better.

The contest ends when the forger's notes are so good that the detective can do no better than guess: "real" half the time, "fake" the other half. At that point the forger has learned to make banknotes, and nobody ever wrote down a rule for what a banknote is.

That is a generative adversarial network, a GAN. The forger is a neural network called the generator; the detective is a second network called the discriminator. Train them against each other on photos of faces and the generator learns to produce new faces that never existed.

This lesson builds both players from scratch in NumPy, trains them on a toy picture you can see at a glance, and then shows the three ways the contest goes wrong: the players chase each other in circles, the forger settles for copying a few examples (mode collapse), and the detective wins so completely that the forger stops learning. Each failure gets a fix you can run. The same goal, making new samples, is reached differently by primer.ml.generative.autoencoders (squeeze and rebuild) and primer.ml.generative.diffusion (remove noise a little at a time).

The two players

Everyday picture. The forger is a recipe that turns a few dice rolls into a banknote: different rolls, a different note. The detective is a machine you feed a note into, and out comes a single number, how sure it is that the note is real.

Tiny worked example. Our "banknotes" are points on a flat page. The real data is a ring of eight little clouds: eight centres evenly spaced on a circle of radius 2, each real point scattered a tiny amount (a spread of 0.05) around one centre, chosen at random. This toy picture comes from the research literature on GAN failures because you can see whether a forger has found all eight clouds.

  • The generator G takes two random numbers z (the dice rolls, drawn from a bell curve) and returns one point on the page. It is a small neural network: 2 numbers in, two layers of 32 tanh units, 2 numbers out.
  • The discriminator D takes a point and returns a probability that it is real. It is another small network: 2 numbers in, 64 tanh units, one score out, then squashed into the range 0 to 1.

The squashing is the sigmoid, the same one primer.ml.neural_net uses to turn a score into a probability:

Level 3: the formula and its symbols

$$ D(x) = \sigma\big(a(x)\big) = \frac{1}{1 + e^{-a(x)}} $$

Symbols

Symbol Meaning here In the example
$x$ a point on the page, real or forged $(2, 0)$
$a(x)$ the detective's raw score for $x$ (also called a logit): any number, large and positive for "surely real" $a(x) = x_1 - 1$
$x_1$ the first coordinate of $x$ 2
$e$ Euler's number, about 2.718 (see primer.notation)
$\sigma$ the sigmoid: squashes any score into a probability between 0 and 1 $\sigma(1) = 0.731$
$D(x)$ the detective's verdict: the probability that $x$ is real 0.731

In words: "the detective computes a score for the point, and the sigmoid turns the score into a probability that the point is real."

With the numbers: take a detective with a single neuron whose score is $a(x) = x_1 - 1$. The point $(2, 0)$, one of the eight real centres, scores $2 - 1 = 1$, and $\sigma(1) = 1 / (1 + e^{-1}) = 0.731$: probably real. The point $(0, 0)$ in the middle of the ring scores $-1$ and gets $\sigma(-1) = 0.269$: probably fake.

Level 3: in Python

In Python:

import math
def sigma(a):
    return 1 / (1 + math.exp(-a))
# a one-neuron detective: score a(x) = x_1 - 1
def a(x):
    return x[0] - 1
round(sigma(a((2, 0))), 3)  # → 0.731
round(sigma(a((0, 0))), 3)  # → 0.269
flowchart LR Z["noise z<br/>2 random numbers"] --> G["Generator G<br/>the forger"] G --> F["fake point G(z)"] R["real point x<br/>from the ring"] --> D["Discriminator D<br/>the detective"] F --> D D --> P["D(point)<br/>probability it is real"] P -. "learns to call real real<br/>and fake fake" .-> D P -. "learns to make D say real<br/>about its fakes" .-> G

Reading it: start at the far left. Noise goes into the generator and comes out as a fake point. Real points come in from the data. Both kinds of point go through the same detective, which gives each one a probability. The two dotted arrows are the learning signals, and they pull in opposite directions: the detective adjusts itself to score real points high and fakes low, and the generator adjusts itself to make the detective score its points high. Notice that the generator never sees a real point. Everything it learns about the data arrives through the detective's verdicts.

Why it matters: that last point is the whole trick. Nobody writes down what a face, a voice or a banknote is. The detective discovers what separates real from fake, and its gradient (the direction that would make it less sure a fake is fake; see primer.ml.neural_net for gradients) tells the forger how to improve.

In code: ring_of_gaussians draws real points around ring_modes, MLP is both players (its MLP.backward also returns the gradient with respect to the input, the channel through which the forger learns), with shapes GENERATOR_SIZES and DISCRIMINATOR_SIZES, and sigmoid turns the detective's score into a probability.

The game: one number both players fight over

Everyday picture. Picture a scoreboard. The detective earns points for confident, correct calls and loses points, heavily, for confident mistakes. The detective wants the score as high as possible. The forger wants it as low as possible. One number, two players pulling it in opposite directions: that is a minimax game.

Tiny worked example. The detective looks at two real notes and two fakes. The score uses the logarithm (log): log 1 = 0, and the log of a small number is a large negative number, so a confident mistake costs far more than a hesitant one (see primer.notation for logs from scratch).

Note Real? Verdict D What is scored Value
1 real 0.9 log D = log 0.9 −0.105
2 real 0.8 log D = log 0.8 −0.223
3 fake 0.2 log (1 − D) = log 0.8 −0.223
4 fake 0.4 log (1 − D) = log 0.6 −0.511

The average over the real notes is −0.164, over the fakes −0.367, and the total is −0.53. A perfect detective (1 on every real note, 0 on every fake) would score log 1 + log 1 = 0, the ceiling. A detective reduced to a coin flip (0.5 on everything) scores log 0.5 + log 0.5 = −1.386.

Level 3: the formula and its symbols

$$ \min_G \, \max_D \; V(D, G) = \mathbb{E}_{x \sim p_{\text{data}}}\big[\log D(x)\big] + \mathbb{E}_{z \sim p_z}\big[\log\big(1 - D(G(z))\big)\big] $$

Symbols

Symbol Meaning here In the example
$V(D, G)$ the value: the scoreboard number −0.53
$\max_D$ "the detective chooses its weights to make what follows as large as possible"
$\min_G$ "the forger chooses its weights to make that best-case value as small as possible"
$\mathbb{E}$ expectation: the average over many draws (see primer.notation, probability notation) an average over 2 notes
$x \sim p_{\text{data}}$ "$x$ drawn from the real data" notes 1 and 2
$z \sim p_z$ "$z$ drawn from the noise the forger starts from" the dice rolls behind notes 3 and 4
$D(x)$ the verdict on a real point 0.9, 0.8
$G(z)$ a fake point made from noise $z$ notes 3 and 4
$D(G(z))$ the verdict on a fake 0.2, 0.4
$\log$ natural logarithm: 0 at 1, very negative near 0 $\log 0.6 = -0.511$

In words: "average the log of the verdicts on real points, add the average log of one minus the verdicts on fakes; the detective pushes this number up, the forger pushes it down."

With the numbers: (log 0.9 + log 0.8) / 2 + (log 0.8 + log 0.6) / 2 = −0.164 + (−0.367) = −0.53.

Level 3: in Python

In Python:

import math
# D(x) on two real notes, D(G(z)) on two fakes
d_real, d_fake = [0.9, 0.8], [0.2, 0.4]
# E over real x of log D(x): an average
real_term = sum(math.log(d) for d in d_real) / len(d_real)
round(real_term, 3)  # → -0.164
# E over noise z of log(1 - D(G(z)))
fake_term = sum(math.log(1 - d) for d in d_fake) / len(d_fake)
round(fake_term, 3)  # → -0.367
round(real_term + fake_term, 2)  # → -0.53
# a coin-flip detective, D = 1/2 on everything
round(math.log(0.5) + math.log(0.5), 3)  # → -1.386

If this looks familiar, it should: −V is exactly the binary cross-entropy of a classifier that labels real points 1 and fakes 0 (see primer.ml.losses). The detective is an ordinary classifier. What is new is that its second class, the fakes, keeps changing underneath it.

Training alternates. Nobody can solve "min over G of max over D" directly, so each round takes one small step for each player:

flowchart TB S["sample a batch of real points<br/>and a batch of noise"] --> F["forger makes fakes G(z)"] F --> DS["detective step:<br/>climb V, so D(real) rises<br/>and D(fake) falls"] DS --> GS["forger step:<br/>change G so the new detective<br/>scores its fakes higher"] GS --> CHK{"trained enough?"} CHK -- no --> S CHK -- yes --> OUT["keep G;<br/>throw D away"]

Reading it: follow one lap from the top. A fresh batch of real points and noise comes in, the forger turns the noise into fakes, and then the two players move in turn: first the detective takes one step of gradient ascent on V (it wants V higher), then the forger takes one step of gradient descent (it wants V lower), judged by the detective as it is after its step. The loop repeats thousands of times. At the bottom, only the generator is kept. The detective was scaffolding: its whole job was to teach.

Why it matters: each player is trained with an ordinary optimizer on an ordinary loss, but the loss surface moves every time the other player steps. That is the root of everything that goes wrong later in this lesson.

In code: value_function computes V from a batch of verdicts, and train_gan runs the loop above on the ring, one detective step then one forger step, each with Adam (primer.ml.optimizers.Adam).

The best possible detective, and where the game ends

Everyday picture. Suppose that, at one spot on the page, real points turn up three times as often as fakes. However clever the detective is, it cannot tell two identical-looking points apart, so the best it can do there is to say "75% real". It should match the local mix, no more and no less.

Tiny worked example. At a point where the real data's density (how thickly its points cover that spot) is 0.3 and the forger's is 0.1, the best verdict is 0.3 / (0.3 + 0.1) = 0.75. Try its neighbours: the detective's expected score at that spot is 0.3 · log D + 0.1 · log(1 − D), which is −0.2274 at D = 0.7, −0.2249 at D = 0.75 and −0.2279 at D = 0.8. The middle one is the highest.

Level 3: the formula and its symbols

$$ D^*(x) = \frac{p_{\text{data}}(x)}{p_{\text{data}}(x) + p_g(x)} $$

Symbols

Symbol Meaning here In the example
$D^*(x)$ the best possible verdict at $x$, for a forger that is held fixed 0.75
$p_{\text{data}}(x)$ how densely the real data covers the spot $x$ 0.3
$p_g(x)$ how densely the forger's fakes cover the spot $x$ 0.1

In words: "the best verdict at each point is the share of the points found there that are real."

With the numbers: 0.3 / (0.3 + 0.1) = 0.75. If the forger matched the data exactly, $p_g = p_{\text{data}}$ everywhere, and every verdict would be 0.3 / (0.3 + 0.3) = 0.5.

Level 3: in Python

In Python:

import math
p_data, p_g = 0.3, 0.1
round(p_data / (p_data + p_g), 2)  # → 0.75
# the detective's expected score at this spot, for a few verdicts D
def score(D):
    return p_data * math.log(D) + p_g * math.log(1 - D)
[round(score(D), 4) for D in (0.7, 0.75, 0.8)]  # → [-0.2274, -0.2249, -0.2279]
# a forger that matches the data exactly
round(0.3 / (0.3 + 0.3), 2)  # → 0.5

Where does the formula come from? At each point the detective is choosing one number D to maximise $p_{\text{data}} \log D + p_g \log(1 - D)$. The slope of that with respect to D is $p_{\text{data}}/D - p_g/(1 - D)$, and setting the slope to zero gives $D^$. This is the central result of the original GAN paper, and it tells you where the game ends. At the equilibrium the forger's distribution equals the data's, and the best detective says 1/2 everywhere, scoring V = −log 4 ≈ −1.386, the coin-flip score from the table above. Plug $D^$ back into V and what remains is −log 4 plus twice the Jensen-Shannon divergence between the real and forged distributions, a measure of how different two piles of probability are that is zero only when they are identical. So a forger facing a perfect detective is really minimising that divergence. Keep that in mind: it returns as a problem in the fixes section.

Top: the real data has two bumps and the forger one wide bump; bottom: the best detective's verdict rises to about 0.8 on the real bumps, drops towards 0 where only fakes live, and would be flat at one half if the forger matched

Reading it: the top panel shows two distributions on a line: the real data (blue) has two narrow bumps at −2 and +2, and the forger (red) spreads one wide bump over the middle. The bottom panel is the best verdict $D^$ at every position. Over the real bumps it rises to about 0.8, because most points found there are real (the forger's wide bump still reaches them). In the middle and at the far edges it falls towards 0, because only fakes live there. The grey dashed line at 0.5 is what $D^$ becomes once the forger matches the data: the detective has nothing left to go on. The shape of the bottom curve is exactly the information the forger needs, which way to move its mass.

After training with balanced learning rates, the forger's green points sit on all eight clouds, and the detective's verdict on the ring is close to one half

Reading it: this is the real ring after 2,500 rounds of training. The background colour is the trained detective's verdict at every spot: blue for "real", red for "fake", white for 0.5. Black dots are real points, green dots are the forger's. The green points sit on all eight clouds, and the background there is pale, close to white: at the eight centres the verdict is between 0.46 and 0.58. That is the equilibrium made visible, the detective reduced to guessing where the forgery is good. Away from the ring, where neither real nor fake points ever appear, the formula says nothing, and the detective's colouring there is arbitrary (deep red in the empty middle).

Why it matters: the detective never needs to model what real data looks like; it only estimates a ratio between two densities. That is why GANs could learn sharp images long before anyone could write down a probability for an image.

In code: optimal_discriminator is the formula, fit_discriminator_table trains the most flexible detective possible (one free score per point) against a fixed forger and arrives at the same verdicts without being told the formula, and detective_verdicts evaluates a trained detective anywhere on the page.

Why the forger uses a different loss

Everyday picture. Early in training, the forger is terrible and the detective is sure every fake is fake. Now picture a game of "hotter, colder" where the helper stops talking once you are far enough away: the further off you are, the quieter the hints. That is the original forger loss. The fix is a helper who shouts louder the further off you are.

Tiny worked example. The detective looks at a fake and is 99% sure it is fake: D(G(z)) = 0.01. How hard does each loss push the forger?

  • The original loss, log(1 − D(G(z))), which the forger minimises: its slope with respect to the detective's score is −D = −0.01. Almost nothing.
  • The non-saturating loss, −log D(G(z)), which the forger also minimises: its slope is −(1 − D) = −0.99. Ninety-nine times more.

When the detective is undecided, D = 0.5, both slopes are −0.5. The two losses agree when it doesn't matter and differ exactly when it does.

Level 3: the formula and its symbols

$$ \frac{\partial}{\partial a}\,\log\big(1 - \sigma(a)\big) = -\sigma(a) = -D \qquad\qquad \frac{\partial}{\partial a}\,\big(-\log \sigma(a)\big) = -\big(1 - \sigma(a)\big) = -(1 - D) $$

Symbols

Symbol Meaning here In the example
$a$ the detective's score for a fake, before the sigmoid $\log(0.01/0.99) = -4.6$
$\sigma(a)$ the verdict D on that fake 0.01
$\frac{\partial}{\partial a}$ the derivative with respect to $a$: how much the loss changes when the score is nudged up a tiny amount (see primer.notation)
$\log(1 - \sigma(a))$ the original, "saturating" forger loss $\log 0.99$
$-\log \sigma(a)$ the non-saturating forger loss $-\log 0.01 = 4.6$

In words: "the original loss changes by only D when the detective's score moves, so it goes quiet when D is near 0; the non-saturating loss changes by 1 − D, so it is loudest exactly when the forger is doing worst."

With the numbers: at D = 0.01, the slopes are −0.01 and −0.99.

Level 3: in Python

In Python:

import math
def sigma(a):
    return 1 / (1 + math.exp(-a))
# a score that makes the detective 99% sure the fake is fake
a = math.log(0.01 / 0.99)
round(sigma(a), 2)  # → 0.01
# measure each slope by nudging the score a tiny bit either way
h = 1e-6
def slope(loss):
    return (loss(a + h) - loss(a - h)) / (2 * h)
def saturating(a):
    return math.log(1 - sigma(a))
def non_saturating(a):
    return -math.log(sigma(a))
round(slope(saturating), 3)  # → -0.01
round(slope(non_saturating), 3)  # → -0.99

Both losses want the same thing, a higher verdict on the fakes, and they share the same equilibrium. Only the strength of the push differs. The original GAN paper already recommended the switch, and essentially every GAN since uses the non-saturating loss or a Wasserstein-style loss (below).

Left: as the detective grows sure a fake is fake, the original loss flattens out while the non-saturating loss keeps a steady slope; right: giving the detective an 800-step head start shrinks the original loss's gradient to 0.09 while the non-saturating one grows to 15

Reading it: on the left, the x-axis is the detective's score for a fake: far left means "certainly fake". The red curve, the original loss, goes flat as you move left, and a flat curve has no slope, so no gradient. The blue curve, the non-saturating loss, keeps a steady downhill slope all the way. On the right is a measurement from our ring: the detective trains alone against a frozen, untrained forger for more and more steps, and we measure how big a gradient each loss would hand the forger (log scale). At the start both are about 0.5 to 0.6. After 800 detective steps its verdict on fakes is 0.003, the original loss's gradient has shrunk to 0.09, and the non-saturating one has grown to 15: over a hundred times larger.

Why it matters: a forger that gets no gradient does not learn, and the detective's lead only widens. Early in training the detective always has the easy job, so this failure is the default, not an edge case.

In code: generator_logit_gradient returns both slopes, and head_start_experiment gives the detective a head start and measures the forger's gradient under each loss.

Why training is unstable, part 1: the players chase each other in circles

Everyday picture. Two people share an office thermostat. One turns the heat up whenever it feels cold; the other opens a window whenever it feels hot. Each reacts to where the room was, so the temperature overshoots one way, then the other, and never settles. A GAN's players do the same: the forger moves to where the detective currently says "real", and by the time it arrives, the detective has moved.

Tiny worked example. The smallest GAN there is, the Dirac GAN, has one number per player. The real data is a single point at x = 0. The forger has one number θ and always outputs x = θ. The detective has one number ψ and scores a point as ψ · x, so D(x) = σ(ψx). Start with the forger at θ = 1 and a detective with ψ = 0 (it cannot tell anything apart), and let both take steps of size 1 at the same time:

Step σ(ψθ) θ (forger) ψ (detective) Distance from (0, 0)
0 1 0 1
1 σ(0) = 0.5 1 −0.5 1.118
2 σ(−0.5) = 0.3775 0.811 −0.878 1.195

After step 1 the detective has learned "bigger x means fake" (ψ turned negative), and after step 2 the forger has moved towards the data at 0. Both moves were sensible, yet the pair got further from the equilibrium, the point (0, 0) where the forger sits on the data and the detective is indifferent.

Level 3: the formula and its symbols

$$ V(\theta, \psi) = \log \sigma(\psi \cdot 0) + \log\big(1 - \sigma(\psi\theta)\big) \qquad \frac{\partial V}{\partial \psi} = -\sigma(\psi\theta)\,\theta \qquad \frac{\partial V}{\partial \theta} = -\sigma(\psi\theta)\,\psi $$

Symbols

Symbol Meaning here In the example
$\theta$ the forger's only number: where it puts its fake 1
$\psi$ the detective's only number: the slope of its score $\psi x$ 0
$\psi \cdot 0$ the detective's score for the real data at $x = 0$: always 0, so $D(0) = 0.5$ 0
$\sigma(\psi\theta)$ the verdict on the fake 0.5
$\frac{\partial V}{\partial \psi}$ the slope the detective climbs: $\psi \leftarrow \psi + \eta\,\frac{\partial V}{\partial \psi}$ $-0.5$
$\frac{\partial V}{\partial \theta}$ the slope the forger descends: $\theta \leftarrow \theta - \eta\,\frac{\partial V}{\partial \theta}$ 0
$\eta$ the step size (learning rate) 1
$\leftarrow$ "is replaced by"

In words: "the detective tilts its slope against wherever the fake sits; the forger slides along whichever way the slope says is more real; each reacts to the other's old position."

With the numbers: step 1: σ(0) = 0.5, so ψ becomes 0 − 1 · 0.5 · 1 = −0.5 and θ stays 1 + 1 · 0.5 · 0 = 1. Step 2: σ(−0.5) = 0.3775, so θ becomes 1 + 0.3775 · (−0.5) = 0.811 and ψ becomes −0.5 − 0.3775 · 1 = −0.878.

Level 3: in Python

In Python:

import math
def sigma(a):
    return 1 / (1 + math.exp(-a))
theta, psi, lr = 1.0, 0.0, 1.0
s = sigma(psi * theta)
s  # → 0.5
# both move at once, each using the other's old position
theta, psi = theta + lr * s * psi, psi - lr * s * theta
(theta, psi)  # → (1.0, -0.5)
s = sigma(psi * theta)
round(s, 4)  # → 0.3775
theta, psi = theta + lr * s * psi, psi - lr * s * theta
(round(theta, 3), round(psi, 3))  # → (0.811, -0.878)
# distance from the equilibrium (0, 0): 1, then 1.118, then 1.195
round(math.hypot(theta, psi), 3)  # → 1.195
flowchart LR A["fake sits to the right<br/>of the data"] --> B["detective tilts:<br/>right means fake"] B --> C["forger slides left,<br/>overshoots past the data"] C --> D["detective tilts back:<br/>left means fake"] D --> E["forger slides right,<br/>overshoots again"] E --> A

Reading it: the boxes form a loop because the game has no downhill direction that both players share. Each arrow is a sensible move for the player making it, but each player only reacts to the other's current position, so the forger arrives just after the detective has changed its mind. The loop is a rotation around the equilibrium, not a descent into it.

In the plane of the forger's theta against the detective's psi, simultaneous steps spiral outward from the start at (1, 0), and alternating steps circle the equilibrium without ever reaching it

Reading it: the horizontal axis is the forger's θ, the vertical axis the detective's ψ, and the black cross at (0, 0) is the equilibrium. Both paths start at the dot, (1, 0), and take 300 steps of size 0.2. The red path, where both players step at once, spirals outwards: after 300 steps it is 2.6 away from the equilibrium, where it started 1 away. The blue path, where the forger waits to see the detective's new ψ before it moves, stays on a closed loop about 1 away forever. Neither ever arrives. Plain gradient steps do not solve this game even in two dimensions.

Why it matters: a real GAN has millions of numbers per player, and this rotation happens in many directions at once. It shows up as losses that swing up and down without trending, and samples that keep changing character instead of steadily improving.

In code: dirac_gan runs this two-number game with simultaneous or alternating steps.

Why training is unstable, part 2: mode collapse

Everyday picture. The forger discovers that one particular note, say the twenty, always fools the detective. So it prints nothing but twenties. The detective eventually learns that twenties are suspicious, so the forger switches to printing only fifties, and so on. Each forgery is good, but the forger never produces the full variety of real money. Each distinct kind of real example is a mode, and settling on a few of them is mode collapse.

Tiny worked example. To measure variety on the ring, generate 1,000 points and count, for each of the eight clouds, how many land within 0.15 of its centre (3 spreads). A cloud is covered if it gets at least 20 of the 1,000 (2%; a perfect forger gives each about 125). A point that lands near any centre counts as high quality. After 2,500 rounds with equal learning rates for both players, the counts for the eight clouds are (0, 68, 0, 20, 0, 51, 0, 17). Clouds 2, 4 and 6 are covered, cloud 8's 17 falls short, so 3 of 8 modes are covered, and only 156 of the 1,000 points are high quality: the rest are strung between clouds.

stateDiagram-v2 direction LR EvenClouds: forger piles onto clouds 1, 3, 5, 7 CatchEven: detective learns those clouds are suspicious OddClouds: forger jumps to clouds 2, 4, 6, 8 CatchOdd: detective learns these are suspicious EvenClouds --> CatchEven CatchEven --> OddClouds OddClouds --> CatchOdd CatchOdd --> EvenClouds

Reading it: each box is a phase of training, and the arrows go round in a circle. The forger does not have to cover every cloud to fool the detective today; it only has to put its points where the detective currently says "real". The detective then learns those spots are over-supplied with fakes, the verdict there drops, and the forger's gradient pushes the whole pile somewhere else. It is the chase from part 1, now played out across the clouds of the data. The alternation between odd and even clouds is what our run actually does, pictured next.

Three rows of four snapshots between steps 1,750 and 2,500: with equal learning rates the green points hop between alternate clouds; with a detective three times faster they cover all eight; with a detective ten times faster they sit on six and never reach the other two

Reading it: each panel is the page, grey circles mark the eight real clouds, and green dots are 1,000 points from the forger at that step (always made from the same noise, so the panels are comparable). The panel titles spell the coverage with X for a covered cloud, counting anticlockwise from the cloud on the right. Top row, equal learning rates: the green points hop from clouds 1, 3 and 7 at step 1,750 to 2, 4, 6 and 8 at step 2,000, back to 1, 3 and 5, then to 2, 4 and 6, never holding more than half the ring. Middle row, a detective that learns three times faster than the forger: all eight clouds, every snapshot. Bottom row, a detective ten times faster: six clouds, and the same six at every snapshot. Two clouds stay empty for good.

Modes covered over training: equal learning rates wander between one and four; a three-times-faster detective reaches all eight by step 1,000 and stays; a ten-times-faster detective locks onto six at step 500

Reading it: the x-axis is the training step, the y-axis the number of clouds covered out of eight. The blue line (detective three times faster) climbs to 8 by step 1,000 and stays there. The red line (equal speeds) zigzags between 1 and 4, the hopping from the top row above. The amber line (detective ten times faster) jumps to 6 by step 500 and never moves again. The same code and the same data give three very different outcomes, and the only thing that changed is how fast the detective learns.

These runs use one fixed seed. With other seeds the equal-speed run sometimes finds all eight clouds late on, and the ten-times run sometimes escapes after a while; the pattern, not the exact counts, is the lesson.

Why it matters: mode collapse is the classic GAN failure on real data. A face generator that has collapsed produces beautiful faces that all look like the same few people. Quality metrics that look at one sample at a time can't see it; you have to measure coverage, as we did here.

In code: mode_coverage counts covered clouds and high-quality points, and train_gan records it every 250 steps.

Why training is unstable, part 3: a detective that wins too fast

Everyday picture. A detective who learns much faster than the forger soon rejects everything outside the few spots the forger has already mastered, with total confidence. Every small experiment the forger tries in a new direction comes back "more fake", so it retreats to what already works and stops exploring.

Tiny worked example. The bottom row above: with the detective's learning rate ten times the forger's, after 2,500 steps the counts per cloud are (199, 181, 117, 0, 132, 160, 72, 0). The six covered clouds are served well: 86% of points are high quality, more than the balanced run's 81%. But clouds 4 and 8 get nothing, from step 500 until the end. Better-looking samples, less variety: a trade-off you will meet again with real image generators.

flowchart LR A["detective learns<br/>much faster"] --> B["gaps between the forger's clouds<br/>become confident fake zones"] B --> C["every small step towards them<br/>lowers the verdict"] C --> D["forger polishes the clouds<br/>it already has"] D --> E["empty clouds stay empty"]

Reading it: read it as a chain of causes. A fast detective turns every region the forger isn't already covering into a wall of confident "fake" verdicts. To reach an empty cloud the forger would have to move points through such a wall. But a gradient only sees the next small step, and every small step into the wall makes the verdict worse, so the gradient pushes the points back onto the clouds they already occupy. The forger spends its effort sharpening those, and the missing ones are never discovered.

Why it matters: "just train the detective harder" sounds like it should give the forger a better teacher. Past a point it gives it a worse one. The detective needs to be good enough to point the way, not so good that it only says no.

In code: train_gan takes both learning rates as arguments; this run gives the detective ten times the forger's.

Fixes that make the game trainable

Everyday picture. A good sparring partner matters more than a strong one. Every stabiliser below is a way of keeping the detective useful: fast enough to teach, smooth enough that its verdicts point somewhere.

flowchart LR P1["players out of step"] --> F1["careful learning rates<br/>two time-scales"] P2["detective builds<br/>steep cliffs"] --> F2["gradient penalty"] P2 --> F3["spectral normalization"] P3["piles don't overlap,<br/>so the verdict gives no direction"] --> F4["Wasserstein distance"] F4 --> F2

Reading it: failures are on the left and the fixes that target them on the right. Players moving out of step are fixed by choosing their speeds. A detective whose verdict jumps from 0 to 1 across a narrow cliff is fixed by penalising or capping its steepness. A yardstick that can't tell "near miss" from "far miss" is replaced by the Wasserstein distance, which in practice needs one of the steepness fixes to compute, hence the last arrow.

Careful learning rates

Everyday picture. Pace the two sparring partners so that neither runs away with the match.

Tiny worked example. The three rows of the ring figure differ only in the detective's learning rate, with the forger's fixed at 0.001:

Detective's rate Clouds covered at step 2,500 High-quality points
0.001 (equal) 3, and hopping 16%
0.003 (three times) 8, steady 81%
0.01 (ten times) 6, stuck 86%

A modestly faster detective keeps its verdicts close to the best-detective formula for the forger it is facing now, so the forger follows a gradient that points at the data rather than at yesterday's detective. Heusel and colleagues called this the two time-scale update rule and proved it reaches an equilibrium under reasonable conditions. The right ratio depends on the problem; the lesson is that it is a dial worth turning.

In code: train_gan with a detective learning rate of 0.003 is the balanced run; the other two rows are the same call with 0.001 and 0.01.

Gradient penalty: fine the detective for steep cliffs

Everyday picture. Charge the detective a fee for how sharply its score changes right on top of the real data. A detective with gentle slopes still ranks real above fake, but it can no longer swing wildly, and the swinging is what fed the spiral in part 1.

Tiny worked example. In the Dirac GAN the detective's score is ψ · x, whose slope with respect to x is just ψ. The R1 gradient penalty adds (γ / 2) · ψ² to what the detective minimises. With γ = 1 and ψ = −0.5 the fee is 0.125, and its gradient, γψ = −0.5, pulls ψ back towards 0 on every step, like friction. With step size 0.2 that friction shrinks ψ by a factor 1 − 0.2 · 1 = 0.8 per step, on top of the game's own push.

Level 3: the formula and its symbols

$$ R_1 = \frac{\gamma}{2}\;\mathbb{E}_{x \sim p_{\text{data}}}\Big[\big\lVert \nabla_x\, a(x) \big\rVert^2\Big] $$

Symbols

Symbol Meaning here In the Dirac example
$R_1$ the penalty added to the detective's loss 0.125
$\gamma$ how heavily steepness is fined 1
$a(x)$ the detective's score at $x$ $\psi x$
$\nabla_x\, a(x)$ the gradient of the score with respect to the input point: which way, and how steeply, the score rises as you move $x$ $\psi = -0.5$
$\lVert \cdot \rVert^2$ the squared length of that gradient (see primer.notation, the length of a vector) 0.25
$\mathbb{E}_{x \sim p_{\text{data}}}$ averaged over real points only the single real point $x = 0$

In words: "measure how steeply the detective's score changes at the real points, square it, average it, and charge the detective γ/2 times that."

With the numbers: (1 / 2) · (−0.5)² = 0.125.

Level 3: in Python

In Python:

gamma, psi = 1.0, -0.5
# slope of the score psi * x with respect to x, at the real point
slope = psi
gamma / 2 * slope ** 2  # → 0.125
# the penalty's pull on psi: its derivative, gamma * psi
gamma * psi  # → -0.5
# one step of size 0.2 against that pull shrinks psi by a factor 0.8
round(psi - 0.2 * gamma * psi, 2)  # → -0.4

Distance from the equilibrium over 300 steps: simultaneous steps grow from 1 to 2.6, alternating steps hold near 1, and with the R1 penalty the distance falls to 0.08 within 50 steps and to zero after

Reading it: the x-axis is the step, the y-axis the distance of (θ, ψ) from the equilibrium (0, 0). The red line (plain simultaneous steps) climbs in steps, one per lap: the outward spiral. The blue line (alternating steps) hovers near 1 forever: the endless orbit. The green line adds the R1 penalty with γ = 1 to the simultaneous steps and dives: 0.08 after 50 steps, essentially 0 by

  1. The friction turns the rotation into a spiral inwards. That is the Mescheder, Geiger and Nowozin result in two numbers, and it is why StyleGAN and many later GANs train with an R1 penalty. The Wasserstein GAN with gradient penalty (WGAN-GP) uses the same idea with a different target: it keeps the slope near 1 at points between real and fake.

In code: dirac_gan takes the penalty weight γ as an argument and adds the R1 penalty's pull, −γψ, to the detective's step.

Wasserstein distance: a yardstick that knows near from far

Everyday picture. Picture the real data and the fakes as two piles of sand. The Wasserstein distance, also called the earth mover's distance, is the least work needed to reshape one pile into the other: the amount of sand moved times how far it travels. The Jensen-Shannon divergence from the best-detective section only asks how much the piles overlap.

Tiny worked example. Put all the real sand at x = 0 and all the fake sand at x = θ. If θ = 1, the piles don't overlap, and the Jensen-Shannon divergence is log 2 = 0.693. If θ = 5, they still don't overlap, and it is still 0.693. It cannot tell a near miss from a far one, so it gives the forger no direction. The Wasserstein distance is 1 for θ = 1 and 5 for θ = 5: it shrinks as the forger approaches, so it always points home.

For two equally sized samples on a line, the cheapest way to move the sand is to pair the smallest with the smallest, the next with the next, and so on:

Level 3: the formula and its symbols

$$ W_1 = \frac{1}{n} \sum_{i=1}^{n} \big\lvert a_{(i)} - b_{(i)} \big\rvert $$

Symbols

Symbol Meaning here In the example
$W_1$ the Wasserstein (earth mover's) distance between two piles on a line 3
$n$ how many grains (samples) each pile has 3
$a_{(i)}$ the $i$-th smallest value in the first pile (the brackets mean "after sorting") $a = (0, 1, 2)$
$b_{(i)}$ the $i$-th smallest value in the second pile $b = (5, 3, 4)$ sorts to $(3, 4, 5)$
$\lvert \cdot \rvert$ the absolute value: the gap, ignoring direction $\lvert 0 - 3 \rvert = 3$

In words: "sort both piles, pair them up in order, and average how far each grain has to travel."

With the numbers: pairs (0, 3), (1, 4), (2, 5), gaps 3, 3, 3, average 3.

Level 3: in Python

In Python:

a = [0, 1, 2]
b = [5, 3, 4]
# pair the smallest with the smallest, and so on
pairs = list(zip(sorted(a), sorted(b)))
pairs  # → [(0, 3), (1, 4), (2, 5)]
sum(abs(x - y) for x, y in pairs) / len(pairs)  # → 3.0

As the forger's pile slides from minus 4 to plus 4, the Jensen-Shannon divergence is flat at log 2 everywhere except exactly 0, while the Wasserstein distance is a V shape pointing at 0

Reading it: the x-axis is where the forger's pile sits; the real pile is at 0. The red line (Jensen-Shannon) is flat at 0.693 everywhere except the single point θ = 0, where it drops to 0: a flat line has no slope, so nothing tells the forger which way to move. The blue V (Wasserstein) slopes down towards 0 from both sides: wherever the forger is, the slope points at the data. On images, where real photos and early fakes barely overlap, this difference is the difference between learning and not.

The Wasserstein GAN replaces the detective with a critic that outputs an unbounded score instead of a probability, and trains it so that the gap between its average score on real and on fake points estimates this distance. The catch: the estimate is only valid if the critic's slope is at most 1 everywhere. The original WGAN enforced that crudely by clipping every weight; WGAN-GP does it with a gradient penalty; spectral normalization, next, does it by capping each layer.

In code: wasserstein_1d is the sort-and-pair formula, and js_divergence computes the Jensen-Shannon divergence between two histograms, flat at log 2 whenever they don't overlap.

Spectral normalization: a speed limit on every layer

Everyday picture. A layer of a network is a matrix, and a matrix stretches some directions more than others. If no layer can stretch anything by more than 1, the whole detective can't change its score faster than the input changes: its cliffs have a speed limit.

Tiny worked example. The matrix W = [[3, 0], [0, 1]] stretches the horizontal direction by 3 and the vertical by 1. Its largest stretch is

  1. Divide W by 3 and no direction is stretched by more than 1.
Level 3: the formula and its symbols

$$ \bar W = \frac{W}{\lVert W \rVert_2} \qquad \lVert W \rVert_2 = \max_{\lVert v \rVert = 1} \lVert W v \rVert $$

Symbols

Symbol Meaning here In the example
$W$ one layer's weight matrix (see primer.notation for matrices) [[3, 0], [0, 1]]
$v$ an input direction: an arrow of length 1 $(1, 0)$
$\lVert v \rVert$ the length of $v$ 1
$Wv$ the arrow after passing through the layer $(3, 0)$
$\max_{\lVert v \rVert = 1}$ "the largest value over every arrow of length 1"
$\lVert W \rVert_2$ the spectral norm: the most W lengthens any arrow 3
$\bar W$ the normalised layer used in the detective [[1, 0], [0, 1/3]]

In words: "find the direction the layer stretches most, and divide the whole layer by that stretch."

With the numbers: $(1, 0)$ becomes $(3, 0)$, length 3; $(0, 1)$ stays length 1; every other direction lands in between. So $\lVert W \rVert_2 = 3$, and $W / 3$ stretches by at most 1.

Level 3: in Python

In Python:

import math
W = [[3, 0], [0, 1]]
def stretch(v):
    Wv = [sum(w * x for w, x in zip(row, v)) for row in W]
    return math.hypot(*Wv) / math.hypot(*v)
# try unit arrows all the way round the circle, one per degree
angles = [2 * math.pi * k / 360 for k in range(360)]
round(max(stretch((math.cos(t), math.sin(t))) for t in angles), 3)  # → 3.0
# divide W by its largest stretch
W = [[w / 3 for w in row] for row in W]
round(max(stretch((math.cos(t), math.sin(t))) for t in angles), 3)  # → 1.0

Trying every direction is fine for a 2 × 2 matrix but not for a layer with a thousand inputs. Power iteration finds the most-stretched direction cheaply:

flowchart LR V["start: any arrow v"] --> U["u = W v,<br/>rescaled to length 1"] U --> B["v = W transpose u,<br/>rescaled to length 1"] B --> Q{"repeat"} Q -- again --> U Q -- done --> S["largest stretch = length of W v"]

Reading it: push an arrow through the layer and back through its transpose, over and over. Each round trip multiplies the arrow's most-stretched component by more than any other, so the arrow swings round to point along that direction, and the stretch is read off at the end. Spectral normalization keeps one arrow per layer and runs a single round trip per training step: the weights change slowly, so the arrow stays nearly right, and the cost is a couple of extra matrix-vector products.

Why it matters: spectral normalization (Miyato and colleagues, 2018) is a one-line change to the detective that made GAN training far less sensitive to learning rates and architecture, and it became a default in large image GANs such as BigGAN.

In code: largest_stretch is power iteration, and spectrally_normalize divides a matrix by it.

Where GANs stand today

Everyday picture. The forger-and-detective idea started as the whole machine. Today it is more often a picky critic inside another machine, brought in for the one thing it does best: making outputs look sharp and real.

Tiny worked example. A GAN makes an image in one pass through the generator. A diffusion model typically runs its network tens of times per image, removing a little noise each time (see primer.ml.generative.diffusion). For a while, that speed plus very sharp samples (StyleGAN's faces, from 2018, are still striking) made GANs the best image generators. Then diffusion models overtook them on image quality and, above all, on coverage: they don't mode-collapse and they train stably, which is everything this lesson has been fighting. Since around 2021 most new image and video generators are diffusion or flow models.

GANs didn't disappear. Their loss did the moving:

flowchart LR X["real image or audio"] --> E["encoder"] E --> C["compact code"] C --> DEC["decoder"] DEC --> Y["reconstruction"] Y --> L1["reconstruction loss<br/>pixel or perceptual"] X --> L1 Y --> DISC["detective:<br/>real or reconstructed?"] X --> DISC L1 --> T["total loss for<br/>encoder and decoder"] DISC --> T

Reading it: this is an autoencoder (see primer.ml.generative.autoencoders) with a detective bolted on. The top path squeezes the input into a compact code and rebuilds it. A reconstruction loss alone rewards averages, and the average of many plausible textures is a blur. The detective on the lower path looks at the reconstruction and the original and asks "which is real?", which punishes blur and rewards crisp detail. The decoder is trained on both losses together. This design is how the image autoencoders inside latent diffusion models are trained, how VQGAN's image tokenizer is trained, and how neural audio decoders such as HiFi-GAN produce clean waveforms. GAN-style losses are also used to distil a slow many-step diffusion model into a fast one- or few-step generator.

Why it matters: when you meet a modern image, audio or video system, expect a GAN loss somewhere in its decoder or its fast sampling path, even when the headline method is diffusion. Everything in this lesson, the non-saturating loss, gradient penalties, spectral normalization and careful learning rates, is still how those critics are kept stable.

In 20 seconds

  • A GAN trains a generator (forger) to turn noise into samples and a discriminator (detective) to tell real samples from fakes, against each other. The forger learns only through the detective's gradient.
  • The game: min over G of max over D of E[log D(x)] + E[log(1 − D(G(z)))]. The best detective says p_data / (p_data + p_g); at the equilibrium the forger matches the data and the detective says 1/2 everywhere.
  • The non-saturating loss (maximise log D(G(z))) keeps the forger's gradient strong when the detective is confident, exactly when the original loss goes quiet.
  • Unstable because each player's target moves: plain gradient steps circle or spiral around the equilibrium, the forger hops between a few modes (mode collapse), and a detective that wins too fast freezes it.
  • Fixes: balanced learning rates, gradient penalties (R1, WGAN-GP), spectral normalization, and the Wasserstein distance, which still points the way when real and fake don't overlap.
  • Today: diffusion has overtaken GANs for generating images, but adversarial losses remain inside image and audio decoders and fast distilled samplers.

Self-test questions

Explain a GAN to a non-engineer in 30 seconds. Two programs play a game. One makes fake pictures; the other looks at real and fake pictures and guesses which is which. Every time the guesser catches a fake, the faker learns what gave it away and improves. After enough rounds the fakes are good enough that the guesser is reduced to guessing, and the faker has learned to make realistic pictures without anyone ever describing what a picture should look like.

Why does the generator never need to see real data? It learns only from the discriminator's gradient with respect to its own outputs: which way to move each fake so the discriminator finds it more real. The discriminator has seen real data, so its verdicts carry that information to the generator.

What does the optimal discriminator compute, and what does it say at equilibrium? D*(x) = p_data(x) / (p_data(x) + p_g(x)): the share of points found at x that are real. When the generator matches the data, p_g = p_data and D* is 1/2 everywhere, the value is −log 4, and the discriminator can do no better than a coin flip.

Why do almost all GANs use the non-saturating generator loss? The original loss log(1 − D(G(z))) has a slope of −D with respect to the discriminator's score, which is near zero when the discriminator confidently rejects fakes, as it does early in training. The non-saturating loss −log D(G(z)) has slope −(1 − D), largest exactly then. Both have the same equilibrium.

What is mode collapse, and how would you detect it? The generator covers only some of the distinct kinds of data (modes), often hopping between them as the discriminator catches up. Per-sample quality can look excellent, so you detect it by measuring coverage: how many modes or classes receive samples, or a distribution-level metric such as FID on real data.

Why doesn't plain gradient descent find the GAN equilibrium? The game has no shared downhill direction. Near the equilibrium the two players' updates form a rotation, so simultaneous steps spiral outward and alternating steps orbit forever, as the two-number Dirac GAN shows. Damping the discriminator with a gradient penalty turns the spiral inward.

Why is the Wasserstein distance a better training signal than Jensen-Shannon divergence? When the real and generated distributions don't overlap, Jensen-Shannon is stuck at log 2 however far apart they are, so it gives no direction. Wasserstein measures how far the mass must move, so it shrinks steadily as the generator approaches the data. Estimating it needs a critic whose slope is capped, which is what weight clipping, gradient penalties and spectral normalization provide.

If diffusion models won, why learn GANs? Adversarial losses are still how many image and audio decoders get sharp output, and how some diffusion models are distilled into one-step generators. And the instabilities here, moving targets and collapsing variety, show up wherever two models are trained against each other.

The papers behind this lesson

  • Goodfellow et al., Generative Adversarial Nets (2014): https://arxiv.org/abs/1406.2661. Introduced the generator-discriminator game, derived the optimal discriminator and the equilibrium where the generator matches the data, and suggested the non-saturating loss. Annotated companion
  • Metz, Poole, Pfau and Sohl-Dickstein, Unrolled Generative Adversarial Networks (2016): https://arxiv.org/abs/1611.02163. Used the ring of eight Gaussians to show a generator hopping between modes, and reduced it by letting the generator look ahead at several discriminator steps.
  • Heusel et al., GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium (2017): https://arxiv.org/abs/1706.08500. Showed that separate learning rates for the two players give provable convergence, and introduced the FID score for judging generated images.
  • Arjovsky, Chintala and Bottou, Wasserstein GAN (2017): https://arxiv.org/abs/1701.07875. Replaced the Jensen-Shannon objective with the earth mover's distance, which still gives a direction when real and generated data don't overlap. Annotated companion
  • Gulrajani et al., Improved Training of Wasserstein GANs (2017): https://arxiv.org/abs/1704.00028. Enforced the Wasserstein critic's slope limit with a gradient penalty instead of weight clipping.
  • Mescheder, Geiger and Nowozin, Which Training Methods for GANs do actually Converge? (2018): https://arxiv.org/abs/1801.04406. Introduced the two-number Dirac GAN to show why plain GAN training circles instead of converging, and the R1 penalty that makes it converge.
  • Miyato et al., Spectral Normalization for Generative Adversarial Networks (2018): https://arxiv.org/abs/1802.05957. Capped every discriminator layer's largest stretch at 1 using one step of power iteration per training step.

Further reading

on GitHub
   1r"""
   2# GANs: a forger against a detective
   3
   4Run: `python -m primer.ml.generative.gans`
   5
   6This lesson builds on the small networks and gradients of
   7`primer.ml.neural_net`, the binary cross-entropy of `primer.ml.losses` and
   8the Adam optimizer of `primer.ml.optimizers`; `primer.notation` explains
   9every symbol from zero.
  10
  11## Level 1: The practitioner's guide
  12
  13**In one sentence.** A generative adversarial network (GAN) trains a
  14generator to turn random noise into samples by pitting it against a
  15discriminator that learns to tell its samples from real ones, which yields
  16sharp output in a single pass at the price of a training game that is hard
  17to keep stable.
  18
  19**When you need it.** Two tells. You need a generator that runs in one
  20pass, because it sits in a real-time loop (a speech vocoder, an interactive
  21tool, a game) and a many-step diffusion sampler is too slow. Or you have a
  22decoder that rebuilds images or audio and its output is blurry: a
  23reconstruction loss rewards averages, and an adversarial loss is the
  24standard cure. You don't need to train a GAN from scratch to generate new
  25images or video today: since 2021 diffusion and flow models are the default
  26there, because they train stably and cover the data, which is exactly what
  27this lesson shows GANs struggling to do. The number that shows the naive
  28approach failing: on this lesson's toy data (eight small clouds on a ring),
  29a GAN trained with the same learning rate for both players covers 3 of the 8
  30clouds after 2,500 rounds, and only 16% of its samples land on any cloud.
  31Change nothing but the discriminator's learning rate to three times the
  32generator's and it covers all 8 with 81% on target.
  33
  34**Your options.** From the least commitment to the most:
  35
  36| Option | What it does | What it guarantees | What it costs | Where it lives |
  37|---|---|---|---|---|
  38| A diffusion or flow model instead | Generates by removing noise over many passes | Stable training, full coverage; matches BigGAN-deep's image quality with 25 passes and better coverage | Tens of network passes per sample | A hosted API or a diffusion library |
  39| A pretrained GAN generator | One pass from noise to sample, with an editable latent space | Real-time generation in its domain (faces, one class of object) | You are limited to domains someone trained; less variety than diffusion | Released checkpoints (StyleGAN family) |
  40| Train a GAN with the standard stabilisers | Non-saturating loss, a faster discriminator, an R1 penalty or spectral normalization | A one-pass generator for your own narrow domain | Two networks, coverage monitoring, learning-rate tuning; a run can still collapse | Your training loop |
  41| Wasserstein critic (WGAN-GP) | Replaces the verdict with a distance that still points home when real and fake don't overlap | A gradient wherever the generator is, and a loss that tracks quality | Several critic steps per generator step, plus a gradient penalty | Your training loop |
  42| Adversarial loss inside a decoder | A discriminator judges reconstructions against originals | Sharp detail where a reconstruction loss alone gives blur | One more network to train and keep stable | The training of VAEs, tokenizers and vocoders |
  43| Adversarial distillation of a diffusion model | A student learns to match a diffusion teacher in 1 to 4 steps, with a discriminator keeping it sharp | Real-time sampling from a foundation model | A teacher, a distillation run, some loss of variety | The fast sampling path of a diffusion system |
  44
  45**How to choose.** Start from the latency you need and the variety you
  46can't lose.
  47
  48- A new image, audio or video generator with no hard latency limit: use
  49  diffusion or flow matching, and put adversarial losses only in the decoder
  50  and in a distilled fast path.
  51- One-pass generation is non-negotiable and the domain is narrow: a GAN
  52  generator, trained with every stabiliser in the table, or a distilled
  53  diffusion model if a teacher exists for your domain.
  54- A blurry decoder: add a discriminator that compares reconstructions with
  55  originals, and expect the same instabilities as any GAN.
  56- A latent space to edit (age a face, change a pose): a style-based
  57  generator, whose latent space was designed for disentangled control.
  58- Whatever you pick, measure coverage, not just per-sample quality. Mode
  59  collapse produces beautiful samples that all look the same, and a metric
  60  that scores one sample at a time cannot see it.
  61
  62**What it costs.** Sampling is the GAN's strength: one generator pass per
  63sample. HiFi-GAN generates 22.05 kHz speech 167.9 times faster than real
  64time on one V100 GPU, and its small version runs 13.4 times faster than
  65real time on a CPU; a diffusion model of the 2021 generation needed 25
  66network passes per image to match BigGAN-deep. Training is where the cost
  67lies: two networks, one loss surface that moves every time the other
  68player steps, and a set of dials whose settings decide the outcome. On the
  69ring, 2,500 rounds with the discriminator learning at 0.003 against the
  70generator's 0.001 covers all eight clouds; at 0.001 it covers three and at
  710.01 it covers six. Each stabiliser has a price. An R1 gradient penalty
  72needs the gradient of the discriminator's own gradient, an extra backward
  73pass on every discriminator step; spectral normalization costs a couple of
  74matrix-vector products per layer per step; a Wasserstein critic is trained
  75for several steps per generator step. Quality and variety trade against
  76each other on every dial: the ten-times-faster discriminator gives 86%
  77on-target samples against the balanced run's 81%, and leaves two clouds
  78empty for good. BigGAN exposed the same trade as a knob, its truncation
  79trick, which trades sample variety for fidelity.
  80
  81**What breaks.**
  82
  83- **Mode collapse.** The generator covers a few kinds of data and hops
  84  between them as the discriminator catches up: on the ring, 3 clouds of 8,
  85  alternating between odd and even ones every 250 steps. Track coverage
  86  (FID or a per-class count) and rebalance the learning rates.
  87- **A silent gradient.** With the original generator loss, a confident
  88  discriminator hands the generator almost nothing: after an 800-step head
  89  start its verdict on fakes is 0.003 and the gradient 0.09, against 15 with
  90  the non-saturating loss. Use the non-saturating loss; every modern GAN
  91  does.
  92- **Oscillation.** Plain gradient steps rotate around the equilibrium
  93  rather than descending into it: in the two-number Dirac GAN, simultaneous
  94  steps drift from 1 away to 2.6 away in 300 steps. Losses that swing
  95  without trending are the symptom; an R1 penalty (0.08 away after 50
  96  steps) is the fix.
  97- **A discriminator that wins too fast.** Every region the generator has
  98  not reached becomes a wall of confident "fake", and the generator
  99  polishes what it has: six clouds from step 500 to the end, at ten times
 100  the generator's rate. A faster discriminator helps up to a point, then
 101  hurts.
 102- **No overlap, no direction.** When real and generated data don't overlap,
 103  the original objective is flat (Jensen-Shannon stuck at 0.693 whether the
 104  generator is 1 or 10 away). The Wasserstein critic's distance still
 105  slopes towards the data.
 106- **A seed that lies.** The same settings can find all eight clouds with one
 107  seed and three with another. Judge a recipe over several seeds.
 108
 109**In the wild.** StyleGAN (Karras, Laine and Aila) is the reference
 110one-pass image generator: a style-based generator that disentangles
 111high-level attributes from stochastic detail, trained on the FFHQ face
 112dataset the paper introduced, and trained with the R1 penalty.
 113BigGAN scaled class-conditional GANs to ImageNet with spectral
 114normalization, reaching an Inception Score of 166.5 and an FID of 7.4 at
 115128 × 128. FID itself, the standard score for generated images, came from
 116the two time-scale paper (Heusel et al.), along with the proof that
 117separate learning rates converge. Dhariwal and Nichol's *Diffusion Models
 118Beat GANs* marked the handover on image quality. Adversarial losses now
 119live inside other systems: the autoencoders of latent diffusion and
 120VQGAN's tokenizer are trained with a discriminator, HiFi-GAN and the neural
 121audio codecs use adversaries to keep waveforms clean, and Adversarial
 122Diffusion Distillation turns a diffusion model into a one-to-four-step
 123sampler by pairing score distillation with an adversarial loss. Every paper
 124is linked at the end of the lesson.
 125
 126**Go deeper.** Level 2 builds both players in NumPy, derives the best
 127possible discriminator and the game's equilibrium, shows why the original
 128generator loss goes silent, then reproduces each failure on the ring (the
 129Dirac GAN's spiral, mode collapse, the over-strong discriminator) and runs
 130each fix. If you only needed to choose, you are done.
 131
 132## Level 2: How it works, from scratch
 133
 134A forger wants to print banknotes that pass as real. A detective wants to
 135catch every fake. At first the forger is hopeless: smudged ink, the wrong
 136colour, and the detective spots every note. But every time the detective
 137rejects a note, the forger learns *what gave it away* and fixes that. Every
 138time the forger improves, the detective has to look more closely. Neither
 139is ever told what a real banknote should look like. They only push against
 140each other, and both keep getting better.
 141
 142The contest ends when the forger's notes are so good that the detective can
 143do no better than guess: "real" half the time, "fake" the other half. At
 144that point the forger has learned to make banknotes, and nobody ever wrote
 145down a rule for what a banknote is.
 146
 147That is a **generative adversarial network**, a **GAN**. The forger is a
 148neural network called the **generator**; the detective is a second network
 149called the **discriminator**. Train them against each other on photos of
 150faces and the generator learns to produce new faces that never existed.
 151
 152This lesson builds both players from scratch in NumPy, trains them on a toy
 153picture you can see at a glance, and then shows the three ways the contest
 154goes wrong: the players chase each other in circles, the forger settles for
 155copying a few examples (**mode collapse**), and the detective wins so
 156completely that the forger stops learning. Each failure gets a fix you can
 157run. The same goal, making new samples, is reached differently by
 158`primer.ml.generative.autoencoders` (squeeze and rebuild) and
 159`primer.ml.generative.diffusion` (remove noise a little at a time).
 160
 161## The two players
 162
 163**Everyday picture.** The forger is a recipe that turns a few dice rolls
 164into a banknote: different rolls, a different note. The detective is a
 165machine you feed a note into, and out comes a single number, how sure it is
 166that the note is real.
 167
 168**Tiny worked example.** Our "banknotes" are points on a flat page. The real
 169data is a **ring of eight little clouds**: eight centres evenly spaced on a
 170circle of radius 2, each real point scattered a tiny amount (a spread of
 1710.05) around one centre, chosen at random. This toy picture comes from the
 172research literature on GAN failures because you can *see* whether a forger
 173has found all eight clouds.
 174
 175- The **generator** G takes two random numbers z (the dice rolls, drawn from
 176  a bell curve) and returns one point on the page. It is a small neural
 177  network: 2 numbers in, two layers of 32 tanh units, 2 numbers out.
 178- The **discriminator** D takes a point and returns a probability that it is
 179  real. It is another small network: 2 numbers in, 64 tanh units, one
 180  **score** out, then squashed into the range 0 to 1.
 181
 182The squashing is the **sigmoid**, the same one `primer.ml.neural_net` uses to
 183turn a score into a probability:
 184
 185$$
 186D(x) = \sigma\big(a(x)\big) = \frac{1}{1 + e^{-a(x)}}
 187$$
 188
 189**Symbols**
 190
 191| Symbol | Meaning here | In the example |
 192|---|---|---|
 193| $x$ | a point on the page, real or forged | $(2, 0)$ |
 194| $a(x)$ | the detective's raw **score** for $x$ (also called a **logit**): any number, large and positive for "surely real" | $a(x) = x_1 - 1$ |
 195| $x_1$ | the first coordinate of $x$ | 2 |
 196| $e$ | Euler's number, about 2.718 (see `primer.notation`) | |
 197| $\sigma$ | the sigmoid: squashes any score into a probability between 0 and 1 | $\sigma(1) = 0.731$ |
 198| $D(x)$ | the detective's verdict: the probability that $x$ is real | 0.731 |
 199
 200**In words:** "the detective computes a score for the point, and the sigmoid
 201turns the score into a probability that the point is real."
 202
 203**With the numbers:** take a detective with a single neuron whose score is
 204$a(x) = x_1 - 1$. The point $(2, 0)$, one of the eight real centres, scores
 205$2 - 1 = 1$, and $\sigma(1) = 1 / (1 + e^{-1}) = 0.731$: probably real. The
 206point $(0, 0)$ in the middle of the ring scores $-1$ and gets
 207$\sigma(-1) = 0.269$: probably fake.
 208
 209**In Python:**
 210
 211```python
 212import math
 213def sigma(a):
 214    return 1 / (1 + math.exp(-a))
 215# a one-neuron detective: score a(x) = x_1 - 1
 216def a(x):
 217    return x[0] - 1
 218round(sigma(a((2, 0))), 3)  # → 0.731
 219round(sigma(a((0, 0))), 3)  # → 0.269
 220```
 221
 222```mermaid
 223flowchart LR
 224  Z["noise z<br/>2 random numbers"] --> G["Generator G<br/>the forger"]
 225  G --> F["fake point G(z)"]
 226  R["real point x<br/>from the ring"] --> D["Discriminator D<br/>the detective"]
 227  F --> D
 228  D --> P["D(point)<br/>probability it is real"]
 229  P -. "learns to call real real<br/>and fake fake" .-> D
 230  P -. "learns to make D say real<br/>about its fakes" .-> G
 231```
 232
 233**Reading it:** start at the far left. Noise goes into the generator and
 234comes out as a fake point. Real points come in from the data. Both kinds of
 235point go through the same detective, which gives each one a probability.
 236The two dotted arrows are the learning signals, and they pull in opposite
 237directions: the detective adjusts itself to score real points high and fakes
 238low, and the generator adjusts itself to make the detective score *its*
 239points high. Notice that the generator never sees a real point. Everything
 240it learns about the data arrives through the detective's verdicts.
 241
 242**Why it matters:** that last point is the whole trick. Nobody writes down
 243what a face, a voice or a banknote is. The detective discovers what separates
 244real from fake, and its gradient (the direction that would make it less
 245sure a fake is fake; see `primer.ml.neural_net` for gradients) tells the
 246forger how to improve.
 247
 248**In code:** `ring_of_gaussians` draws real points around `ring_modes`, `MLP` is both players (its `MLP.backward` also returns the gradient with respect to the input, the channel through which the forger learns), with shapes `GENERATOR_SIZES` and `DISCRIMINATOR_SIZES`, and `sigmoid` turns the detective's score into a probability.
 249
 250## The game: one number both players fight over
 251
 252**Everyday picture.** Picture a scoreboard. The detective earns points for
 253confident, correct calls and loses points, heavily, for confident mistakes.
 254The detective wants the score as high as possible. The forger wants it as
 255low as possible. One number, two players pulling it in opposite directions:
 256that is a **minimax** game.
 257
 258**Tiny worked example.** The detective looks at two real notes and two
 259fakes. The score uses the **logarithm** (log): log 1 = 0, and the log of a
 260small number is a large negative number, so a confident mistake costs far
 261more than a hesitant one (see `primer.notation` for logs from scratch).
 262
 263| Note | Real? | Verdict D | What is scored | Value |
 264|---|---|---|---|---|
 265| 1 | real | 0.9 | log D = log 0.9 | −0.105 |
 266| 2 | real | 0.8 | log D = log 0.8 | −0.223 |
 267| 3 | fake | 0.2 | log (1 − D) = log 0.8 | −0.223 |
 268| 4 | fake | 0.4 | log (1 − D) = log 0.6 | −0.511 |
 269
 270The average over the real notes is −0.164, over the fakes −0.367, and the
 271total is **−0.53**. A perfect detective (1 on every real note, 0 on every
 272fake) would score log 1 + log 1 = 0, the ceiling. A detective reduced to a
 273coin flip (0.5 on everything) scores log 0.5 + log 0.5 = −1.386.
 274
 275$$
 276\min_G \, \max_D \; V(D, G) = \mathbb{E}_{x \sim p_{\text{data}}}\big[\log D(x)\big] + \mathbb{E}_{z \sim p_z}\big[\log\big(1 - D(G(z))\big)\big]
 277$$
 278
 279**Symbols**
 280
 281| Symbol | Meaning here | In the example |
 282|---|---|---|
 283| $V(D, G)$ | the **value**: the scoreboard number | −0.53 |
 284| $\max_D$ | "the detective chooses its weights to make what follows as large as possible" | |
 285| $\min_G$ | "the forger chooses its weights to make that best-case value as small as possible" | |
 286| $\mathbb{E}$ | **expectation**: the average over many draws (see `primer.notation`, probability notation) | an average over 2 notes |
 287| $x \sim p_{\text{data}}$ | "$x$ drawn from the real data" | notes 1 and 2 |
 288| $z \sim p_z$ | "$z$ drawn from the noise the forger starts from" | the dice rolls behind notes 3 and 4 |
 289| $D(x)$ | the verdict on a real point | 0.9, 0.8 |
 290| $G(z)$ | a fake point made from noise $z$ | notes 3 and 4 |
 291| $D(G(z))$ | the verdict on a fake | 0.2, 0.4 |
 292| $\log$ | natural logarithm: 0 at 1, very negative near 0 | $\log 0.6 = -0.511$ |
 293
 294**In words:** "average the log of the verdicts on real points, add the
 295average log of one minus the verdicts on fakes; the detective pushes this
 296number up, the forger pushes it down."
 297
 298**With the numbers:** (log 0.9 + log 0.8) / 2 + (log 0.8 + log 0.6) / 2 =
 299−0.164 + (−0.367) = −0.53.
 300
 301**In Python:**
 302
 303```python
 304import math
 305# D(x) on two real notes, D(G(z)) on two fakes
 306d_real, d_fake = [0.9, 0.8], [0.2, 0.4]
 307# E over real x of log D(x): an average
 308real_term = sum(math.log(d) for d in d_real) / len(d_real)
 309round(real_term, 3)  # → -0.164
 310# E over noise z of log(1 - D(G(z)))
 311fake_term = sum(math.log(1 - d) for d in d_fake) / len(d_fake)
 312round(fake_term, 3)  # → -0.367
 313round(real_term + fake_term, 2)  # → -0.53
 314# a coin-flip detective, D = 1/2 on everything
 315round(math.log(0.5) + math.log(0.5), 3)  # → -1.386
 316```
 317
 318If this looks familiar, it should: −V is exactly the **binary
 319cross-entropy** of a classifier that labels real points 1 and fakes 0 (see
 320`primer.ml.losses`). The detective is an ordinary classifier. What is new
 321is that its second class, the fakes, keeps changing underneath it.
 322
 323Training alternates. Nobody can solve "min over G of max over D" directly,
 324so each round takes one small step for each player:
 325
 326```mermaid
 327flowchart TB
 328  S["sample a batch of real points<br/>and a batch of noise"] --> F["forger makes fakes G(z)"]
 329  F --> DS["detective step:<br/>climb V, so D(real) rises<br/>and D(fake) falls"]
 330  DS --> GS["forger step:<br/>change G so the new detective<br/>scores its fakes higher"]
 331  GS --> CHK{"trained enough?"}
 332  CHK -- no --> S
 333  CHK -- yes --> OUT["keep G;<br/>throw D away"]
 334```
 335
 336**Reading it:** follow one lap from the top. A fresh batch of real points
 337and noise comes in, the forger turns the noise into fakes, and then the two
 338players move in turn: first the detective takes one step of gradient
 339**ascent** on V (it wants V higher), then the forger takes one step of
 340gradient **descent** (it wants V lower), judged by the detective as it is
 341*after* its step. The loop repeats thousands of times. At the bottom, only
 342the generator is kept. The detective was scaffolding: its whole job was to
 343teach.
 344
 345**Why it matters:** each player is trained with an ordinary optimizer on an
 346ordinary loss, but the loss surface moves every time the other player steps.
 347That is the root of everything that goes wrong later in this lesson.
 348
 349**In code:** `value_function` computes V from a batch of verdicts, and `train_gan` runs the loop above on the ring, one detective step then one forger step, each with Adam (`primer.ml.optimizers.Adam`).
 350
 351## The best possible detective, and where the game ends
 352
 353**Everyday picture.** Suppose that, at one spot on the page, real points
 354turn up three times as often as fakes. However clever the detective is, it
 355cannot tell two identical-looking points apart, so the best it can do there
 356is to say "75% real". It should match the local mix, no more and no less.
 357
 358**Tiny worked example.** At a point where the real data's **density** (how
 359thickly its points cover that spot) is 0.3 and the forger's is 0.1, the
 360best verdict is 0.3 / (0.3 + 0.1) = 0.75. Try its neighbours: the
 361detective's expected score at that spot is
 3620.3 · log D + 0.1 · log(1 − D), which is −0.2274 at D = 0.7, −0.2249 at
 363D = 0.75 and −0.2279 at D = 0.8. The middle one is the highest.
 364
 365$$
 366D^*(x) = \frac{p_{\text{data}}(x)}{p_{\text{data}}(x) + p_g(x)}
 367$$
 368
 369**Symbols**
 370
 371| Symbol | Meaning here | In the example |
 372|---|---|---|
 373| $D^*(x)$ | the best possible verdict at $x$, for a forger that is held fixed | 0.75 |
 374| $p_{\text{data}}(x)$ | how densely the real data covers the spot $x$ | 0.3 |
 375| $p_g(x)$ | how densely the forger's fakes cover the spot $x$ | 0.1 |
 376
 377**In words:** "the best verdict at each point is the share of the points
 378found there that are real."
 379
 380**With the numbers:** 0.3 / (0.3 + 0.1) = 0.75. If the forger matched the
 381data exactly, $p_g = p_{\text{data}}$ everywhere, and every verdict would be
 3820.3 / (0.3 + 0.3) = 0.5.
 383
 384**In Python:**
 385
 386```python
 387import math
 388p_data, p_g = 0.3, 0.1
 389round(p_data / (p_data + p_g), 2)  # → 0.75
 390# the detective's expected score at this spot, for a few verdicts D
 391def score(D):
 392    return p_data * math.log(D) + p_g * math.log(1 - D)
 393[round(score(D), 4) for D in (0.7, 0.75, 0.8)]  # → [-0.2274, -0.2249, -0.2279]
 394# a forger that matches the data exactly
 395round(0.3 / (0.3 + 0.3), 2)  # → 0.5
 396```
 397
 398Where does the formula come from? At each point the detective is choosing
 399one number D to maximise $p_{\text{data}} \log D + p_g \log(1 - D)$. The
 400slope of that with respect to D is $p_{\text{data}}/D - p_g/(1 - D)$, and
 401setting the slope to zero gives $D^*$. This is the central result of the
 402original GAN paper, and it tells you where the game ends. **At the
 403equilibrium the forger's distribution equals the data's, and the best
 404detective says 1/2 everywhere**, scoring V = −log 4 ≈ −1.386, the
 405coin-flip score from the table above. Plug $D^*$ back into V and what
 406remains is −log 4 plus twice the **Jensen-Shannon divergence** between the
 407real and forged distributions, a measure of how different two piles of
 408probability are that is zero only when they are identical. So a forger
 409facing a perfect detective is really minimising that divergence. Keep that
 410in mind: it returns as a problem in the fixes section.
 411
 412![Top: the real data has two bumps and the forger one wide bump; bottom: the best detective's verdict rises to about 0.8 on the real bumps, drops towards 0 where only fakes live, and would be flat at one half if the forger matched](figures/primer.ml.generative.gans.optimal_detective.svg)
 413
 414**Reading it:** the top panel shows two distributions on a line: the real
 415data (blue) has two narrow bumps at −2 and +2, and the forger (red) spreads
 416one wide bump over the middle. The bottom panel is the best verdict
 417$D^*$ at every position. Over the real bumps it rises to about 0.8,
 418because most points found there are real (the forger's wide bump still
 419reaches them). In the middle and at the far edges it falls towards 0,
 420because only fakes live there. The grey dashed line at 0.5 is what $D^*$
 421becomes once the forger matches the data: the detective has nothing left to
 422go on. The shape of the bottom curve is exactly the information the forger
 423needs, which way to move its mass.
 424
 425![After training with balanced learning rates, the forger's green points sit on all eight clouds, and the detective's verdict on the ring is close to one half](figures/primer.ml.generative.gans.detective_view.svg)
 426
 427**Reading it:** this is the real ring after 2,500 rounds of training. The
 428background colour is the trained detective's verdict at every spot: blue for
 429"real", red for "fake", white for 0.5. Black dots are real points, green
 430dots are the forger's. The green points sit on all eight clouds, and the
 431background there is pale, close to white: at the eight centres the verdict
 432is between 0.46 and 0.58. That is the equilibrium made visible, the detective
 433reduced to guessing where the forgery is good. Away from the ring, where
 434neither real nor fake points ever appear, the formula says nothing, and the
 435detective's colouring there is arbitrary (deep red in the empty middle).
 436
 437**Why it matters:** the detective never needs to model what real data looks
 438like; it only estimates a *ratio* between two densities. That is why GANs
 439could learn sharp images long before anyone could write down a probability
 440for an image.
 441
 442**In code:** `optimal_discriminator` is the formula, `fit_discriminator_table` trains the most flexible detective possible (one free score per point) against a fixed forger and arrives at the same verdicts without being told the formula, and `detective_verdicts` evaluates a trained detective anywhere on the page.
 443
 444## Why the forger uses a different loss
 445
 446**Everyday picture.** Early in training, the forger is terrible and the
 447detective is sure every fake is fake. Now picture a game of "hotter,
 448colder" where the helper stops talking once you are far enough away: the
 449further off you are, the quieter the hints. That is the original forger
 450loss. The fix is a helper who shouts louder the further off you are.
 451
 452**Tiny worked example.** The detective looks at a fake and is 99% sure it
 453is fake: D(G(z)) = 0.01. How hard does each loss push the forger?
 454
 455- The original loss, log(1 − D(G(z))), which the forger minimises: its slope
 456  with respect to the detective's score is −D = **−0.01**. Almost nothing.
 457- The **non-saturating** loss, −log D(G(z)), which the forger also
 458  minimises: its slope is −(1 − D) = **−0.99**. Ninety-nine times more.
 459
 460When the detective is undecided, D = 0.5, both slopes are −0.5. The two
 461losses agree when it doesn't matter and differ exactly when it does.
 462
 463$$
 464\frac{\partial}{\partial a}\,\log\big(1 - \sigma(a)\big) = -\sigma(a) = -D
 465\qquad\qquad
 466\frac{\partial}{\partial a}\,\big(-\log \sigma(a)\big) = -\big(1 - \sigma(a)\big) = -(1 - D)
 467$$
 468
 469**Symbols**
 470
 471| Symbol | Meaning here | In the example |
 472|---|---|---|
 473| $a$ | the detective's score for a fake, before the sigmoid | $\log(0.01/0.99) = -4.6$ |
 474| $\sigma(a)$ | the verdict D on that fake | 0.01 |
 475| $\frac{\partial}{\partial a}$ | the **derivative** with respect to $a$: how much the loss changes when the score is nudged up a tiny amount (see `primer.notation`) | |
 476| $\log(1 - \sigma(a))$ | the original, "saturating" forger loss | $\log 0.99$ |
 477| $-\log \sigma(a)$ | the non-saturating forger loss | $-\log 0.01 = 4.6$ |
 478
 479**In words:** "the original loss changes by only D when the detective's
 480score moves, so it goes quiet when D is near 0; the non-saturating loss
 481changes by 1 − D, so it is loudest exactly when the forger is doing worst."
 482
 483**With the numbers:** at D = 0.01, the slopes are −0.01 and −0.99.
 484
 485**In Python:**
 486
 487```python
 488import math
 489def sigma(a):
 490    return 1 / (1 + math.exp(-a))
 491# a score that makes the detective 99% sure the fake is fake
 492a = math.log(0.01 / 0.99)
 493round(sigma(a), 2)  # → 0.01
 494# measure each slope by nudging the score a tiny bit either way
 495h = 1e-6
 496def slope(loss):
 497    return (loss(a + h) - loss(a - h)) / (2 * h)
 498def saturating(a):
 499    return math.log(1 - sigma(a))
 500def non_saturating(a):
 501    return -math.log(sigma(a))
 502round(slope(saturating), 3)  # → -0.01
 503round(slope(non_saturating), 3)  # → -0.99
 504```
 505
 506Both losses want the same thing, a higher verdict on the fakes, and they
 507share the same equilibrium. Only the strength of the push differs. The
 508original GAN paper already recommended the switch, and essentially every
 509GAN since uses the non-saturating loss or a Wasserstein-style loss (below).
 510
 511![Left: as the detective grows sure a fake is fake, the original loss flattens out while the non-saturating loss keeps a steady slope; right: giving the detective an 800-step head start shrinks the original loss's gradient to 0.09 while the non-saturating one grows to 15](figures/primer.ml.generative.gans.forger_losses.svg)
 512
 513**Reading it:** on the left, the x-axis is the detective's score for a
 514fake: far left means "certainly fake". The red curve, the original loss,
 515goes flat as you move left, and a flat curve has no slope, so no gradient.
 516The blue curve, the non-saturating loss, keeps a steady downhill slope all
 517the way. On the right is a measurement from our ring: the detective trains
 518alone against a frozen, untrained forger for more and more steps, and we
 519measure how big a gradient each loss would hand the forger (log scale).
 520At the start both are about 0.5 to 0.6. After 800 detective steps its
 521verdict on fakes is 0.003, the original loss's gradient has shrunk to 0.09,
 522and the non-saturating one has grown to 15: over a hundred times larger.
 523
 524**Why it matters:** a forger that gets no gradient does not learn, and the
 525detective's lead only widens. Early in training the detective *always* has
 526the easy job, so this failure is the default, not an edge case.
 527
 528**In code:** `generator_logit_gradient` returns both slopes, and `head_start_experiment` gives the detective a head start and measures the forger's gradient under each loss.
 529
 530## Why training is unstable, part 1: the players chase each other in circles
 531
 532**Everyday picture.** Two people share an office thermostat. One turns the
 533heat up whenever it feels cold; the other opens a window whenever it feels
 534hot. Each reacts to where the room *was*, so the temperature overshoots one
 535way, then the other, and never settles. A GAN's players do the same: the
 536forger moves to where the detective currently says "real", and by the time
 537it arrives, the detective has moved.
 538
 539**Tiny worked example.** The smallest GAN there is, the **Dirac GAN**, has
 540one number per player. The real data is a single point at x = 0. The forger
 541has one number θ and always outputs x = θ. The detective has one number ψ
 542and scores a point as ψ · x, so D(x) = σ(ψx). Start with the forger at
 543θ = 1 and a detective with ψ = 0 (it cannot tell anything apart), and let
 544both take steps of size 1 at the same time:
 545
 546| Step | σ(ψθ) | θ (forger) | ψ (detective) | Distance from (0, 0) |
 547|---|---|---|---|---|
 548| 0 | | 1 | 0 | 1 |
 549| 1 | σ(0) = 0.5 | 1 | −0.5 | 1.118 |
 550| 2 | σ(−0.5) = 0.3775 | 0.811 | −0.878 | 1.195 |
 551
 552After step 1 the detective has learned "bigger x means fake" (ψ turned
 553negative), and after step 2 the forger has moved towards the data at 0.
 554Both moves were sensible, yet the pair got *further* from the equilibrium,
 555the point (0, 0) where the forger sits on the data and the detective is
 556indifferent.
 557
 558$$
 559V(\theta, \psi) = \log \sigma(\psi \cdot 0) + \log\big(1 - \sigma(\psi\theta)\big)
 560\qquad
 561\frac{\partial V}{\partial \psi} = -\sigma(\psi\theta)\,\theta
 562\qquad
 563\frac{\partial V}{\partial \theta} = -\sigma(\psi\theta)\,\psi
 564$$
 565
 566**Symbols**
 567
 568| Symbol | Meaning here | In the example |
 569|---|---|---|
 570| $\theta$ | the forger's only number: where it puts its fake | 1 |
 571| $\psi$ | the detective's only number: the slope of its score $\psi x$ | 0 |
 572| $\psi \cdot 0$ | the detective's score for the real data at $x = 0$: always 0, so $D(0) = 0.5$ | 0 |
 573| $\sigma(\psi\theta)$ | the verdict on the fake | 0.5 |
 574| $\frac{\partial V}{\partial \psi}$ | the slope the detective climbs: $\psi \leftarrow \psi + \eta\,\frac{\partial V}{\partial \psi}$ | $-0.5$ |
 575| $\frac{\partial V}{\partial \theta}$ | the slope the forger descends: $\theta \leftarrow \theta - \eta\,\frac{\partial V}{\partial \theta}$ | 0 |
 576| $\eta$ | the step size (learning rate) | 1 |
 577| $\leftarrow$ | "is replaced by" | |
 578
 579**In words:** "the detective tilts its slope against wherever the fake sits;
 580the forger slides along whichever way the slope says is more real; each
 581reacts to the other's old position."
 582
 583**With the numbers:** step 1: σ(0) = 0.5, so ψ becomes 0 − 1 · 0.5 · 1 =
 584−0.5 and θ stays 1 + 1 · 0.5 · 0 = 1. Step 2: σ(−0.5) = 0.3775, so θ
 585becomes 1 + 0.3775 · (−0.5) = 0.811 and ψ becomes −0.5 − 0.3775 · 1 = −0.878.
 586
 587**In Python:**
 588
 589```python
 590import math
 591def sigma(a):
 592    return 1 / (1 + math.exp(-a))
 593theta, psi, lr = 1.0, 0.0, 1.0
 594s = sigma(psi * theta)
 595s  # → 0.5
 596# both move at once, each using the other's old position
 597theta, psi = theta + lr * s * psi, psi - lr * s * theta
 598(theta, psi)  # → (1.0, -0.5)
 599s = sigma(psi * theta)
 600round(s, 4)  # → 0.3775
 601theta, psi = theta + lr * s * psi, psi - lr * s * theta
 602(round(theta, 3), round(psi, 3))  # → (0.811, -0.878)
 603# distance from the equilibrium (0, 0): 1, then 1.118, then 1.195
 604round(math.hypot(theta, psi), 3)  # → 1.195
 605```
 606
 607```mermaid
 608flowchart LR
 609  A["fake sits to the right<br/>of the data"] --> B["detective tilts:<br/>right means fake"]
 610  B --> C["forger slides left,<br/>overshoots past the data"]
 611  C --> D["detective tilts back:<br/>left means fake"]
 612  D --> E["forger slides right,<br/>overshoots again"]
 613  E --> A
 614```
 615
 616**Reading it:** the boxes form a loop because the game has no downhill
 617direction that both players share. Each arrow is a sensible move for the
 618player making it, but each player only reacts to the other's current
 619position, so the forger arrives just after the detective has changed its
 620mind. The loop is a rotation around the equilibrium, not a descent into it.
 621
 622![In the plane of the forger's theta against the detective's psi, simultaneous steps spiral outward from the start at (1, 0), and alternating steps circle the equilibrium without ever reaching it](figures/primer.ml.generative.gans.dirac_oscillation.svg)
 623
 624**Reading it:** the horizontal axis is the forger's θ, the vertical axis
 625the detective's ψ, and the black cross at (0, 0) is the equilibrium. Both
 626paths start at the dot, (1, 0), and take 300 steps of size 0.2. The red
 627path, where both players step at once, spirals *outwards*: after 300 steps
 628it is 2.6 away from the equilibrium, where it started 1 away. The blue path,
 629where the forger waits to see the detective's new ψ before it moves, stays
 630on a closed loop about 1 away forever. Neither ever arrives. Plain gradient
 631steps do not solve this game even in two dimensions.
 632
 633**Why it matters:** a real GAN has millions of numbers per player, and this
 634rotation happens in many directions at once. It shows up as losses that
 635swing up and down without trending, and samples that keep changing
 636character instead of steadily improving.
 637
 638**In code:** `dirac_gan` runs this two-number game with simultaneous or alternating steps.
 639
 640## Why training is unstable, part 2: mode collapse
 641
 642**Everyday picture.** The forger discovers that one particular note, say
 643the twenty, always fools the detective. So it prints nothing but twenties.
 644The detective eventually learns that twenties are suspicious, so the forger
 645switches to printing only fifties, and so on. Each forgery is good, but the
 646forger never produces the full *variety* of real money. Each distinct kind
 647of real example is a **mode**, and settling on a few of them is **mode
 648collapse**.
 649
 650**Tiny worked example.** To measure variety on the ring, generate 1,000
 651points and count, for each of the eight clouds, how many land within 0.15
 652of its centre (3 spreads). A cloud is **covered** if it gets at least 20
 653of the 1,000 (2%; a perfect forger gives each about 125). A point that lands
 654near any centre counts as **high quality**. After 2,500 rounds with equal
 655learning rates for both players, the counts for the eight clouds are
 656(0, 68, 0, 20, 0, 51, 0, 17). Clouds 2, 4 and 6 are covered, cloud 8's 17
 657falls short, so **3 of 8** modes are covered, and only 156 of the 1,000
 658points are high quality: the rest are strung between clouds.
 659
 660```mermaid
 661stateDiagram-v2
 662  direction LR
 663  EvenClouds: forger piles onto clouds 1, 3, 5, 7
 664  CatchEven: detective learns those clouds are suspicious
 665  OddClouds: forger jumps to clouds 2, 4, 6, 8
 666  CatchOdd: detective learns these are suspicious
 667  EvenClouds --> CatchEven
 668  CatchEven --> OddClouds
 669  OddClouds --> CatchOdd
 670  CatchOdd --> EvenClouds
 671```
 672
 673**Reading it:** each box is a phase of training, and the arrows go round in
 674a circle. The forger does not have to cover every cloud to fool the
 675detective *today*; it only has to put its points where the detective
 676currently says "real". The detective then learns those spots are
 677over-supplied with fakes, the verdict there drops, and the forger's
 678gradient pushes the whole pile somewhere else. It is the chase from part 1,
 679now played out across the clouds of the data. The alternation between odd
 680and even clouds is what our run actually does, pictured next.
 681
 682![Three rows of four snapshots between steps 1,750 and 2,500: with equal learning rates the green points hop between alternate clouds; with a detective three times faster they cover all eight; with a detective ten times faster they sit on six and never reach the other two](figures/primer.ml.generative.gans.ring_snapshots.svg)
 683
 684**Reading it:** each panel is the page, grey circles mark the eight real
 685clouds, and green dots are 1,000 points from the forger at that step (always
 686made from the same noise, so the panels are comparable). The panel titles
 687spell the coverage with X for a covered cloud, counting anticlockwise from
 688the cloud on the right. Top row, equal learning rates: the green points hop
 689from clouds 1, 3 and 7 at step 1,750 to 2, 4, 6 and 8 at step 2,000, back
 690to 1, 3 and 5, then to 2, 4 and 6, never holding more than half the ring.
 691Middle row, a detective that learns three times faster than the forger: all
 692eight clouds, every snapshot. Bottom row, a detective ten times faster: six
 693clouds, and the same six at every snapshot. Two clouds stay empty for good.
 694
 695![Modes covered over training: equal learning rates wander between one and four; a three-times-faster detective reaches all eight by step 1,000 and stays; a ten-times-faster detective locks onto six at step 500](figures/primer.ml.generative.gans.mode_coverage.svg)
 696
 697**Reading it:** the x-axis is the training step, the y-axis the number of
 698clouds covered out of eight. The blue line (detective three times faster)
 699climbs to 8 by step 1,000 and stays there. The red line (equal speeds)
 700zigzags between 1 and 4, the hopping from the top row above. The amber line
 701(detective ten times faster) jumps to 6 by step 500 and never moves again.
 702The same code and the same data give three very different outcomes, and the
 703only thing that changed is how fast the detective learns.
 704
 705These runs use one fixed seed. With other seeds the equal-speed run
 706sometimes finds all eight clouds late on, and the ten-times run sometimes
 707escapes after a while; the pattern, not the exact counts, is the lesson.
 708
 709**Why it matters:** mode collapse is the classic GAN failure on real data. A
 710face generator that has collapsed produces beautiful faces that all look
 711like the same few people. Quality metrics that look at one sample at a time
 712can't see it; you have to measure *coverage*, as we did here.
 713
 714**In code:** `mode_coverage` counts covered clouds and high-quality points, and `train_gan` records it every 250 steps.
 715
 716## Why training is unstable, part 3: a detective that wins too fast
 717
 718**Everyday picture.** A detective who learns much faster than the forger
 719soon rejects everything outside the few spots the forger has already
 720mastered, with total confidence. Every small experiment the forger tries in
 721a new direction comes back "more fake", so it retreats to what already works
 722and stops exploring.
 723
 724**Tiny worked example.** The bottom row above: with the detective's learning
 725rate ten times the forger's, after 2,500 steps the counts per cloud are
 726(199, 181, 117, 0, 132, 160, 72, 0). The six covered clouds are served well:
 72786% of points are high quality, *more* than the balanced run's 81%. But
 728clouds 4 and 8 get nothing, from step 500 until the end. Better-looking
 729samples, less variety: a trade-off you will meet again with real image
 730generators.
 731
 732```mermaid
 733flowchart LR
 734  A["detective learns<br/>much faster"] --> B["gaps between the forger's clouds<br/>become confident fake zones"]
 735  B --> C["every small step towards them<br/>lowers the verdict"]
 736  C --> D["forger polishes the clouds<br/>it already has"]
 737  D --> E["empty clouds stay empty"]
 738```
 739
 740**Reading it:** read it as a chain of causes. A fast detective turns every
 741region the forger isn't already covering into a wall of confident "fake"
 742verdicts. To reach an empty cloud the forger would have to move points
 743*through* such a wall. But a gradient only sees the next small step, and
 744every small step into the wall makes the verdict worse, so the gradient
 745pushes the points back onto the clouds they already occupy. The forger
 746spends its effort sharpening those, and the missing ones are never
 747discovered.
 748
 749**Why it matters:** "just train the detective harder" sounds like it should
 750give the forger a better teacher. Past a point it gives it a worse one. The
 751detective needs to be good enough to point the way, not so good that it only
 752says no.
 753
 754**In code:** `train_gan` takes both learning rates as arguments; this run gives the detective ten times the forger's.
 755
 756## Fixes that make the game trainable
 757
 758**Everyday picture.** A good sparring partner matters more than a strong
 759one. Every stabiliser below is a way of keeping the detective useful: fast
 760enough to teach, smooth enough that its verdicts point somewhere.
 761
 762```mermaid
 763flowchart LR
 764  P1["players out of step"] --> F1["careful learning rates<br/>two time-scales"]
 765  P2["detective builds<br/>steep cliffs"] --> F2["gradient penalty"]
 766  P2 --> F3["spectral normalization"]
 767  P3["piles don't overlap,<br/>so the verdict gives no direction"] --> F4["Wasserstein distance"]
 768  F4 --> F2
 769```
 770
 771**Reading it:** failures are on the left and the fixes that target them on
 772the right. Players moving out of step are fixed by choosing their speeds.
 773A detective whose verdict jumps from 0 to 1 across a narrow cliff is fixed
 774by penalising or capping its steepness. A yardstick that can't tell "near
 775miss" from "far miss" is replaced by the Wasserstein distance, which in
 776practice needs one of the steepness fixes to compute, hence the last arrow.
 777
 778### Careful learning rates
 779
 780**Everyday picture.** Pace the two sparring partners so that neither runs
 781away with the match.
 782
 783**Tiny worked example.** The three rows of the ring figure differ only in
 784the detective's learning rate, with the forger's fixed at 0.001:
 785
 786| Detective's rate | Clouds covered at step 2,500 | High-quality points |
 787|---|---|---|
 788| 0.001 (equal) | 3, and hopping | 16% |
 789| 0.003 (three times) | **8**, steady | 81% |
 790| 0.01 (ten times) | 6, stuck | 86% |
 791
 792A modestly faster detective keeps its verdicts close to the best-detective
 793formula for the forger it is facing *now*, so the forger follows a gradient
 794that points at the data rather than at yesterday's detective. Heusel and
 795colleagues called this the **two time-scale update rule** and proved it
 796reaches an equilibrium under reasonable conditions. The right ratio depends
 797on the problem; the lesson is that it is a dial worth turning.
 798
 799**In code:** `train_gan` with a detective learning rate of 0.003 is the balanced run; the other two rows are the same call with 0.001 and 0.01.
 800
 801### Gradient penalty: fine the detective for steep cliffs
 802
 803**Everyday picture.** Charge the detective a fee for how sharply its score
 804changes right on top of the real data. A detective with gentle slopes still
 805ranks real above fake, but it can no longer swing wildly, and the swinging is
 806what fed the spiral in part 1.
 807
 808**Tiny worked example.** In the Dirac GAN the detective's score is ψ · x,
 809whose slope with respect to x is just ψ. The **R1 gradient penalty** adds
 810(γ / 2) · ψ² to what the detective minimises. With γ = 1 and ψ = −0.5 the
 811fee is 0.125, and its gradient, γψ = −0.5, pulls ψ back towards 0 on every
 812step, like friction. With step size 0.2 that friction shrinks ψ by a factor
 8131 − 0.2 · 1 = 0.8 per step, on top of the game's own push.
 814
 815$$
 816R_1 = \frac{\gamma}{2}\;\mathbb{E}_{x \sim p_{\text{data}}}\Big[\big\lVert \nabla_x\, a(x) \big\rVert^2\Big]
 817$$
 818
 819**Symbols**
 820
 821| Symbol | Meaning here | In the Dirac example |
 822|---|---|---|
 823| $R_1$ | the penalty added to the detective's loss | 0.125 |
 824| $\gamma$ | how heavily steepness is fined | 1 |
 825| $a(x)$ | the detective's score at $x$ | $\psi x$ |
 826| $\nabla_x\, a(x)$ | the **gradient** of the score with respect to the input point: which way, and how steeply, the score rises as you move $x$ | $\psi = -0.5$ |
 827| $\lVert \cdot \rVert^2$ | the squared length of that gradient (see `primer.notation`, the length of a vector) | 0.25 |
 828| $\mathbb{E}_{x \sim p_{\text{data}}}$ | averaged over real points only | the single real point $x = 0$ |
 829
 830**In words:** "measure how steeply the detective's score changes at the real
 831points, square it, average it, and charge the detective γ/2 times that."
 832
 833**With the numbers:** (1 / 2) · (−0.5)² = 0.125.
 834
 835**In Python:**
 836
 837```python
 838gamma, psi = 1.0, -0.5
 839# slope of the score psi * x with respect to x, at the real point
 840slope = psi
 841gamma / 2 * slope ** 2  # → 0.125
 842# the penalty's pull on psi: its derivative, gamma * psi
 843gamma * psi  # → -0.5
 844# one step of size 0.2 against that pull shrinks psi by a factor 0.8
 845round(psi - 0.2 * gamma * psi, 2)  # → -0.4
 846```
 847
 848![Distance from the equilibrium over 300 steps: simultaneous steps grow from 1 to 2.6, alternating steps hold near 1, and with the R1 penalty the distance falls to 0.08 within 50 steps and to zero after](figures/primer.ml.generative.gans.dirac_r1.svg)
 849
 850**Reading it:** the x-axis is the step, the y-axis the distance of (θ, ψ)
 851from the equilibrium (0, 0). The red line (plain simultaneous steps) climbs
 852in steps, one per lap: the outward spiral. The blue line (alternating steps) hovers near
 8531 forever: the endless orbit. The green line adds the R1 penalty with γ = 1
 854to the simultaneous steps and dives: 0.08 after 50 steps, essentially 0 by
 855100. The friction turns the rotation into a spiral *inwards*. That is the
 856Mescheder, Geiger and Nowozin result in two numbers, and it is why StyleGAN
 857and many later GANs train with an R1 penalty. The Wasserstein GAN with
 858gradient penalty (WGAN-GP) uses the same idea with a different target: it
 859keeps the slope near 1 at points between real and fake.
 860
 861**In code:** `dirac_gan` takes the penalty weight γ as an argument and adds the R1 penalty's pull, −γψ, to the detective's step.
 862
 863### Wasserstein distance: a yardstick that knows near from far
 864
 865**Everyday picture.** Picture the real data and the fakes as two piles of
 866sand. The **Wasserstein distance**, also called the **earth mover's
 867distance**, is the least work needed to reshape one pile into the other:
 868the amount of sand moved times how far it travels. The Jensen-Shannon
 869divergence from the best-detective section only asks how much the piles
 870*overlap*.
 871
 872**Tiny worked example.** Put all the real sand at x = 0 and all the fake sand
 873at x = θ. If θ = 1, the piles don't overlap, and the Jensen-Shannon
 874divergence is log 2 = 0.693. If θ = 5, they still don't overlap, and it is
 875still 0.693. It cannot tell a near miss from a far one, so it gives the
 876forger no direction. The Wasserstein distance is 1 for θ = 1 and 5 for
 877θ = 5: it shrinks as the forger approaches, so it always points home.
 878
 879For two equally sized samples on a line, the cheapest way to move the sand
 880is to pair the smallest with the smallest, the next with the next, and so
 881on:
 882
 883$$
 884W_1 = \frac{1}{n} \sum_{i=1}^{n} \big\lvert a_{(i)} - b_{(i)} \big\rvert
 885$$
 886
 887**Symbols**
 888
 889| Symbol | Meaning here | In the example |
 890|---|---|---|
 891| $W_1$ | the Wasserstein (earth mover's) distance between two piles on a line | 3 |
 892| $n$ | how many grains (samples) each pile has | 3 |
 893| $a_{(i)}$ | the $i$-th smallest value in the first pile (the brackets mean "after sorting") | $a = (0, 1, 2)$ |
 894| $b_{(i)}$ | the $i$-th smallest value in the second pile | $b = (5, 3, 4)$ sorts to $(3, 4, 5)$ |
 895| $\lvert \cdot \rvert$ | the absolute value: the gap, ignoring direction | $\lvert 0 - 3 \rvert = 3$ |
 896
 897**In words:** "sort both piles, pair them up in order, and average how far
 898each grain has to travel."
 899
 900**With the numbers:** pairs (0, 3), (1, 4), (2, 5), gaps 3, 3, 3, average
 901**3**.
 902
 903**In Python:**
 904
 905```python
 906a = [0, 1, 2]
 907b = [5, 3, 4]
 908# pair the smallest with the smallest, and so on
 909pairs = list(zip(sorted(a), sorted(b)))
 910pairs  # → [(0, 3), (1, 4), (2, 5)]
 911sum(abs(x - y) for x, y in pairs) / len(pairs)  # → 3.0
 912```
 913
 914![As the forger's pile slides from minus 4 to plus 4, the Jensen-Shannon divergence is flat at log 2 everywhere except exactly 0, while the Wasserstein distance is a V shape pointing at 0](figures/primer.ml.generative.gans.wasserstein_vs_js.svg)
 915
 916**Reading it:** the x-axis is where the forger's pile sits; the real pile
 917is at 0. The red line (Jensen-Shannon) is flat at 0.693 everywhere except
 918the single point θ = 0, where it drops to 0: a flat line has no slope, so
 919nothing tells the forger which way to move. The blue V (Wasserstein) slopes
 920down towards 0 from both sides: wherever the forger is, the slope points at
 921the data. On images, where real photos and early fakes barely overlap, this
 922difference is the difference between learning and not.
 923
 924The **Wasserstein GAN** replaces the detective with a **critic** that
 925outputs an unbounded score instead of a probability, and trains it so that
 926the gap between its average score on real and on fake points estimates this
 927distance. The catch: the estimate is only valid if the critic's slope is at
 928most 1 everywhere. The original WGAN enforced that crudely by clipping every
 929weight; WGAN-GP does it with a gradient penalty; spectral normalization,
 930next, does it by capping each layer.
 931
 932**In code:** `wasserstein_1d` is the sort-and-pair formula, and `js_divergence` computes the Jensen-Shannon divergence between two histograms, flat at log 2 whenever they don't overlap.
 933
 934### Spectral normalization: a speed limit on every layer
 935
 936**Everyday picture.** A layer of a network is a matrix, and a matrix
 937stretches some directions more than others. If no layer can stretch
 938anything by more than 1, the whole detective can't change its score faster
 939than the input changes: its cliffs have a speed limit.
 940
 941**Tiny worked example.** The matrix W = [[3, 0], [0, 1]] stretches the
 942horizontal direction by 3 and the vertical by 1. Its **largest stretch** is
 9433. Divide W by 3 and no direction is stretched by more than 1.
 944
 945$$
 946\bar W = \frac{W}{\lVert W \rVert_2}
 947\qquad
 948\lVert W \rVert_2 = \max_{\lVert v \rVert = 1} \lVert W v \rVert
 949$$
 950
 951**Symbols**
 952
 953| Symbol | Meaning here | In the example |
 954|---|---|---|
 955| $W$ | one layer's weight matrix (see `primer.notation` for matrices) | [[3, 0], [0, 1]] |
 956| $v$ | an input direction: an arrow of length 1 | $(1, 0)$ |
 957| $\lVert v \rVert$ | the length of $v$ | 1 |
 958| $Wv$ | the arrow after passing through the layer | $(3, 0)$ |
 959| $\max_{\lVert v \rVert = 1}$ | "the largest value over every arrow of length 1" | |
 960| $\lVert W \rVert_2$ | the **spectral norm**: the most W lengthens any arrow | 3 |
 961| $\bar W$ | the normalised layer used in the detective | [[1, 0], [0, 1/3]] |
 962
 963**In words:** "find the direction the layer stretches most, and divide the
 964whole layer by that stretch."
 965
 966**With the numbers:** $(1, 0)$ becomes $(3, 0)$, length 3; $(0, 1)$ stays
 967length 1; every other direction lands in between. So $\lVert W \rVert_2 = 3$,
 968and $W / 3$ stretches by at most 1.
 969
 970**In Python:**
 971
 972```python
 973import math
 974W = [[3, 0], [0, 1]]
 975def stretch(v):
 976    Wv = [sum(w * x for w, x in zip(row, v)) for row in W]
 977    return math.hypot(*Wv) / math.hypot(*v)
 978# try unit arrows all the way round the circle, one per degree
 979angles = [2 * math.pi * k / 360 for k in range(360)]
 980round(max(stretch((math.cos(t), math.sin(t))) for t in angles), 3)  # → 3.0
 981# divide W by its largest stretch
 982W = [[w / 3 for w in row] for row in W]
 983round(max(stretch((math.cos(t), math.sin(t))) for t in angles), 3)  # → 1.0
 984```
 985
 986Trying every direction is fine for a 2 × 2 matrix but not for a layer with
 987a thousand inputs. **Power iteration** finds the most-stretched direction
 988cheaply:
 989
 990```mermaid
 991flowchart LR
 992  V["start: any arrow v"] --> U["u = W v,<br/>rescaled to length 1"]
 993  U --> B["v = W transpose u,<br/>rescaled to length 1"]
 994  B --> Q{"repeat"}
 995  Q -- again --> U
 996  Q -- done --> S["largest stretch = length of W v"]
 997```
 998
 999**Reading it:** push an arrow through the layer and back through its
1000transpose, over and over. Each round trip multiplies the arrow's
1001most-stretched component by more than any other, so the arrow swings round
1002to point along that direction, and the stretch is read off at the end.
1003Spectral normalization keeps one arrow per layer and runs a single round
1004trip per training step: the weights change slowly, so the arrow stays
1005nearly right, and the cost is a couple of extra matrix-vector products.
1006
1007**Why it matters:** spectral normalization (Miyato and colleagues, 2018) is a
1008one-line change to the detective that made GAN training far less sensitive
1009to learning rates and architecture, and it became a default in large image
1010GANs such as BigGAN.
1011
1012**In code:** `largest_stretch` is power iteration, and `spectrally_normalize` divides a matrix by it.
1013
1014## Where GANs stand today
1015
1016**Everyday picture.** The forger-and-detective idea started as the whole
1017machine. Today it is more often a picky critic *inside* another machine,
1018brought in for the one thing it does best: making outputs look sharp and
1019real.
1020
1021**Tiny worked example.** A GAN makes an image in **one** pass through the
1022generator. A diffusion model typically runs its network tens of times per
1023image, removing a little noise each time (see
1024`primer.ml.generative.diffusion`). For a while, that speed plus very sharp
1025samples (StyleGAN's faces, from 2018, are still striking) made GANs the best
1026image generators. Then diffusion models overtook them on image quality and,
1027above all, on *coverage*: they don't mode-collapse and they train stably,
1028which is everything this lesson has been fighting. Since around 2021 most
1029new image and video generators are diffusion or flow models.
1030
1031GANs didn't disappear. Their loss did the moving:
1032
1033```mermaid
1034flowchart LR
1035  X["real image or audio"] --> E["encoder"]
1036  E --> C["compact code"]
1037  C --> DEC["decoder"]
1038  DEC --> Y["reconstruction"]
1039  Y --> L1["reconstruction loss<br/>pixel or perceptual"]
1040  X --> L1
1041  Y --> DISC["detective:<br/>real or reconstructed?"]
1042  X --> DISC
1043  L1 --> T["total loss for<br/>encoder and decoder"]
1044  DISC --> T
1045```
1046
1047**Reading it:** this is an autoencoder (see
1048`primer.ml.generative.autoencoders`) with a detective bolted on. The top
1049path squeezes the input into a compact code and rebuilds it. A
1050reconstruction loss alone rewards *averages*, and the average of many
1051plausible textures is a blur. The detective on the lower path looks at the
1052reconstruction and the original and asks "which is real?", which punishes
1053blur and rewards crisp detail. The decoder is trained on both losses
1054together. This design is how the image autoencoders inside latent diffusion
1055models are trained, how VQGAN's image tokenizer is trained, and how neural
1056audio decoders such as HiFi-GAN produce clean waveforms. GAN-style losses
1057are also used to distil a slow many-step diffusion model into a fast one-
1058or few-step generator.
1059
1060**Why it matters:** when you meet a modern image, audio or video system,
1061expect a GAN loss somewhere in its decoder or its fast sampling path, even
1062when the headline method is diffusion. Everything in this lesson, the
1063non-saturating loss, gradient penalties, spectral normalization and careful
1064learning rates, is still how those critics are kept stable.
1065
1066## In 20 seconds
1067
1068- **A GAN** trains a generator (forger) to turn noise into samples and a
1069  discriminator (detective) to tell real samples from fakes, against each
1070  other. The forger learns only through the detective's gradient.
1071- **The game:** min over G of max over D of E[log D(x)] + E[log(1 − D(G(z)))].
1072  The best detective says p_data / (p_data + p_g); at the equilibrium the
1073  forger matches the data and the detective says 1/2 everywhere.
1074- **The non-saturating loss** (maximise log D(G(z))) keeps the forger's
1075  gradient strong when the detective is confident, exactly when the
1076  original loss goes quiet.
1077- **Unstable because** each player's target moves: plain gradient steps
1078  circle or spiral around the equilibrium, the forger hops between a few
1079  modes (mode collapse), and a detective that wins too fast freezes it.
1080- **Fixes:** balanced learning rates, gradient penalties (R1, WGAN-GP),
1081  spectral normalization, and the Wasserstein distance, which still points
1082  the way when real and fake don't overlap.
1083- **Today:** diffusion has overtaken GANs for generating images, but
1084  adversarial losses remain inside image and audio decoders and fast
1085  distilled samplers.
1086
1087## Self-test questions
1088
1089**Explain a GAN to a non-engineer in 30 seconds.**
1090Two programs play a game. One makes fake pictures; the other looks at real
1091and fake pictures and guesses which is which. Every time the guesser catches
1092a fake, the faker learns what gave it away and improves. After enough
1093rounds the fakes are good enough that the guesser is reduced to guessing,
1094and the faker has learned to make realistic pictures without anyone ever
1095describing what a picture should look like.
1096
1097**Why does the generator never need to see real data?**
1098It learns only from the discriminator's gradient with respect to its own
1099outputs: which way to move each fake so the discriminator finds it more
1100real. The discriminator has seen real data, so its verdicts carry that
1101information to the generator.
1102
1103**What does the optimal discriminator compute, and what does it say at equilibrium?**
1104D*(x) = p_data(x) / (p_data(x) + p_g(x)): the share of points found at x
1105that are real. When the generator matches the data, p_g = p_data and D* is
11061/2 everywhere, the value is −log 4, and the discriminator can do no better
1107than a coin flip.
1108
1109**Why do almost all GANs use the non-saturating generator loss?**
1110The original loss log(1 − D(G(z))) has a slope of −D with respect to the
1111discriminator's score, which is near zero when the discriminator confidently
1112rejects fakes, as it does early in training. The non-saturating loss
1113−log D(G(z)) has slope −(1 − D), largest exactly then. Both have the same
1114equilibrium.
1115
1116**What is mode collapse, and how would you detect it?**
1117The generator covers only some of the distinct kinds of data (modes), often
1118hopping between them as the discriminator catches up. Per-sample quality can
1119look excellent, so you detect it by measuring coverage: how many modes or
1120classes receive samples, or a distribution-level metric such as FID on real
1121data.
1122
1123**Why doesn't plain gradient descent find the GAN equilibrium?**
1124The game has no shared downhill direction. Near the equilibrium the two
1125players' updates form a rotation, so simultaneous steps spiral outward and
1126alternating steps orbit forever, as the two-number Dirac GAN shows. Damping
1127the discriminator with a gradient penalty turns the spiral inward.
1128
1129**Why is the Wasserstein distance a better training signal than Jensen-Shannon divergence?**
1130When the real and generated distributions don't overlap, Jensen-Shannon is
1131stuck at log 2 however far apart they are, so it gives no direction.
1132Wasserstein measures how far the mass must move, so it shrinks steadily as
1133the generator approaches the data. Estimating it needs a critic whose slope
1134is capped, which is what weight clipping, gradient penalties and spectral
1135normalization provide.
1136
1137**If diffusion models won, why learn GANs?**
1138Adversarial losses are still how many image and audio decoders get sharp
1139output, and how some diffusion models are distilled into one-step
1140generators. And the instabilities here, moving targets and collapsing
1141variety, show up wherever two models are trained against each other.
1142
1143## The papers behind this lesson
1144
1145- **Goodfellow et al., *Generative Adversarial Nets* (2014)**:
1146  https://arxiv.org/abs/1406.2661. Introduced the generator-discriminator
1147  game, derived the optimal discriminator and the equilibrium where the
1148  generator matches the data, and suggested the non-saturating loss.
1149  [Annotated companion](../../../papers/generative-adversarial-nets.html)
1150- **Metz, Poole, Pfau and Sohl-Dickstein, *Unrolled Generative Adversarial
1151  Networks* (2016)**: https://arxiv.org/abs/1611.02163. Used the ring of
1152  eight Gaussians to show a generator hopping between modes, and reduced it
1153  by letting the generator look ahead at several discriminator steps.
1154- **Heusel et al., *GANs Trained by a Two Time-Scale Update Rule Converge to
1155  a Local Nash Equilibrium* (2017)**: https://arxiv.org/abs/1706.08500.
1156  Showed that separate learning rates for the two players give provable
1157  convergence, and introduced the FID score for judging generated images.
1158- **Arjovsky, Chintala and Bottou, *Wasserstein GAN* (2017)**:
1159  https://arxiv.org/abs/1701.07875. Replaced the Jensen-Shannon objective
1160  with the earth mover's distance, which still gives a direction when real
1161  and generated data don't overlap.
1162  [Annotated companion](../../../papers/wasserstein-gan.html)
1163- **Gulrajani et al., *Improved Training of Wasserstein GANs* (2017)**:
1164  https://arxiv.org/abs/1704.00028. Enforced the Wasserstein critic's
1165  slope limit with a gradient penalty instead of weight clipping.
1166- **Mescheder, Geiger and Nowozin, *Which Training Methods for GANs do
1167  actually Converge?* (2018)**: https://arxiv.org/abs/1801.04406.
1168  Introduced the two-number Dirac GAN to show why plain GAN training
1169  circles instead of converging, and the R1 penalty that makes it converge.
1170- **Miyato et al., *Spectral Normalization for Generative Adversarial
1171  Networks* (2018)**: https://arxiv.org/abs/1802.05957. Capped every
1172  discriminator layer's largest stretch at 1 using one step of power
1173  iteration per training step.
1174
1175## Further reading
1176
1177- Goodfellow, *NIPS 2016 Tutorial: Generative Adversarial Networks*: https://arxiv.org/abs/1701.00160
1178- Radford, Metz and Chintala, *Unsupervised Representation Learning with Deep Convolutional GANs* (DCGAN, 2015): https://arxiv.org/abs/1511.06434
1179- Karras, Laine and Aila, *A Style-Based Generator Architecture for Generative Adversarial Networks* (StyleGAN, 2018): https://arxiv.org/abs/1812.04948
1180- Dhariwal and Nichol, *Diffusion Models Beat GANs on Image Synthesis* (2021): https://arxiv.org/abs/2105.05233
1181- Esser, Rombach and Ommer, *Taming Transformers for High-Resolution Image Synthesis* (VQGAN, 2020): https://arxiv.org/abs/2012.09841
1182- Rombach et al., *High-Resolution Image Synthesis with Latent Diffusion Models* (2021): https://arxiv.org/abs/2112.10752
1183- Kong, Kim and Bae, *HiFi-GAN: Generative Adversarial Networks for Efficient and High Fidelity Speech Synthesis* (2020): https://arxiv.org/abs/2010.05646
1184- Sauer et al., *Adversarial Diffusion Distillation* (2023): https://arxiv.org/abs/2311.17042
1185- PyTorch, *DCGAN Tutorial*: https://pytorch.org/tutorials/beginner/dcgan_faces_tutorial.html
1186"""
1187
1188from __future__ import annotations
1189
1190import functools
1191import math
1192
1193import numpy as np
1194
1195from primer._show import banner, say, table, takeaway
1196from primer.ml.optimizers import Adam
1197
1198# ---------------------------------------------------------------------------
1199# 1. The toy data: a ring of eight little clouds
1200# ---------------------------------------------------------------------------
1201
1202N_MODES, RADIUS, MODE_STD = 8, 2.0, 0.05
1203
1204
1205def ring_modes(n_modes: int = N_MODES, radius: float = RADIUS) -> np.ndarray:
1206    """The centres of the clouds: (n_modes, 2) points evenly spaced on a circle, the first on the x-axis."""
1207    angles = 2 * np.pi * np.arange(n_modes) / n_modes
1208    return radius * np.stack([np.cos(angles), np.sin(angles)], axis=1)
1209
1210
1211def ring_of_gaussians(n: int, seed: int | np.random.Generator = 0, std: float = MODE_STD) -> np.ndarray:
1212    """n real samples: pick one of the eight modes at random, then add a little round noise. Shape (n, 2)."""
1213    rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
1214    centres = ring_modes()[rng.integers(0, N_MODES, n)]
1215    return centres + std * rng.standard_normal((n, 2))
1216
1217
1218# ---------------------------------------------------------------------------
1219# 2. The two players: small multi-layer perceptrons, forward and backward by hand
1220# ---------------------------------------------------------------------------
1221
1222GENERATOR_SIZES = (2, 32, 32, 2)  # 2 numbers of noise in, two tanh layers, one 2-D point out
1223DISCRIMINATOR_SIZES = (2, 64, 1)  # a 2-D point in, one tanh layer, one score (a logit) out
1224
1225
1226def sigmoid(a):
1227    """1 / (1 + e^-a), written with tanh so it never overflows for large |a|."""
1228    return 0.5 * (1.0 + np.tanh(0.5 * np.asarray(a, dtype=float)))
1229
1230
1231class MLP:
1232    """A stack of layers `x W + b`, with tanh between them and nothing after the last.
1233
1234    Every weight and bias lives in one flat vector, `params`; `W` and `b` are
1235    views into it. That lets one optimizer step update the whole network in
1236    one line, and lets a test nudge any single number. The same recipe as
1237    `primer.ml.neural_net.MLP`, generalised to any depth and to a backward
1238    pass that also returns the gradient with respect to the *input*, which is
1239    the only route by which the forger ever learns anything.
1240    """
1241
1242    def __init__(self, sizes: tuple[int, ...], seed: int = 0):
1243        rng = np.random.default_rng(seed)
1244        shapes = [(a, b) for a, b in zip(sizes[:-1], sizes[1:])]
1245        self.params = np.zeros(sum(a * b + b for a, b in shapes))
1246        self.W, self.b = self._views(self.params, shapes)
1247        self._shapes = shapes
1248        for W in self.W:
1249            # Variance 1/fan_in keeps tanh out of its flat ends at the start (see primer.ml.deep_nets).
1250            W[...] = rng.normal(0.0, 1.0 / np.sqrt(W.shape[0]), W.shape)
1251
1252    @staticmethod
1253    def _views(flat: np.ndarray, shapes) -> tuple[list[np.ndarray], list[np.ndarray]]:
1254        Ws, bs, i = [], [], 0
1255        for a, b in shapes:
1256            Ws.append(flat[i : i + a * b].reshape(a, b))
1257            i += a * b
1258            bs.append(flat[i : i + b])
1259            i += b
1260        return Ws, bs
1261
1262    def forward(self, x: np.ndarray) -> np.ndarray:
1263        """x: (n, sizes[0]) -> (n, sizes[-1]). Caches every layer's output for `backward`."""
1264        self._acts = [x]
1265        for i, (W, b) in enumerate(zip(self.W, self.b)):
1266            z = self._acts[-1] @ W + b
1267            self._acts.append(np.tanh(z) if i < len(self.W) - 1 else z)
1268        return self._acts[-1]
1269
1270    def backward(self, grad_out: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
1271        """Chain rule from d(something)/d(output) back to (every parameter, the input).
1272
1273        Returns (flat gradient laid out like `params`, gradient w.r.t. the input x).
1274        """
1275        grad = np.zeros_like(self.params)
1276        gW, gb = self._views(grad, self._shapes)
1277        g = grad_out
1278        for i in reversed(range(len(self.W))):
1279            gW[i][...] = self._acts[i].T @ g  # each weight's gradient: its input times the error arriving
1280            gb[i][...] = g.sum(axis=0)
1281            g = g @ self.W[i].T  # send the error back through the weights
1282            if i > 0:
1283                g = g * (1.0 - self._acts[i] ** 2)  # tanh'(z) = 1 - tanh(z)^2, reusing the cached output
1284        return grad, g
1285
1286
1287# ---------------------------------------------------------------------------
1288# 3. The game: the value function and the best possible detective
1289# ---------------------------------------------------------------------------
1290
1291
1292def value_function(d_real, d_fake, eps: float = 1e-12) -> float:
1293    """V = mean log D(x) over real samples + mean log(1 - D(G(z))) over fakes.
1294
1295    The detective wants V high (ceiling 0); the forger wants it low. -V is
1296    exactly the detective's binary cross-entropy with real = 1 and fake = 0.
1297    """
1298    d_real = np.clip(np.asarray(d_real, dtype=float), eps, 1.0)
1299    d_fake = np.clip(np.asarray(d_fake, dtype=float), 0.0, 1.0 - eps)
1300    return float(np.mean(np.log(d_real)) + np.mean(np.log(1.0 - d_fake)))
1301
1302
1303def optimal_discriminator(p_data, p_g):
1304    """D*(x) = p_data(x) / (p_data(x) + p_g(x)): the best any detective can do against a fixed forger."""
1305    p_data, p_g = np.asarray(p_data, dtype=float), np.asarray(p_g, dtype=float)
1306    return p_data / (p_data + p_g)
1307
1308
1309def fit_discriminator_table(p_data, p_g, steps: int = 4000, lr: float = 5.0) -> np.ndarray:
1310    """Train the most flexible detective there is, one free score per point, and return its verdicts.
1311
1312    The data and the forger live on the same few points with the given
1313    probabilities. Gradient ascent on the exact expected V: for a point with
1314    score a and verdict D = sigmoid(a), dV/da = p_data (1 - D) - p_g D. It is
1315    zero exactly when D = p_data / (p_data + p_g), so the trained table
1316    arrives at `optimal_discriminator` without ever being told the formula.
1317    """
1318    p_data, p_g = np.asarray(p_data, dtype=float), np.asarray(p_g, dtype=float)
1319    a = np.zeros_like(p_data)
1320    for _ in range(steps):
1321        d = sigmoid(a)
1322        a += lr * (p_data * (1.0 - d) - p_g * d)
1323    return sigmoid(a)
1324
1325
1326# ---------------------------------------------------------------------------
1327# 4. Saturation: the original forger loss versus the non-saturating one
1328# ---------------------------------------------------------------------------
1329
1330
1331def generator_logit_gradient(d_fake, loss: str = "non_saturating"):
1332    """The forger's loss gradient with respect to the detective's score a, where D(G(z)) = sigmoid(a).
1333
1334    saturating (the original):  loss = log(1 - D)  ->  d loss / da = -D
1335    non_saturating:             loss = -log D      ->  d loss / da = -(1 - D)
1336
1337    When the detective is sure a fake is fake (D near 0), the first is near 0
1338    and the second near -1: only the second still tells the forger which way to move.
1339    """
1340    d = np.asarray(d_fake, dtype=float)
1341    out = -d if loss == "saturating" else -(1.0 - d)
1342    return float(out) if out.ndim == 0 else out
1343
1344
1345def _forger_gradient(G: MLP, D: MLP, z: np.ndarray, loss: str) -> np.ndarray:
1346    """The gradient of the forger's loss with respect to every forger parameter (one flat vector)."""
1347    fakes = G.forward(z)
1348    a = D.forward(fakes)
1349    _, grad_fakes = D.backward(generator_logit_gradient(sigmoid(a), loss) / len(z))
1350    grad_params, _ = G.backward(grad_fakes)
1351    return grad_params
1352
1353
1354def head_start_experiment(head_starts=(0, 25, 50, 100, 200, 300), lr: float = 3e-3, batch: int = 128, seed: int = 0) -> list[dict]:
1355    """Train only the detective against a frozen, untrained forger, then measure what the forger would learn.
1356
1357    For each head start (detective steps), returns the average verdict on
1358    fakes and the size (norm) of the forger's parameter gradient under each
1359    loss. The untrained forger's points sit near the centre, far from the
1360    ring, so the detective quickly becomes certain, and the original loss's
1361    gradient fades while the non-saturating one stays large.
1362    """
1363    rng = np.random.default_rng(seed)
1364    G, D = MLP(GENERATOR_SIZES, seed=seed), MLP(DISCRIMINATOR_SIZES, seed=seed + 1)
1365    opt = Adam(lr=lr, beta1=0.5)
1366    probe_z = np.random.default_rng(seed + 100).standard_normal((batch, 2))
1367    rows, done = [], 0
1368    for target in sorted(head_starts):
1369        while done < target:
1370            _detective_step(D, opt, ring_of_gaussians(batch, rng), G.forward(rng.standard_normal((batch, 2))))
1371            done += 1
1372        rows.append(
1373            dict(
1374                head_start=target,
1375                d_fake=float(sigmoid(D.forward(G.forward(probe_z))).mean()),
1376                saturating=float(np.linalg.norm(_forger_gradient(G, D, probe_z, "saturating"))),
1377                non_saturating=float(np.linalg.norm(_forger_gradient(G, D, probe_z, "non_saturating"))),
1378            )
1379        )
1380    return rows
1381
1382
1383# ---------------------------------------------------------------------------
1384# 5. Training: alternate one detective step and one forger step
1385# ---------------------------------------------------------------------------
1386
1387
1388def _detective_step(D: MLP, opt: Adam, real: np.ndarray, fakes: np.ndarray) -> None:
1389    """One step up the value function: push D(real) towards 1 and D(fake) towards 0."""
1390    n = len(real)
1391    # d(-V)/da is D - 1 for a real point and D for a fake: binary cross-entropy's "prediction minus target".
1392    grad_real, _ = D.backward((sigmoid(D.forward(real)) - 1.0) / n)
1393    grad_fake, _ = D.backward(sigmoid(D.forward(fakes)) / n)
1394    D.params[:] = opt.step(D.params, grad_real + grad_fake)
1395
1396
1397def mode_coverage(samples: np.ndarray, min_share: float = 0.02, std: float = MODE_STD) -> dict:
1398    """How much of the ring the samples cover.
1399
1400    A sample is *high quality* if it lands within 3 standard deviations of a
1401    mode's centre. A mode is *covered* if it holds at least `min_share` of
1402    all samples (2% by default: 20 of 1,000, where a perfect forger puts 125).
1403    Returns covered (a count), which (a string such as ".X...X.." marking the
1404    covered modes), counts per mode, and quality (the high-quality fraction).
1405    """
1406    distances = np.linalg.norm(samples[:, None, :] - ring_modes()[None], axis=2)
1407    nearest, close = distances.argmin(axis=1), distances.min(axis=1) < 3 * std
1408    counts = np.bincount(nearest[close], minlength=N_MODES)
1409    hit = counts >= min_share * len(samples)
1410    return dict(
1411        covered=int(hit.sum()),
1412        which="".join("X" if h else "." for h in hit),
1413        counts=counts,
1414        quality=float(close.mean()),
1415    )
1416
1417
1418@functools.lru_cache(maxsize=None)
1419def train_gan(
1420    steps: int = 2500,
1421    lr_g: float = 1e-3,
1422    lr_d: float = 3e-3,
1423    batch: int = 128,
1424    snapshot_every: int = 250,
1425    loss: str = "non_saturating",
1426    seed: int = 1,
1427) -> dict:
1428    """Train a forger and a detective on the ring, alternating one step each.
1429
1430    Returns snapshots every `snapshot_every` steps: the step, 1,000 generated
1431    points (always from the same noise, so pictures are comparable), the
1432    detective's parameters, and the `mode_coverage` of the points. Cached,
1433    because the lesson, its figures and its tests all look at the same runs.
1434    """
1435    rng = np.random.default_rng(seed)
1436    G, D = MLP(GENERATOR_SIZES, seed=seed), MLP(DISCRIMINATOR_SIZES, seed=seed + 1)
1437    # beta1 = 0.5 rather than 0.9: less momentum, the usual choice for GANs since DCGAN, because the target keeps moving.
1438    opt_g, opt_d = Adam(lr=lr_g, beta1=0.5), Adam(lr=lr_d, beta1=0.5)
1439    probe_z = np.random.default_rng(seed + 100).standard_normal((1000, 2))
1440    snapshots = []
1441    for step in range(steps + 1):
1442        if step % snapshot_every == 0:
1443            points = G.forward(probe_z).copy()
1444            snapshots.append(dict(step=step, samples=points, detective=D.params.copy(), coverage=mode_coverage(points)))
1445        if step == steps:
1446            break
1447        z = rng.standard_normal((batch, 2))
1448        fakes = G.forward(z)
1449        # 1. Detective: one step towards telling this batch of real points from this batch of fakes.
1450        _detective_step(D, opt_d, ring_of_gaussians(batch, rng), fakes)
1451        # 2. Forger: one step towards fooling the detective as it is *now*.
1452        G.params[:] = opt_g.step(G.params, _forger_gradient(G, D, z, loss))
1453    return dict(snapshots=snapshots, coverage=[s["coverage"] for s in snapshots], lr_g=lr_g, lr_d=lr_d)
1454
1455
1456def detective_verdicts(params: np.ndarray, points: np.ndarray) -> np.ndarray:
1457    """D(x) for each point, using a detective's saved parameters."""
1458    D = MLP(DISCRIMINATOR_SIZES)
1459    D.params[:] = params
1460    return sigmoid(D.forward(points))[:, 0]
1461
1462
1463# ---------------------------------------------------------------------------
1464# 6. Oscillation in miniature: the Dirac GAN
1465# ---------------------------------------------------------------------------
1466
1467
1468def dirac_gan(theta: float = 1.0, psi: float = 0.0, lr: float = 0.2, steps: int = 300, gamma: float = 0.0, alternating: bool = False) -> np.ndarray:
1469    """The smallest GAN there is: two numbers, followed step by step. Returns the path, shape (steps + 1, 2).
1470
1471    Real data is always the single point x = 0. The forger has one number,
1472    theta, and always outputs x = theta. The detective has one number, psi,
1473    and says D(x) = sigmoid(psi x). The value is
1474    V = log D(0) + log(1 - D(theta)) = log 0.5 + log(1 - sigmoid(psi theta)).
1475
1476        dV/dpsi   = -sigmoid(psi theta) theta   (the detective climbs this)
1477        dV/dtheta = -sigmoid(psi theta) psi     (the forger descends this)
1478
1479    `gamma` > 0 adds the R1 gradient penalty (gamma / 2)(dD-score/dx at the
1480    real data)^2 = (gamma / 2) psi^2 to what the detective minimises.
1481    `alternating` lets the forger see the detective's new psi before it moves.
1482    """
1483    path = [(theta, psi)]
1484    for _ in range(steps):
1485        s = sigmoid(psi * theta)
1486        new_psi = psi + lr * (-s * theta - gamma * psi)
1487        if alternating:
1488            psi = new_psi
1489            s = sigmoid(psi * theta)
1490        theta = theta + lr * s * psi  # descend dV/dtheta = -s psi
1491        psi = new_psi
1492        path.append((float(theta), float(psi)))
1493    return np.array(path)
1494
1495
1496# ---------------------------------------------------------------------------
1497# 7. Stabilizers: distances between piles, and capping the detective's steepness
1498# ---------------------------------------------------------------------------
1499
1500
1501def js_divergence(p, q) -> float:
1502    """Jensen-Shannon divergence between two histograms on the same bins (natural log).
1503
1504    JS = 1/2 KL(p || m) + 1/2 KL(q || m), with m the average of p and q. It
1505    is 0 when p = q and log 2 when they share no bin at all, *however far
1506    apart* the bins are. That flatness is why a GAN's detective can stop
1507    giving useful directions.
1508    """
1509    p, q = np.asarray(p, dtype=float), np.asarray(q, dtype=float)
1510    p, q = p / p.sum(), q / q.sum()
1511    m = 0.5 * (p + q)
1512
1513    def kl(a, b):
1514        mask = a > 0  # 0 log 0 counts as 0
1515        return float(np.sum(a[mask] * np.log(a[mask] / b[mask])))
1516
1517    return 0.5 * kl(p, m) + 0.5 * kl(q, m)
1518
1519
1520def wasserstein_1d(a, b) -> float:
1521    """Earth mover's distance between two equal-sized 1-D samples: sort both, pair them up, average the gaps."""
1522    a, b = np.sort(np.asarray(a, dtype=float)), np.sort(np.asarray(b, dtype=float))
1523    return float(np.mean(np.abs(a - b)))
1524
1525
1526def largest_stretch(W: np.ndarray, iters: int = 100, seed: int = 0) -> float:
1527    """The most W can lengthen any vector (its spectral norm), by power iteration.
1528
1529    Push a vector through W and back through W transpose, again and again;
1530    it swings round to the direction W stretches most, and the stretch
1531    factor is read off at the end. Spectral normalization runs one step of
1532    this per training step instead of a full singular value decomposition.
1533    """
1534    v = np.random.default_rng(seed).standard_normal(W.shape[1])
1535    for _ in range(iters):
1536        u = W @ v
1537        u /= np.linalg.norm(u)
1538        v = W.T @ u
1539        v /= np.linalg.norm(v)
1540    return float(np.linalg.norm(W @ v))
1541
1542
1543def spectrally_normalize(W: np.ndarray) -> np.ndarray:
1544    """W divided by its largest stretch, so no input direction is stretched by more than 1."""
1545    return W / largest_stretch(W)
1546
1547
1548# ---------------------------------------------------------------------------
1549# 8. Figures (rendered into the HTML docs by `make figures`)
1550# ---------------------------------------------------------------------------
1551
1552# The detective's learning rate for each of the three ring runs; the forger always uses 1e-3.
1553RING_RUNS = {"equal speeds (0.001)": 1e-3, "detective 3x faster (0.003)": 3e-3, "detective 10x faster (0.01)": 1e-2}
1554
1555
1556def _normal_pdf(x, mean, std):
1557    return np.exp(-0.5 * ((x - mean) / std) ** 2) / (std * math.sqrt(2 * math.pi))
1558
1559
1560def figures() -> dict:
1561    """Plot this lesson's data. matplotlib is imported here, and only here,
1562    so the lesson itself needs nothing beyond NumPy."""
1563    import matplotlib
1564
1565    matplotlib.use("Agg")
1566    import matplotlib.pyplot as plt
1567
1568    REAL, FAKE, GOOD, MUTED, THIRD = "#2563eb", "#dc2626", "#059669", "#9ca3af", "#d97706"
1569    run_colours = dict(zip(RING_RUNS, (FAKE, REAL, THIRD)))
1570    runs = {name: train_gan(lr_d=lr_d) for name, lr_d in RING_RUNS.items()}
1571    figs = {}
1572
1573    # --- 1. The best detective on a line -------------------------------------
1574    xs = np.linspace(-5, 5, 400)
1575    p_data = 0.5 * _normal_pdf(xs, -2, 0.5) + 0.5 * _normal_pdf(xs, 2, 0.5)
1576    p_g = _normal_pdf(xs, 0, 1.5)
1577    fig, (top, bottom) = plt.subplots(2, 1, figsize=(6.4, 5), sharex=True)
1578    top.fill_between(xs, p_data, color=REAL, alpha=0.25)
1579    top.plot(xs, p_data, color=REAL, label="real data  p_data")
1580    top.fill_between(xs, p_g, color=FAKE, alpha=0.15)
1581    top.plot(xs, p_g, color=FAKE, label="forger  p_g")
1582    top.set_ylabel("density")
1583    top.set_title("Where each distribution puts its points")
1584    top.legend(frameon=False, loc="upper center")
1585    bottom.plot(xs, optimal_discriminator(p_data, p_g), color="#111827", label="best verdict  D*")
1586    bottom.axhline(0.5, color=MUTED, ls="--", label="D* once the forger matches the data")
1587    bottom.set_ylim(-0.03, 1.03)
1588    bottom.set_xlabel("position x")
1589    bottom.set_ylabel("D*(x)")
1590    bottom.set_title("The best detective: the share of points at x that are real")
1591    bottom.legend(frameon=False, loc="lower right", fontsize=8)
1592    fig.tight_layout()
1593    figs["optimal_detective"] = fig
1594
1595    # --- 2. The trained detective's view of the ring -------------------------
1596    last = runs["detective 3x faster (0.003)"]["snapshots"][-1]
1597    grid = np.linspace(-3, 3, 161)
1598    gx, gy = np.meshgrid(grid, grid)
1599    verdict = detective_verdicts(last["detective"], np.stack([gx.ravel(), gy.ravel()], axis=1)).reshape(gx.shape)
1600    real = ring_of_gaussians(1000, seed=5)
1601    fig, ax = plt.subplots(figsize=(5.6, 4.8))
1602    im = ax.imshow(verdict, origin="lower", extent=(-3, 3, -3, 3), cmap="RdBu", vmin=0, vmax=1)
1603    ax.scatter(real[:, 0], real[:, 1], s=3, color="#111827", label="real")
1604    ax.scatter(last["samples"][:, 0], last["samples"][:, 1], s=3, color=GOOD, alpha=0.6, label="forger")
1605    ax.set_aspect("equal")
1606    ax.grid(False)
1607    ax.set_title(f"The detective's verdict after {last['step']:,} steps")
1608    ax.legend(frameon=True, loc="upper right", markerscale=3)
1609    fig.colorbar(im, ax=ax, fraction=0.046, label="D(x): probability real")
1610    figs["detective_view"] = fig
1611
1612    # --- 3. Saturating vs non-saturating forger loss -------------------------
1613    scores = np.linspace(-7, 3, 300)
1614    rows = head_start_experiment(head_starts=(0, 50, 100, 200, 300, 400, 600, 800))
1615    fig, (a1, a2) = plt.subplots(1, 2, figsize=(10, 3.8))
1616    a1.plot(scores, np.log(1 - sigmoid(scores)), color=FAKE, label="original: log(1 - D)")
1617    a1.plot(scores, -np.log(sigmoid(scores)), color=REAL, label="non-saturating: -log D")
1618    a1.axvline(math.log(0.01 / 0.99), color=MUTED, ls="--")
1619    a1.text(math.log(0.01 / 0.99) + 0.15, 5.6, "D = 0.01", color="#4b5563")
1620    a1.set_xlabel("detective's score a for a fake  (left: sure it is fake)")
1621    a1.set_ylabel("forger's loss")
1622    a1.set_title("The original loss goes flat when the detective is sure")
1623    a1.legend(frameon=False, loc="upper right")
1624    steps = [r["head_start"] for r in rows]
1625    a2.semilogy(steps, [r["saturating"] for r in rows], "o-", color=FAKE, label="original loss")
1626    a2.semilogy(steps, [r["non_saturating"] for r in rows], "o-", color=REAL, label="non-saturating loss")
1627    for r in rows[::2]:
1628        a2.annotate(f"D={r['d_fake']:.3f}", (r["head_start"], r["non_saturating"]), textcoords="offset points", xytext=(-10, 8), fontsize=7, color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1629    a2.set_xlabel("detective's head start (training steps)")
1630    a2.set_ylabel("size of the forger's gradient")
1631    a2.set_title("Measured on the ring: a head start starves one loss")
1632    a2.legend(frameon=False, loc="center right")
1633    fig.tight_layout()
1634    figs["forger_losses"] = fig
1635
1636    # --- 4. The Dirac GAN: spiral out, or orbit ------------------------------
1637    simultaneous, alternating, penalised = dirac_gan(), dirac_gan(alternating=True), dirac_gan(gamma=1.0)
1638    fig, ax = plt.subplots(figsize=(5.4, 5))
1639    ax.plot(simultaneous[:, 0], simultaneous[:, 1], color=FAKE, lw=1.2, label="simultaneous steps: spiral out")
1640    ax.plot(alternating[:, 0], alternating[:, 1], color=REAL, lw=1.2, label="alternating steps: endless orbit")
1641    ax.plot([1], [0], "o", color="#111827")
1642    ax.annotate("start (1, 0)", (1, 0), textcoords="offset points", xytext=(6, 6), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1643    ax.plot([0], [0], "x", color="#111827", ms=10, mew=2)
1644    ax.annotate("equilibrium", (0, 0), textcoords="offset points", xytext=(6, -12), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1645    ax.set_aspect("equal")
1646    ax.set_xlabel("forger's theta (where the fake sits)")
1647    ax.set_ylabel("detective's psi (its slope)")
1648    ax.set_title("The Dirac GAN never settles")
1649    ax.legend(frameon=False, loc="upper left", fontsize=8)
1650    figs["dirac_oscillation"] = fig
1651
1652    # --- 5. Ring snapshots: hopping, covering, stuck -------------------------
1653    shown = (1750, 2000, 2250, 2500)
1654    modes = ring_modes()
1655    fig, axes = plt.subplots(3, 4, figsize=(10, 7.8), sharex=True, sharey=True)
1656    for row, (name, run) in zip(axes, runs.items()):
1657        by_step = {snap["step"]: snap for snap in run["snapshots"]}
1658        for ax, step in zip(row, shown):
1659            snap = by_step[step]
1660            for centre in modes:
1661                ax.add_patch(plt.Circle(centre, 3 * MODE_STD, color=MUTED, alpha=0.5))
1662            ax.scatter(snap["samples"][:, 0], snap["samples"][:, 1], s=2, color=GOOD, alpha=0.5)
1663            ax.set_title(f"step {step:,}  {snap['coverage']['which']}", fontsize=9, family="monospace")
1664            ax.set_xlim(-3, 3)
1665            ax.set_ylim(-3, 3)
1666            ax.set_aspect("equal")
1667            ax.set_xticks([])
1668            ax.set_yticks([])
1669        row[0].set_ylabel(name, fontsize=9)
1670    fig.suptitle("The forger's points (green) against the eight real clouds (grey)", fontweight="bold")
1671    fig.tight_layout()
1672    figs["ring_snapshots"] = fig
1673
1674    # --- 6. Mode coverage over training ----------------------------------------
1675    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1676    for name, run in runs.items():
1677        ax.step([s["step"] for s in run["snapshots"]], [c["covered"] for c in run["coverage"]], where="post", color=run_colours[name], lw=2, label=name)
1678    ax.set_ylim(-0.3, 8.5)
1679    ax.set_yticks(range(0, 9))
1680    ax.set_xlabel("training step")
1681    ax.set_ylabel("clouds covered (of 8)")
1682    ax.set_title("How much of the ring the forger covers")
1683    ax.legend(frameon=False, loc="upper left", fontsize=8)
1684    figs["mode_coverage"] = fig
1685
1686    # --- 7. R1 penalty: the spiral turns inward --------------------------------
1687    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1688    for path, colour, label in (
1689        (simultaneous, FAKE, "simultaneous steps"),
1690        (alternating, REAL, "alternating steps"),
1691        (penalised, GOOD, "simultaneous + R1 penalty (gamma = 1)"),
1692    ):
1693        ax.plot(np.linalg.norm(path, axis=1), color=colour, lw=2, label=label)
1694    ax.set_xlabel("step")
1695    ax.set_ylabel("distance from equilibrium")
1696    ax.set_title("A gradient penalty damps the chase")
1697    ax.legend(frameon=False, loc="upper left", fontsize=8)
1698    figs["dirac_r1"] = fig
1699
1700    # --- 8. Jensen-Shannon vs Wasserstein --------------------------------------
1701    bins = np.round(np.arange(-40, 41) / 10, 1)  # positions -4.0 ... 4.0 in steps of 0.1
1702    data = (bins == 0.0).astype(float)
1703    thetas = bins
1704    js = [js_divergence(data, (bins == t).astype(float)) for t in thetas]
1705    w = [wasserstein_1d([0.0], [t]) for t in thetas]
1706    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1707    ax.plot(thetas, js, ".", color=FAKE, ms=4, label="Jensen-Shannon divergence")
1708    ax.plot(thetas, w, color=REAL, lw=2, label="Wasserstein distance")
1709    ax.axhline(math.log(2), color=MUTED, ls="--")
1710    ax.text(-3.9, math.log(2) + 0.12, "log 2 = 0.693", color="#4b5563")
1711    ax.set_xlabel("where the forger's pile sits (the real pile is at 0)")
1712    ax.set_ylabel("distance between the piles")
1713    ax.set_title("Only one yardstick says which way to move")
1714    ax.legend(frameon=False, loc="upper center")
1715    figs["wasserstein_vs_js"] = fig
1716
1717    return figs
1718
1719
1720# ---------------------------------------------------------------------------
1721# 9. Narrated walkthrough
1722# ---------------------------------------------------------------------------
1723
1724
1725def demo() -> None:
1726    banner("1. The two players")
1727    G, D = MLP(GENERATOR_SIZES, seed=1), MLP(DISCRIMINATOR_SIZES, seed=2)
1728    fakes = G.forward(np.random.default_rng(0).standard_normal((1000, 2)))
1729    real = ring_of_gaussians(1000, seed=0)
1730    say(
1731        f"""
1732        The real data is a ring of eight little clouds of radius {RADIUS}. The
1733        forger maps 2 random numbers to a point through layers {GENERATOR_SIZES};
1734        the detective maps a point to a probability through {DISCRIMINATOR_SIZES}
1735        and a sigmoid. Untrained, the forger's points sit an average of
1736        {np.linalg.norm(fakes, axis=1).mean():.2f} from the centre, while real
1737        points sit {np.linalg.norm(real, axis=1).mean():.2f} away, and the
1738        untrained detective says {sigmoid(D.forward(real)).mean():.2f} on real
1739        points and {sigmoid(D.forward(fakes)).mean():.2f} on fakes: it knows nothing yet.
1740        """
1741    )
1742
1743    banner("2. The game: one number both players fight over")
1744    table(
1745        ["detective", "D on real notes", "D on fakes", "V"],
1746        [
1747            ("the worked example", "0.9, 0.8", "0.2, 0.4", value_function([0.9, 0.8], [0.2, 0.4])),
1748            ("never wrong", "1, 1", "0, 0", value_function([1.0, 1.0], [0.0, 0.0])),
1749            ("coin flip", "0.5, 0.5", "0.5, 0.5", value_function([0.5, 0.5], [0.5, 0.5])),
1750        ],
1751        floatfmt=".3f",
1752    )
1753    takeaway("The detective pushes V up towards 0; the forger pushes it down. -V is the detective's binary cross-entropy.")
1754
1755    banner("3. The best possible detective")
1756    p_data, p_g = [0.4, 0.3, 0.2, 0.1], [0.1, 0.1, 0.4, 0.4]
1757    trained = fit_discriminator_table(p_data, p_g)
1758    table(
1759        ["point", "p_data", "p_g", "formula D*", "trained detective"],
1760        [(i + 1, pd, pg, optimal_discriminator(pd, pg), t) for i, (pd, pg, t) in enumerate(zip(p_data, p_g, trained))],
1761        floatfmt=".3f",
1762    )
1763    say("A detective trained against a fixed forger lands on p_data / (p_data + p_g) without being told the formula.")
1764    takeaway("At equilibrium the forger matches the data, so the best detective says 1/2 everywhere and V = -log 4.")
1765
1766    banner("4. Why the forger uses the non-saturating loss")
1767    table(
1768        ["D(G(z))", "original loss slope", "non-saturating slope"],
1769        [(d, generator_logit_gradient(d, "saturating"), generator_logit_gradient(d, "non_saturating")) for d in (0.5, 0.1, 0.01, 0.001)],
1770        floatfmt=".3f",
1771    )
1772    table(
1773        ["detective head start", "D on fakes", "forger gradient, original", "forger gradient, non-saturating"],
1774        [(r["head_start"], r["d_fake"], r["saturating"], r["non_saturating"]) for r in head_start_experiment(head_starts=(0, 200, 800))],
1775        floatfmt=".3f",
1776    )
1777    takeaway("When the detective is sure, the original loss goes quiet; the non-saturating one is loudest exactly then.")
1778
1779    banner("5. Oscillation: the two-number Dirac GAN")
1780    table(
1781        ["steps of size 0.2", "distance after 50", "after 100", "after 300"],
1782        [
1783            (name, *(float(np.linalg.norm(path[i])) for i in (50, 100, 300)))
1784            for name, path in (
1785                ("simultaneous", dirac_gan()),
1786                ("alternating", dirac_gan(alternating=True)),
1787                ("simultaneous + R1 penalty", dirac_gan(gamma=1.0)),
1788            )
1789        ],
1790        floatfmt=".3f",
1791    )
1792    say("Start 1 away from the equilibrium. Plain steps spiral out or orbit; the R1 penalty's friction brings them home.")
1793
1794    banner("6. Mode collapse on the ring")
1795    for name, lr_d in RING_RUNS.items():
1796        run = train_gan(lr_d=lr_d)
1797        print(f"{name:28s} " + " ".join(c["which"] for c in run["coverage"][5:]))
1798    print()
1799    table(
1800        ["detective learning rate", "clouds covered at the end", "high-quality points"],
1801        [(name, train_gan(lr_d=lr_d)["coverage"][-1]["covered"], train_gan(lr_d=lr_d)["coverage"][-1]["quality"]) for name, lr_d in RING_RUNS.items()],
1802        floatfmt=".2f",
1803    )
1804    say(
1805        """
1806        Each row shows coverage every 250 steps from step 1,250 (X = covered).
1807        Equal speeds: the forger hops between alternate clouds. A detective
1808        three times faster: all eight. Ten times faster: six, and stuck.
1809        """
1810    )
1811    takeaway("Mode collapse is invisible sample by sample; you have to measure coverage.")
1812
1813    banner("7. Yardsticks and speed limits")
1814    grid = np.arange(11)
1815    table(
1816        ["forger's pile at", "Jensen-Shannon", "Wasserstein"],
1817        [(t, js_divergence(grid == 0, grid == t), wasserstein_1d([0.0], [float(t)])) for t in (1, 5, 10)],
1818        floatfmt=".3f",
1819    )
1820    W = np.array([[3.0, 0.0], [0.0, 1.0]])
1821    say(
1822        f"""
1823        Jensen-Shannon can't tell a near miss from a far one; Wasserstein can.
1824        Spectral normalization: W = [[3, 0], [0, 1]] has largest stretch
1825        {largest_stretch(W):.3f}; after dividing by it, {largest_stretch(spectrally_normalize(W)):.3f}.
1826        """
1827    )
1828    takeaway("Keep the detective useful: paced, smooth and never infinitely sure.")
1829
1830
1831if __name__ == "__main__":
1832    demo()
Level 3: the code, function by function.
def ring_modes(n_modes: int = 8, radius: float = 2.0) -> numpy.ndarray: on GitHub
1206def ring_modes(n_modes: int = N_MODES, radius: float = RADIUS) -> np.ndarray:
1207    """The centres of the clouds: (n_modes, 2) points evenly spaced on a circle, the first on the x-axis."""
1208    angles = 2 * np.pi * np.arange(n_modes) / n_modes
1209    return radius * np.stack([np.cos(angles), np.sin(angles)], axis=1)

The centres of the clouds: (n_modes, 2) points evenly spaced on a circle, the first on the x-axis.

def ring_of_gaussians( n: int, seed: int | numpy.random._generator.Generator = 0, std: float = 0.05) -> numpy.ndarray: on GitHub
1212def ring_of_gaussians(n: int, seed: int | np.random.Generator = 0, std: float = MODE_STD) -> np.ndarray:
1213    """n real samples: pick one of the eight modes at random, then add a little round noise. Shape (n, 2)."""
1214    rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
1215    centres = ring_modes()[rng.integers(0, N_MODES, n)]
1216    return centres + std * rng.standard_normal((n, 2))

n real samples: pick one of the eight modes at random, then add a little round noise. Shape (n, 2).

GENERATOR_SIZES = (2, 32, 32, 2)
DISCRIMINATOR_SIZES = (2, 64, 1)
def sigmoid(a): on GitHub
1227def sigmoid(a):
1228    """1 / (1 + e^-a), written with tanh so it never overflows for large |a|."""
1229    return 0.5 * (1.0 + np.tanh(0.5 * np.asarray(a, dtype=float)))

1 / (1 + e^-a), written with tanh so it never overflows for large |a|.

class MLP: on GitHub
1232class MLP:
1233    """A stack of layers `x W + b`, with tanh between them and nothing after the last.
1234
1235    Every weight and bias lives in one flat vector, `params`; `W` and `b` are
1236    views into it. That lets one optimizer step update the whole network in
1237    one line, and lets a test nudge any single number. The same recipe as
1238    `primer.ml.neural_net.MLP`, generalised to any depth and to a backward
1239    pass that also returns the gradient with respect to the *input*, which is
1240    the only route by which the forger ever learns anything.
1241    """
1242
1243    def __init__(self, sizes: tuple[int, ...], seed: int = 0):
1244        rng = np.random.default_rng(seed)
1245        shapes = [(a, b) for a, b in zip(sizes[:-1], sizes[1:])]
1246        self.params = np.zeros(sum(a * b + b for a, b in shapes))
1247        self.W, self.b = self._views(self.params, shapes)
1248        self._shapes = shapes
1249        for W in self.W:
1250            # Variance 1/fan_in keeps tanh out of its flat ends at the start (see primer.ml.deep_nets).
1251            W[...] = rng.normal(0.0, 1.0 / np.sqrt(W.shape[0]), W.shape)
1252
1253    @staticmethod
1254    def _views(flat: np.ndarray, shapes) -> tuple[list[np.ndarray], list[np.ndarray]]:
1255        Ws, bs, i = [], [], 0
1256        for a, b in shapes:
1257            Ws.append(flat[i : i + a * b].reshape(a, b))
1258            i += a * b
1259            bs.append(flat[i : i + b])
1260            i += b
1261        return Ws, bs
1262
1263    def forward(self, x: np.ndarray) -> np.ndarray:
1264        """x: (n, sizes[0]) -> (n, sizes[-1]). Caches every layer's output for `backward`."""
1265        self._acts = [x]
1266        for i, (W, b) in enumerate(zip(self.W, self.b)):
1267            z = self._acts[-1] @ W + b
1268            self._acts.append(np.tanh(z) if i < len(self.W) - 1 else z)
1269        return self._acts[-1]
1270
1271    def backward(self, grad_out: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
1272        """Chain rule from d(something)/d(output) back to (every parameter, the input).
1273
1274        Returns (flat gradient laid out like `params`, gradient w.r.t. the input x).
1275        """
1276        grad = np.zeros_like(self.params)
1277        gW, gb = self._views(grad, self._shapes)
1278        g = grad_out
1279        for i in reversed(range(len(self.W))):
1280            gW[i][...] = self._acts[i].T @ g  # each weight's gradient: its input times the error arriving
1281            gb[i][...] = g.sum(axis=0)
1282            g = g @ self.W[i].T  # send the error back through the weights
1283            if i > 0:
1284                g = g * (1.0 - self._acts[i] ** 2)  # tanh'(z) = 1 - tanh(z)^2, reusing the cached output
1285        return grad, g

A stack of layers x W + b, with tanh between them and nothing after the last.

Every weight and bias lives in one flat vector, params; W and b are views into it. That lets one optimizer step update the whole network in one line, and lets a test nudge any single number. The same recipe as primer.ml.neural_net.MLP, generalised to any depth and to a backward pass that also returns the gradient with respect to the input, which is the only route by which the forger ever learns anything.

MLP(sizes: tuple[int, ...], seed: int = 0) on GitHub
1243    def __init__(self, sizes: tuple[int, ...], seed: int = 0):
1244        rng = np.random.default_rng(seed)
1245        shapes = [(a, b) for a, b in zip(sizes[:-1], sizes[1:])]
1246        self.params = np.zeros(sum(a * b + b for a, b in shapes))
1247        self.W, self.b = self._views(self.params, shapes)
1248        self._shapes = shapes
1249        for W in self.W:
1250            # Variance 1/fan_in keeps tanh out of its flat ends at the start (see primer.ml.deep_nets).
1251            W[...] = rng.normal(0.0, 1.0 / np.sqrt(W.shape[0]), W.shape)
params
def forward(self, x: numpy.ndarray) -> numpy.ndarray: on GitHub
1263    def forward(self, x: np.ndarray) -> np.ndarray:
1264        """x: (n, sizes[0]) -> (n, sizes[-1]). Caches every layer's output for `backward`."""
1265        self._acts = [x]
1266        for i, (W, b) in enumerate(zip(self.W, self.b)):
1267            z = self._acts[-1] @ W + b
1268            self._acts.append(np.tanh(z) if i < len(self.W) - 1 else z)
1269        return self._acts[-1]

x: (n, sizes[0]) -> (n, sizes[-1]). Caches every layer's output for backward.

def backward(self, grad_out: numpy.ndarray) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
1271    def backward(self, grad_out: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
1272        """Chain rule from d(something)/d(output) back to (every parameter, the input).
1273
1274        Returns (flat gradient laid out like `params`, gradient w.r.t. the input x).
1275        """
1276        grad = np.zeros_like(self.params)
1277        gW, gb = self._views(grad, self._shapes)
1278        g = grad_out
1279        for i in reversed(range(len(self.W))):
1280            gW[i][...] = self._acts[i].T @ g  # each weight's gradient: its input times the error arriving
1281            gb[i][...] = g.sum(axis=0)
1282            g = g @ self.W[i].T  # send the error back through the weights
1283            if i > 0:
1284                g = g * (1.0 - self._acts[i] ** 2)  # tanh'(z) = 1 - tanh(z)^2, reusing the cached output
1285        return grad, g

Chain rule from d(something)/d(output) back to (every parameter, the input).

Returns (flat gradient laid out like params, gradient w.r.t. the input x).

def value_function(d_real, d_fake, eps: float = 1e-12) -> float: on GitHub
1293def value_function(d_real, d_fake, eps: float = 1e-12) -> float:
1294    """V = mean log D(x) over real samples + mean log(1 - D(G(z))) over fakes.
1295
1296    The detective wants V high (ceiling 0); the forger wants it low. -V is
1297    exactly the detective's binary cross-entropy with real = 1 and fake = 0.
1298    """
1299    d_real = np.clip(np.asarray(d_real, dtype=float), eps, 1.0)
1300    d_fake = np.clip(np.asarray(d_fake, dtype=float), 0.0, 1.0 - eps)
1301    return float(np.mean(np.log(d_real)) + np.mean(np.log(1.0 - d_fake)))

V = mean log D(x) over real samples + mean log(1 - D(G(z))) over fakes.

The detective wants V high (ceiling 0); the forger wants it low. -V is exactly the detective's binary cross-entropy with real = 1 and fake = 0.

def optimal_discriminator(p_data, p_g): on GitHub
1304def optimal_discriminator(p_data, p_g):
1305    """D*(x) = p_data(x) / (p_data(x) + p_g(x)): the best any detective can do against a fixed forger."""
1306    p_data, p_g = np.asarray(p_data, dtype=float), np.asarray(p_g, dtype=float)
1307    return p_data / (p_data + p_g)

D*(x) = p_data(x) / (p_data(x) + p_g(x)): the best any detective can do against a fixed forger.

def fit_discriminator_table(p_data, p_g, steps: int = 4000, lr: float = 5.0) -> numpy.ndarray: on GitHub
1310def fit_discriminator_table(p_data, p_g, steps: int = 4000, lr: float = 5.0) -> np.ndarray:
1311    """Train the most flexible detective there is, one free score per point, and return its verdicts.
1312
1313    The data and the forger live on the same few points with the given
1314    probabilities. Gradient ascent on the exact expected V: for a point with
1315    score a and verdict D = sigmoid(a), dV/da = p_data (1 - D) - p_g D. It is
1316    zero exactly when D = p_data / (p_data + p_g), so the trained table
1317    arrives at `optimal_discriminator` without ever being told the formula.
1318    """
1319    p_data, p_g = np.asarray(p_data, dtype=float), np.asarray(p_g, dtype=float)
1320    a = np.zeros_like(p_data)
1321    for _ in range(steps):
1322        d = sigmoid(a)
1323        a += lr * (p_data * (1.0 - d) - p_g * d)
1324    return sigmoid(a)

Train the most flexible detective there is, one free score per point, and return its verdicts.

The data and the forger live on the same few points with the given probabilities. Gradient ascent on the exact expected V: for a point with score a and verdict D = sigmoid(a), dV/da = p_data (1 - D) - p_g D. It is zero exactly when D = p_data / (p_data + p_g), so the trained table arrives at optimal_discriminator without ever being told the formula.

def generator_logit_gradient(d_fake, loss: str = 'non_saturating'): on GitHub
1332def generator_logit_gradient(d_fake, loss: str = "non_saturating"):
1333    """The forger's loss gradient with respect to the detective's score a, where D(G(z)) = sigmoid(a).
1334
1335    saturating (the original):  loss = log(1 - D)  ->  d loss / da = -D
1336    non_saturating:             loss = -log D      ->  d loss / da = -(1 - D)
1337
1338    When the detective is sure a fake is fake (D near 0), the first is near 0
1339    and the second near -1: only the second still tells the forger which way to move.
1340    """
1341    d = np.asarray(d_fake, dtype=float)
1342    out = -d if loss == "saturating" else -(1.0 - d)
1343    return float(out) if out.ndim == 0 else out

The forger's loss gradient with respect to the detective's score a, where D(G(z)) = sigmoid(a).

saturating (the original): loss = log(1 - D) -> d loss / da = -D non_saturating: loss = -log D -> d loss / da = -(1 - D)

When the detective is sure a fake is fake (D near 0), the first is near 0 and the second near -1: only the second still tells the forger which way to move.

def head_start_experiment( head_starts=(0, 25, 50, 100, 200, 300), lr: float = 0.003, batch: int = 128, seed: int = 0) -> list[dict]: on GitHub
1355def head_start_experiment(head_starts=(0, 25, 50, 100, 200, 300), lr: float = 3e-3, batch: int = 128, seed: int = 0) -> list[dict]:
1356    """Train only the detective against a frozen, untrained forger, then measure what the forger would learn.
1357
1358    For each head start (detective steps), returns the average verdict on
1359    fakes and the size (norm) of the forger's parameter gradient under each
1360    loss. The untrained forger's points sit near the centre, far from the
1361    ring, so the detective quickly becomes certain, and the original loss's
1362    gradient fades while the non-saturating one stays large.
1363    """
1364    rng = np.random.default_rng(seed)
1365    G, D = MLP(GENERATOR_SIZES, seed=seed), MLP(DISCRIMINATOR_SIZES, seed=seed + 1)
1366    opt = Adam(lr=lr, beta1=0.5)
1367    probe_z = np.random.default_rng(seed + 100).standard_normal((batch, 2))
1368    rows, done = [], 0
1369    for target in sorted(head_starts):
1370        while done < target:
1371            _detective_step(D, opt, ring_of_gaussians(batch, rng), G.forward(rng.standard_normal((batch, 2))))
1372            done += 1
1373        rows.append(
1374            dict(
1375                head_start=target,
1376                d_fake=float(sigmoid(D.forward(G.forward(probe_z))).mean()),
1377                saturating=float(np.linalg.norm(_forger_gradient(G, D, probe_z, "saturating"))),
1378                non_saturating=float(np.linalg.norm(_forger_gradient(G, D, probe_z, "non_saturating"))),
1379            )
1380        )
1381    return rows

Train only the detective against a frozen, untrained forger, then measure what the forger would learn.

For each head start (detective steps), returns the average verdict on fakes and the size (norm) of the forger's parameter gradient under each loss. The untrained forger's points sit near the centre, far from the ring, so the detective quickly becomes certain, and the original loss's gradient fades while the non-saturating one stays large.

def mode_coverage( samples: numpy.ndarray, min_share: float = 0.02, std: float = 0.05) -> dict: on GitHub
1398def mode_coverage(samples: np.ndarray, min_share: float = 0.02, std: float = MODE_STD) -> dict:
1399    """How much of the ring the samples cover.
1400
1401    A sample is *high quality* if it lands within 3 standard deviations of a
1402    mode's centre. A mode is *covered* if it holds at least `min_share` of
1403    all samples (2% by default: 20 of 1,000, where a perfect forger puts 125).
1404    Returns covered (a count), which (a string such as ".X...X.." marking the
1405    covered modes), counts per mode, and quality (the high-quality fraction).
1406    """
1407    distances = np.linalg.norm(samples[:, None, :] - ring_modes()[None], axis=2)
1408    nearest, close = distances.argmin(axis=1), distances.min(axis=1) < 3 * std
1409    counts = np.bincount(nearest[close], minlength=N_MODES)
1410    hit = counts >= min_share * len(samples)
1411    return dict(
1412        covered=int(hit.sum()),
1413        which="".join("X" if h else "." for h in hit),
1414        counts=counts,
1415        quality=float(close.mean()),
1416    )

How much of the ring the samples cover.

A sample is high quality if it lands within 3 standard deviations of a mode's centre. A mode is covered if it holds at least min_share of all samples (2% by default: 20 of 1,000, where a perfect forger puts 125). Returns covered (a count), which (a string such as ".X...X.." marking the covered modes), counts per mode, and quality (the high-quality fraction).

@functools.lru_cache(maxsize=None)
def train_gan( steps: int = 2500, lr_g: float = 0.001, lr_d: float = 0.003, batch: int = 128, snapshot_every: int = 250, loss: str = 'non_saturating', seed: int = 1) -> dict: on GitHub
1419@functools.lru_cache(maxsize=None)
1420def train_gan(
1421    steps: int = 2500,
1422    lr_g: float = 1e-3,
1423    lr_d: float = 3e-3,
1424    batch: int = 128,
1425    snapshot_every: int = 250,
1426    loss: str = "non_saturating",
1427    seed: int = 1,
1428) -> dict:
1429    """Train a forger and a detective on the ring, alternating one step each.
1430
1431    Returns snapshots every `snapshot_every` steps: the step, 1,000 generated
1432    points (always from the same noise, so pictures are comparable), the
1433    detective's parameters, and the `mode_coverage` of the points. Cached,
1434    because the lesson, its figures and its tests all look at the same runs.
1435    """
1436    rng = np.random.default_rng(seed)
1437    G, D = MLP(GENERATOR_SIZES, seed=seed), MLP(DISCRIMINATOR_SIZES, seed=seed + 1)
1438    # beta1 = 0.5 rather than 0.9: less momentum, the usual choice for GANs since DCGAN, because the target keeps moving.
1439    opt_g, opt_d = Adam(lr=lr_g, beta1=0.5), Adam(lr=lr_d, beta1=0.5)
1440    probe_z = np.random.default_rng(seed + 100).standard_normal((1000, 2))
1441    snapshots = []
1442    for step in range(steps + 1):
1443        if step % snapshot_every == 0:
1444            points = G.forward(probe_z).copy()
1445            snapshots.append(dict(step=step, samples=points, detective=D.params.copy(), coverage=mode_coverage(points)))
1446        if step == steps:
1447            break
1448        z = rng.standard_normal((batch, 2))
1449        fakes = G.forward(z)
1450        # 1. Detective: one step towards telling this batch of real points from this batch of fakes.
1451        _detective_step(D, opt_d, ring_of_gaussians(batch, rng), fakes)
1452        # 2. Forger: one step towards fooling the detective as it is *now*.
1453        G.params[:] = opt_g.step(G.params, _forger_gradient(G, D, z, loss))
1454    return dict(snapshots=snapshots, coverage=[s["coverage"] for s in snapshots], lr_g=lr_g, lr_d=lr_d)

Train a forger and a detective on the ring, alternating one step each.

Returns snapshots every snapshot_every steps: the step, 1,000 generated points (always from the same noise, so pictures are comparable), the detective's parameters, and the mode_coverage of the points. Cached, because the lesson, its figures and its tests all look at the same runs.

def detective_verdicts(params: numpy.ndarray, points: numpy.ndarray) -> numpy.ndarray: on GitHub
1457def detective_verdicts(params: np.ndarray, points: np.ndarray) -> np.ndarray:
1458    """D(x) for each point, using a detective's saved parameters."""
1459    D = MLP(DISCRIMINATOR_SIZES)
1460    D.params[:] = params
1461    return sigmoid(D.forward(points))[:, 0]

D(x) for each point, using a detective's saved parameters.

def dirac_gan( theta: float = 1.0, psi: float = 0.0, lr: float = 0.2, steps: int = 300, gamma: float = 0.0, alternating: bool = False) -> numpy.ndarray: on GitHub
1469def dirac_gan(theta: float = 1.0, psi: float = 0.0, lr: float = 0.2, steps: int = 300, gamma: float = 0.0, alternating: bool = False) -> np.ndarray:
1470    """The smallest GAN there is: two numbers, followed step by step. Returns the path, shape (steps + 1, 2).
1471
1472    Real data is always the single point x = 0. The forger has one number,
1473    theta, and always outputs x = theta. The detective has one number, psi,
1474    and says D(x) = sigmoid(psi x). The value is
1475    V = log D(0) + log(1 - D(theta)) = log 0.5 + log(1 - sigmoid(psi theta)).
1476
1477        dV/dpsi   = -sigmoid(psi theta) theta   (the detective climbs this)
1478        dV/dtheta = -sigmoid(psi theta) psi     (the forger descends this)
1479
1480    `gamma` > 0 adds the R1 gradient penalty (gamma / 2)(dD-score/dx at the
1481    real data)^2 = (gamma / 2) psi^2 to what the detective minimises.
1482    `alternating` lets the forger see the detective's new psi before it moves.
1483    """
1484    path = [(theta, psi)]
1485    for _ in range(steps):
1486        s = sigmoid(psi * theta)
1487        new_psi = psi + lr * (-s * theta - gamma * psi)
1488        if alternating:
1489            psi = new_psi
1490            s = sigmoid(psi * theta)
1491        theta = theta + lr * s * psi  # descend dV/dtheta = -s psi
1492        psi = new_psi
1493        path.append((float(theta), float(psi)))
1494    return np.array(path)

The smallest GAN there is: two numbers, followed step by step. Returns the path, shape (steps + 1, 2).

Real data is always the single point x = 0. The forger has one number, theta, and always outputs x = theta. The detective has one number, psi, and says D(x) = sigmoid(psi x). The value is V = log D(0) + log(1 - D(theta)) = log 0.5 + log(1 - sigmoid(psi theta)).

dV/dpsi   = -sigmoid(psi theta) theta   (the detective climbs this)
dV/dtheta = -sigmoid(psi theta) psi     (the forger descends this)

gamma > 0 adds the R1 gradient penalty (gamma / 2)(dD-score/dx at the real data)^2 = (gamma / 2) psi^2 to what the detective minimises. alternating lets the forger see the detective's new psi before it moves.

def js_divergence(p, q) -> float: on GitHub
1502def js_divergence(p, q) -> float:
1503    """Jensen-Shannon divergence between two histograms on the same bins (natural log).
1504
1505    JS = 1/2 KL(p || m) + 1/2 KL(q || m), with m the average of p and q. It
1506    is 0 when p = q and log 2 when they share no bin at all, *however far
1507    apart* the bins are. That flatness is why a GAN's detective can stop
1508    giving useful directions.
1509    """
1510    p, q = np.asarray(p, dtype=float), np.asarray(q, dtype=float)
1511    p, q = p / p.sum(), q / q.sum()
1512    m = 0.5 * (p + q)
1513
1514    def kl(a, b):
1515        mask = a > 0  # 0 log 0 counts as 0
1516        return float(np.sum(a[mask] * np.log(a[mask] / b[mask])))
1517
1518    return 0.5 * kl(p, m) + 0.5 * kl(q, m)

Jensen-Shannon divergence between two histograms on the same bins (natural log).

JS = 1/2 KL(p || m) + 1/2 KL(q || m), with m the average of p and q. It is 0 when p = q and log 2 when they share no bin at all, however far apart the bins are. That flatness is why a GAN's detective can stop giving useful directions.

def wasserstein_1d(a, b) -> float: on GitHub
1521def wasserstein_1d(a, b) -> float:
1522    """Earth mover's distance between two equal-sized 1-D samples: sort both, pair them up, average the gaps."""
1523    a, b = np.sort(np.asarray(a, dtype=float)), np.sort(np.asarray(b, dtype=float))
1524    return float(np.mean(np.abs(a - b)))

Earth mover's distance between two equal-sized 1-D samples: sort both, pair them up, average the gaps.

def largest_stretch(W: numpy.ndarray, iters: int = 100, seed: int = 0) -> float: on GitHub
1527def largest_stretch(W: np.ndarray, iters: int = 100, seed: int = 0) -> float:
1528    """The most W can lengthen any vector (its spectral norm), by power iteration.
1529
1530    Push a vector through W and back through W transpose, again and again;
1531    it swings round to the direction W stretches most, and the stretch
1532    factor is read off at the end. Spectral normalization runs one step of
1533    this per training step instead of a full singular value decomposition.
1534    """
1535    v = np.random.default_rng(seed).standard_normal(W.shape[1])
1536    for _ in range(iters):
1537        u = W @ v
1538        u /= np.linalg.norm(u)
1539        v = W.T @ u
1540        v /= np.linalg.norm(v)
1541    return float(np.linalg.norm(W @ v))

The most W can lengthen any vector (its spectral norm), by power iteration.

Push a vector through W and back through W transpose, again and again; it swings round to the direction W stretches most, and the stretch factor is read off at the end. Spectral normalization runs one step of this per training step instead of a full singular value decomposition.

def spectrally_normalize(W: numpy.ndarray) -> numpy.ndarray: on GitHub
1544def spectrally_normalize(W: np.ndarray) -> np.ndarray:
1545    """W divided by its largest stretch, so no input direction is stretched by more than 1."""
1546    return W / largest_stretch(W)

W divided by its largest stretch, so no input direction is stretched by more than 1.

RING_RUNS = {'equal speeds (0.001)': 0.001, 'detective 3x faster (0.003)': 0.003, 'detective 10x faster (0.01)': 0.01}
def figures() -> dict: on GitHub
1561def figures() -> dict:
1562    """Plot this lesson's data. matplotlib is imported here, and only here,
1563    so the lesson itself needs nothing beyond NumPy."""
1564    import matplotlib
1565
1566    matplotlib.use("Agg")
1567    import matplotlib.pyplot as plt
1568
1569    REAL, FAKE, GOOD, MUTED, THIRD = "#2563eb", "#dc2626", "#059669", "#9ca3af", "#d97706"
1570    run_colours = dict(zip(RING_RUNS, (FAKE, REAL, THIRD)))
1571    runs = {name: train_gan(lr_d=lr_d) for name, lr_d in RING_RUNS.items()}
1572    figs = {}
1573
1574    # --- 1. The best detective on a line -------------------------------------
1575    xs = np.linspace(-5, 5, 400)
1576    p_data = 0.5 * _normal_pdf(xs, -2, 0.5) + 0.5 * _normal_pdf(xs, 2, 0.5)
1577    p_g = _normal_pdf(xs, 0, 1.5)
1578    fig, (top, bottom) = plt.subplots(2, 1, figsize=(6.4, 5), sharex=True)
1579    top.fill_between(xs, p_data, color=REAL, alpha=0.25)
1580    top.plot(xs, p_data, color=REAL, label="real data  p_data")
1581    top.fill_between(xs, p_g, color=FAKE, alpha=0.15)
1582    top.plot(xs, p_g, color=FAKE, label="forger  p_g")
1583    top.set_ylabel("density")
1584    top.set_title("Where each distribution puts its points")
1585    top.legend(frameon=False, loc="upper center")
1586    bottom.plot(xs, optimal_discriminator(p_data, p_g), color="#111827", label="best verdict  D*")
1587    bottom.axhline(0.5, color=MUTED, ls="--", label="D* once the forger matches the data")
1588    bottom.set_ylim(-0.03, 1.03)
1589    bottom.set_xlabel("position x")
1590    bottom.set_ylabel("D*(x)")
1591    bottom.set_title("The best detective: the share of points at x that are real")
1592    bottom.legend(frameon=False, loc="lower right", fontsize=8)
1593    fig.tight_layout()
1594    figs["optimal_detective"] = fig
1595
1596    # --- 2. The trained detective's view of the ring -------------------------
1597    last = runs["detective 3x faster (0.003)"]["snapshots"][-1]
1598    grid = np.linspace(-3, 3, 161)
1599    gx, gy = np.meshgrid(grid, grid)
1600    verdict = detective_verdicts(last["detective"], np.stack([gx.ravel(), gy.ravel()], axis=1)).reshape(gx.shape)
1601    real = ring_of_gaussians(1000, seed=5)
1602    fig, ax = plt.subplots(figsize=(5.6, 4.8))
1603    im = ax.imshow(verdict, origin="lower", extent=(-3, 3, -3, 3), cmap="RdBu", vmin=0, vmax=1)
1604    ax.scatter(real[:, 0], real[:, 1], s=3, color="#111827", label="real")
1605    ax.scatter(last["samples"][:, 0], last["samples"][:, 1], s=3, color=GOOD, alpha=0.6, label="forger")
1606    ax.set_aspect("equal")
1607    ax.grid(False)
1608    ax.set_title(f"The detective's verdict after {last['step']:,} steps")
1609    ax.legend(frameon=True, loc="upper right", markerscale=3)
1610    fig.colorbar(im, ax=ax, fraction=0.046, label="D(x): probability real")
1611    figs["detective_view"] = fig
1612
1613    # --- 3. Saturating vs non-saturating forger loss -------------------------
1614    scores = np.linspace(-7, 3, 300)
1615    rows = head_start_experiment(head_starts=(0, 50, 100, 200, 300, 400, 600, 800))
1616    fig, (a1, a2) = plt.subplots(1, 2, figsize=(10, 3.8))
1617    a1.plot(scores, np.log(1 - sigmoid(scores)), color=FAKE, label="original: log(1 - D)")
1618    a1.plot(scores, -np.log(sigmoid(scores)), color=REAL, label="non-saturating: -log D")
1619    a1.axvline(math.log(0.01 / 0.99), color=MUTED, ls="--")
1620    a1.text(math.log(0.01 / 0.99) + 0.15, 5.6, "D = 0.01", color="#4b5563")
1621    a1.set_xlabel("detective's score a for a fake  (left: sure it is fake)")
1622    a1.set_ylabel("forger's loss")
1623    a1.set_title("The original loss goes flat when the detective is sure")
1624    a1.legend(frameon=False, loc="upper right")
1625    steps = [r["head_start"] for r in rows]
1626    a2.semilogy(steps, [r["saturating"] for r in rows], "o-", color=FAKE, label="original loss")
1627    a2.semilogy(steps, [r["non_saturating"] for r in rows], "o-", color=REAL, label="non-saturating loss")
1628    for r in rows[::2]:
1629        a2.annotate(f"D={r['d_fake']:.3f}", (r["head_start"], r["non_saturating"]), textcoords="offset points", xytext=(-10, 8), fontsize=7, color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1630    a2.set_xlabel("detective's head start (training steps)")
1631    a2.set_ylabel("size of the forger's gradient")
1632    a2.set_title("Measured on the ring: a head start starves one loss")
1633    a2.legend(frameon=False, loc="center right")
1634    fig.tight_layout()
1635    figs["forger_losses"] = fig
1636
1637    # --- 4. The Dirac GAN: spiral out, or orbit ------------------------------
1638    simultaneous, alternating, penalised = dirac_gan(), dirac_gan(alternating=True), dirac_gan(gamma=1.0)
1639    fig, ax = plt.subplots(figsize=(5.4, 5))
1640    ax.plot(simultaneous[:, 0], simultaneous[:, 1], color=FAKE, lw=1.2, label="simultaneous steps: spiral out")
1641    ax.plot(alternating[:, 0], alternating[:, 1], color=REAL, lw=1.2, label="alternating steps: endless orbit")
1642    ax.plot([1], [0], "o", color="#111827")
1643    ax.annotate("start (1, 0)", (1, 0), textcoords="offset points", xytext=(6, 6), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1644    ax.plot([0], [0], "x", color="#111827", ms=10, mew=2)
1645    ax.annotate("equilibrium", (0, 0), textcoords="offset points", xytext=(6, -12), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1646    ax.set_aspect("equal")
1647    ax.set_xlabel("forger's theta (where the fake sits)")
1648    ax.set_ylabel("detective's psi (its slope)")
1649    ax.set_title("The Dirac GAN never settles")
1650    ax.legend(frameon=False, loc="upper left", fontsize=8)
1651    figs["dirac_oscillation"] = fig
1652
1653    # --- 5. Ring snapshots: hopping, covering, stuck -------------------------
1654    shown = (1750, 2000, 2250, 2500)
1655    modes = ring_modes()
1656    fig, axes = plt.subplots(3, 4, figsize=(10, 7.8), sharex=True, sharey=True)
1657    for row, (name, run) in zip(axes, runs.items()):
1658        by_step = {snap["step"]: snap for snap in run["snapshots"]}
1659        for ax, step in zip(row, shown):
1660            snap = by_step[step]
1661            for centre in modes:
1662                ax.add_patch(plt.Circle(centre, 3 * MODE_STD, color=MUTED, alpha=0.5))
1663            ax.scatter(snap["samples"][:, 0], snap["samples"][:, 1], s=2, color=GOOD, alpha=0.5)
1664            ax.set_title(f"step {step:,}  {snap['coverage']['which']}", fontsize=9, family="monospace")
1665            ax.set_xlim(-3, 3)
1666            ax.set_ylim(-3, 3)
1667            ax.set_aspect("equal")
1668            ax.set_xticks([])
1669            ax.set_yticks([])
1670        row[0].set_ylabel(name, fontsize=9)
1671    fig.suptitle("The forger's points (green) against the eight real clouds (grey)", fontweight="bold")
1672    fig.tight_layout()
1673    figs["ring_snapshots"] = fig
1674
1675    # --- 6. Mode coverage over training ----------------------------------------
1676    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1677    for name, run in runs.items():
1678        ax.step([s["step"] for s in run["snapshots"]], [c["covered"] for c in run["coverage"]], where="post", color=run_colours[name], lw=2, label=name)
1679    ax.set_ylim(-0.3, 8.5)
1680    ax.set_yticks(range(0, 9))
1681    ax.set_xlabel("training step")
1682    ax.set_ylabel("clouds covered (of 8)")
1683    ax.set_title("How much of the ring the forger covers")
1684    ax.legend(frameon=False, loc="upper left", fontsize=8)
1685    figs["mode_coverage"] = fig
1686
1687    # --- 7. R1 penalty: the spiral turns inward --------------------------------
1688    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1689    for path, colour, label in (
1690        (simultaneous, FAKE, "simultaneous steps"),
1691        (alternating, REAL, "alternating steps"),
1692        (penalised, GOOD, "simultaneous + R1 penalty (gamma = 1)"),
1693    ):
1694        ax.plot(np.linalg.norm(path, axis=1), color=colour, lw=2, label=label)
1695    ax.set_xlabel("step")
1696    ax.set_ylabel("distance from equilibrium")
1697    ax.set_title("A gradient penalty damps the chase")
1698    ax.legend(frameon=False, loc="upper left", fontsize=8)
1699    figs["dirac_r1"] = fig
1700
1701    # --- 8. Jensen-Shannon vs Wasserstein --------------------------------------
1702    bins = np.round(np.arange(-40, 41) / 10, 1)  # positions -4.0 ... 4.0 in steps of 0.1
1703    data = (bins == 0.0).astype(float)
1704    thetas = bins
1705    js = [js_divergence(data, (bins == t).astype(float)) for t in thetas]
1706    w = [wasserstein_1d([0.0], [t]) for t in thetas]
1707    fig, ax = plt.subplots(figsize=(6.4, 3.6))
1708    ax.plot(thetas, js, ".", color=FAKE, ms=4, label="Jensen-Shannon divergence")
1709    ax.plot(thetas, w, color=REAL, lw=2, label="Wasserstein distance")
1710    ax.axhline(math.log(2), color=MUTED, ls="--")
1711    ax.text(-3.9, math.log(2) + 0.12, "log 2 = 0.693", color="#4b5563")
1712    ax.set_xlabel("where the forger's pile sits (the real pile is at 0)")
1713    ax.set_ylabel("distance between the piles")
1714    ax.set_title("Only one yardstick says which way to move")
1715    ax.legend(frameon=False, loc="upper center")
1716    figs["wasserstein_vs_js"] = fig
1717
1718    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1726def demo() -> None:
1727    banner("1. The two players")
1728    G, D = MLP(GENERATOR_SIZES, seed=1), MLP(DISCRIMINATOR_SIZES, seed=2)
1729    fakes = G.forward(np.random.default_rng(0).standard_normal((1000, 2)))
1730    real = ring_of_gaussians(1000, seed=0)
1731    say(
1732        f"""
1733        The real data is a ring of eight little clouds of radius {RADIUS}. The
1734        forger maps 2 random numbers to a point through layers {GENERATOR_SIZES};
1735        the detective maps a point to a probability through {DISCRIMINATOR_SIZES}
1736        and a sigmoid. Untrained, the forger's points sit an average of
1737        {np.linalg.norm(fakes, axis=1).mean():.2f} from the centre, while real
1738        points sit {np.linalg.norm(real, axis=1).mean():.2f} away, and the
1739        untrained detective says {sigmoid(D.forward(real)).mean():.2f} on real
1740        points and {sigmoid(D.forward(fakes)).mean():.2f} on fakes: it knows nothing yet.
1741        """
1742    )
1743
1744    banner("2. The game: one number both players fight over")
1745    table(
1746        ["detective", "D on real notes", "D on fakes", "V"],
1747        [
1748            ("the worked example", "0.9, 0.8", "0.2, 0.4", value_function([0.9, 0.8], [0.2, 0.4])),
1749            ("never wrong", "1, 1", "0, 0", value_function([1.0, 1.0], [0.0, 0.0])),
1750            ("coin flip", "0.5, 0.5", "0.5, 0.5", value_function([0.5, 0.5], [0.5, 0.5])),
1751        ],
1752        floatfmt=".3f",
1753    )
1754    takeaway("The detective pushes V up towards 0; the forger pushes it down. -V is the detective's binary cross-entropy.")
1755
1756    banner("3. The best possible detective")
1757    p_data, p_g = [0.4, 0.3, 0.2, 0.1], [0.1, 0.1, 0.4, 0.4]
1758    trained = fit_discriminator_table(p_data, p_g)
1759    table(
1760        ["point", "p_data", "p_g", "formula D*", "trained detective"],
1761        [(i + 1, pd, pg, optimal_discriminator(pd, pg), t) for i, (pd, pg, t) in enumerate(zip(p_data, p_g, trained))],
1762        floatfmt=".3f",
1763    )
1764    say("A detective trained against a fixed forger lands on p_data / (p_data + p_g) without being told the formula.")
1765    takeaway("At equilibrium the forger matches the data, so the best detective says 1/2 everywhere and V = -log 4.")
1766
1767    banner("4. Why the forger uses the non-saturating loss")
1768    table(
1769        ["D(G(z))", "original loss slope", "non-saturating slope"],
1770        [(d, generator_logit_gradient(d, "saturating"), generator_logit_gradient(d, "non_saturating")) for d in (0.5, 0.1, 0.01, 0.001)],
1771        floatfmt=".3f",
1772    )
1773    table(
1774        ["detective head start", "D on fakes", "forger gradient, original", "forger gradient, non-saturating"],
1775        [(r["head_start"], r["d_fake"], r["saturating"], r["non_saturating"]) for r in head_start_experiment(head_starts=(0, 200, 800))],
1776        floatfmt=".3f",
1777    )
1778    takeaway("When the detective is sure, the original loss goes quiet; the non-saturating one is loudest exactly then.")
1779
1780    banner("5. Oscillation: the two-number Dirac GAN")
1781    table(
1782        ["steps of size 0.2", "distance after 50", "after 100", "after 300"],
1783        [
1784            (name, *(float(np.linalg.norm(path[i])) for i in (50, 100, 300)))
1785            for name, path in (
1786                ("simultaneous", dirac_gan()),
1787                ("alternating", dirac_gan(alternating=True)),
1788                ("simultaneous + R1 penalty", dirac_gan(gamma=1.0)),
1789            )
1790        ],
1791        floatfmt=".3f",
1792    )
1793    say("Start 1 away from the equilibrium. Plain steps spiral out or orbit; the R1 penalty's friction brings them home.")
1794
1795    banner("6. Mode collapse on the ring")
1796    for name, lr_d in RING_RUNS.items():
1797        run = train_gan(lr_d=lr_d)
1798        print(f"{name:28s} " + " ".join(c["which"] for c in run["coverage"][5:]))
1799    print()
1800    table(
1801        ["detective learning rate", "clouds covered at the end", "high-quality points"],
1802        [(name, train_gan(lr_d=lr_d)["coverage"][-1]["covered"], train_gan(lr_d=lr_d)["coverage"][-1]["quality"]) for name, lr_d in RING_RUNS.items()],
1803        floatfmt=".2f",
1804    )
1805    say(
1806        """
1807        Each row shows coverage every 250 steps from step 1,250 (X = covered).
1808        Equal speeds: the forger hops between alternate clouds. A detective
1809        three times faster: all eight. Ten times faster: six, and stuck.
1810        """
1811    )
1812    takeaway("Mode collapse is invisible sample by sample; you have to measure coverage.")
1813
1814    banner("7. Yardsticks and speed limits")
1815    grid = np.arange(11)
1816    table(
1817        ["forger's pile at", "Jensen-Shannon", "Wasserstein"],
1818        [(t, js_divergence(grid == 0, grid == t), wasserstein_1d([0.0], [float(t)])) for t in (1, 5, 10)],
1819        floatfmt=".3f",
1820    )
1821    W = np.array([[3.0, 0.0], [0.0, 1.0]])
1822    say(
1823        f"""
1824        Jensen-Shannon can't tell a near miss from a far one; Wasserstein can.
1825        Spectral normalization: W = [[3, 0], [0, 1]] has largest stretch
1826        {largest_stretch(W):.3f}; after dividing by it, {largest_stretch(spectrally_normalize(W)):.3f}.
1827        """
1828    )
1829    takeaway("Keep the detective useful: paced, smooth and never infinitely sure.")