Image Generation and Restoration
Why GANs are hard to train
A GAN is two networks fighting, and the fight can stall or collapse in ways that ordinary training never does.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A GAN trains two networks against each other, and the two have to stay evenly matched.
The analogy
Picture a village cricket match. One side bats, the other bowls. Both improve only because the other keeps testing them.
Now imagine the bowler becomes unplayable overnight. The batter is out first ball, every ball, and learns nothing. The match has not become harder. It has stopped teaching.
A GAN works the same way. A GAN is a pair of networks: one invents pictures, one tries to spot fakes. Take away the balance and both stop learning.
The two players
The generator invents pictures from random numbers. It has never seen a real photo.
The discriminator is shown one picture at a time. It answers one question: real or invented?
Everything the generator knows arrives second hand. It only ever hears whether its last attempt fooled the judge.
random numbers -> [ generator ] -> a picture
|
a real photo ------------------>|
v
[ discriminator ] -> "real" or "fake"
|
the verdict is the only teacher the generator hasWhy this is fragile
Ordinary training has a fixed target. You show a model a mango and the right answer stays "mango" forever.
A GAN has no fixed target. The generator chases a judge that is itself changing. Both aim at something that moves when they move.
That is the root of every failure on this page. Read that paragraph twice if it feels slippery. It confuses almost everyone the first time.
The four ways it goes wrong
The judge wins too hard. If the discriminator becomes perfect, every fake is rejected outright. The verdict stops carrying any hint about which fake was closer.
The judge is too soft. If the discriminator cannot tell anything apart, its verdict is noise. The generator wanders.
Mode collapse. The generator finds one picture that fools the judge, then produces that picture forever. It has stopped inventing.
Oscillation. The two chase each other in circles. Both losses look busy. Nothing improves.
How you would notice
Loss curves lie here, and that surprises people. In normal training, falling loss means progress. In a GAN, a falling generator loss can mean the discriminator has gone soft.
The honest check is to look at the pictures. Are they varied? Do you keep seeing the same face?
Where you have seen this
- Face generators that produce endless slight variations of one face.
- Photo upscalers that invent plausible detail rather than recovering real detail.
- Fashion and interior tools that show you six options that are secretly one option.
Remember this
- A GAN is a contest, and a contest teaches only while both sides stay competitive.
- Mode collapse means the generator repeats itself instead of inventing.
- Loss curves are a poor progress report. Look at the output.
What to learn next
- StyleGAN and latent space — what a stabilised GAN can do once it trains.
- GANs — the architecture from the beginning, if this page moved too fast.
- Diffusion models — the family that replaced GANs for text-to-image.
Developer — Code and libraries.
Setup
pip install torch==2.5.1Everything here runs on a CPU in well under a minute.
Two failures, produced on purpose
The 2014 paper writes the generator's objective as minimising log(1 - D(G(z))). Almost no implementation uses it, and the reason is visible in three lines of arithmetic.
import torch, torch.nn as nn
torch.manual_seed(0)
# ---------- part 1: why the textbook generator loss stops giving gradients ----------
print("D(G(z)) saturating d/dlogit non-saturating d/dlogit")
for p in [0.5, 0.2, 0.05, 0.01, 0.001]:
logit = torch.tensor([[torch.logit(torch.tensor(p)).item()]], requires_grad=True)
# saturating: minimise log(1 - D(G(z)))
torch.log1p(-torch.sigmoid(logit)).backward()
sat = logit.grad.item(); logit.grad = None
# non-saturating: maximise log D(G(z))
(-torch.log(torch.sigmoid(logit))).backward()
non = logit.grad.item()
print(f"{p:7.3f} {sat:19.5f} {non:22.5f}")
# ---------- part 2: mode collapse you can watch ----------
MODES = torch.tensor([-2.0, 0.0, 2.0])
def real(n): return MODES[torch.randint(0, 3, (n,))].unsqueeze(1) + 0.08*torch.randn(n,1)
def mlp(i,o,h=32): return nn.Sequential(nn.Linear(i,h), nn.ReLU(), nn.Linear(h,h), nn.ReLU(), nn.Linear(h,o))
bce = nn.BCEWithLogitsLoss()
def train(d_lr, steps=800, bs=128):
torch.manual_seed(1)
G, D = mlp(1,1), mlp(1,1)
oG = torch.optim.Adam(G.parameters(), lr=1e-3, betas=(0.5,0.9))
oD = torch.optim.Adam(D.parameters(), lr=d_lr, betas=(0.5,0.9))
for _ in range(steps):
x, z = real(bs), torch.randn(bs,1)
oD.zero_grad()
(bce(D(x), torch.ones(bs,1)) + bce(D(G(z).detach()), torch.zeros(bs,1))).backward(); oD.step()
oG.zero_grad(); bce(D(G(torch.randn(bs,1))), torch.ones(bs,1)).backward(); oG.step()
with torch.no_grad(): f = G(torch.randn(4000,1))
hit = (f - MODES).abs().min(1).values < 0.3 # landed on some real mode
which = (f - MODES).abs().argmin(1)[hit]
counts = [int((which==k).sum()) for k in range(3)]
return counts, hit.float().mean().item(), f.std().item()
print("\nD learning rate samples per mode (-2, 0, +2) on a mode spread")
for lr in (1e-4, 1e-3, 1e-2):
c, hit, sd = train(lr)
print(f"{lr:<17.0e} {str(c):<30s} {hit:8.1%} {sd:8.3f}")
print(f"\n(real data spread for comparison: {real(4000).std().item():.3f})")D(G(z)) saturating d/dlogit non-saturating d/dlogit 0.500 -0.50000 -0.50000 0.200 -0.20000 -0.80000 0.050 -0.05000 -0.95000 0.010 -0.01000 -0.99000 0.001 -0.00100 -0.99900 D learning rate samples per mode (-2, 0, +2) on a mode spread 1e-04 [0, 0, 0] 0.0% 0.585 1e-03 [1469, 516, 1141] 78.1% 1.672 1e-02 [1626, 166, 1757] 88.7% 1.890 (real data spread for comparison: 1.624)
Part one is exact arithmetic and reproduces anywhere. Part two trains real networks, so the precise counts depend on your PyTorch build and platform. The pattern below is what reproduces.
Reading the output
The saturating loss dies exactly when you need it. When the discriminator is sure a sample is fake, D(G(z)) sits near zero. The saturating gradient there is -0.001. The non-saturating gradient is -0.999, a thousand times larger.
That is the whole argument. Early in training the generator is terrible, so D(G(z)) is tiny, so the textbook loss delivers almost nothing. The non-saturating form pushes hardest at exactly that moment.
A weak discriminator is not a safe discriminator. At a learning rate of 1e-4 the generator put nothing on any mode. Its spread was 0.585 against the real 1.624. It collapsed to a blob near zero, because the judge could not tell it that this was wrong.
A strong discriminator drops modes. At 1e-2 the middle mode got 166 samples against 1626 and 1757 for the outer two. Coverage is falling apart while the "on a mode" score climbs. That is mode collapse before it becomes total.
Only the middle setting is healthy. All three modes present, spread 1.672 against the real 1.624.
The fixes people actually use
Non-saturating loss. Write bce(D(G(z)), ones), never -bce(D(G(z)), zeros). One line, and not optional.
Two-timescale updates. Give the discriminator and generator different learning rates, usually a faster discriminator. Heusel et al. (2017) showed this converges where equal rates do not.
Zero-centred gradient penalties. Penalise the discriminator's input gradient on real data (R1) and on fake data (R2). This is the most reliable stabiliser in modern GAN code.
Spectral normalisation. Divide each discriminator weight matrix by its largest singular value. That caps how fast the discriminator can change, keeping the contest alive.
betas=(0.5, 0.9) on Adam. The default first beta of 0.9 is too much momentum in a game where the opponent moves. Nearly all GAN code lowers it.
Common mistakes
Reading loss curves as progress. Two networks in a fair fight produce roughly flat, noisy losses forever. A generator loss that plunges usually means the discriminator stopped trying.
Calling .detach() in the wrong place. When training the discriminator you must detach the generator's output. Forget it and the discriminator's loss flows back into the generator and pushes it the wrong way.
Forgetting model.eval() before sampling. Batch-norm in training mode makes each generated sample depend on the rest of the batch. Your single-image demo then looks nothing like your grid.
Judging by one grid of samples. Sixteen faces hide mode collapse well. Sample two thousand and measure spread, as the code above does.
Fixing instability by lowering both learning rates. That freezes the contest instead of balancing it. Change the ratio, not the scale.
Try it yourself
Add a fourth mode at +4.0 and rerun. Watch which mode is dropped first, and check whether it is the new one or the middle one. Then set betas=(0.9, 0.999) and see how much worse the same run becomes.
What to learn next
- StyleGAN and latent space — what a stabilised GAN can do once it trains.
- GANs — the architecture from the beginning, if this page moved too fast.
- Diffusion models — the family that replaced GANs for text-to-image.
Researcher — Mathematics and papers.
The original objective
Goodfellow et al. (2014) define the minimax game:
$$ \min_G \max_D \; V(D, G) = \mathbb{E}{x \sim p{\text{data}}}[\log D(x)] + \mathbb{E}_{z \sim p_z}[\log (1 - D(G(z)))] $$
$D$ is the discriminator, mapping an image to a probability. $G$ is the generator, mapping a noise vector $z$ to an image. $p_{\text{data}}$ is the data distribution, and $p_z$ the prior on $z$, conventionally $\mathcal{N}(0, I)$.
For a fixed $G$ the optimal discriminator is:
$$ D^*(x) = \frac{p_{\text{data}}(x)}{p_{\text{data}}(x) + p_g(x)} $$
where $p_g$ is the generator's implied distribution. Substituting it back gives:
$$ V(D^*, G) = 2 \cdot \mathrm{JSD}(p_{\text{data}} \,|\, p_g) - 2\log 2 $$
$\mathrm{JSD}$ is the Jensen-Shannon divergence. At the optimum the generator is minimising JSD.
Why that result is the problem, not the solution
If $p_{\text{data}}$ and $p_g$ have disjoint supports, $\mathrm{JSD} = \log 2$ everywhere and its gradient vanishes. Natural images lie on a low-dimensional manifold in pixel space. So early in training the supports really are close to disjoint. See Arjovsky and Bottou (2017), Towards Principled Methods for Training GANs.
The saturating generator loss $\log(1 - D(G(z)))$ therefore has vanishing gradient exactly when $D$ is confident. The non-saturating alternative $-\log D(G(z))$ shares the fixed point but has a gradient that grows as $D(G(z)) \to 0$. Arjovsky and Bottou show it is equivalent to minimising $\mathrm{KL}(p_g | p_{\text{data}}) - 2\,\mathrm{JSD}(p_{\text{data}} | p_g)$. The second term carries the wrong sign, and rewards mode-seeking behaviour. Non-saturating loss trades vanishing gradients for a bias toward mode collapse.
Convergence is not guaranteed by the objective
The game is not a convex-concave saddle-point problem in parameter space. Mescheder et al. (2018) ask Which Training Methods for GANs do actually Converge?. They show unregularised GAN training need not converge locally, even on an absolutely continuous data distribution. Eigenvalues of the Jacobian at equilibrium can have zero real part, producing orbits rather than convergence.
Their fix is the R1 penalty, a zero-centred gradient penalty applied on real data only:
$$ R_1(\psi) = \frac{\gamma}{2}\, \mathbb{E}{x \sim p{\text{data}}}!\left[|\nabla_x D_\psi(x)|^2\right] $$
$\psi$ are the discriminator parameters and $\gamma$ the penalty weight. They prove local convergence for the regularised game.
The main stabilisers, and what each buys
| Method | Mechanism | Reference |
|---|---|---|
| Non-saturating loss | Rescales the generator gradient | Goodfellow et al., 2014 |
| Minibatch discrimination | Lets $D$ see batch statistics, so repeats become detectable | Salimans et al., 2016 |
| WGAN | Replaces JSD with the Wasserstein-1 distance | Arjovsky et al., 2017 |
| WGAN-GP | Enforces the Lipschitz constraint by a gradient penalty | Gulrajani et al., 2017 |
| Spectral normalisation | Divides each weight by its spectral norm | Miyato et al., 2018 |
| TTUR | Separate learning rates for $D$ and $G$ | Heusel et al., 2017 |
| R1 / R2 | Zero-centred gradient penalty, proven local convergence | Mescheder et al., 2018 |
| ADA | Differentiable augmentation with an adaptive rate | Karras et al., 2020 |
WGAN has the strongest theoretical motivation. The Wasserstein-1 distance stays continuous and yields useful gradients even for disjoint supports:
$$ W_1(p_{\text{data}}, p_g) = \sup_{|f|L \le 1} \; \mathbb{E}{x \sim p_{\text{data}}}[f(x)] - \mathbb{E}_{x \sim p_g}[f(x)] $$
The supremum runs over 1-Lipschitz functions $f$. Enforcing that constraint is the entire practical difficulty. Weight clipping (the original proposal) distorts the critic; the gradient penalty of Gulrajani et al. instead penalises $(|\nabla_{\hat{x}} f(\hat{x})|_2 - 1)^2$ at points $\hat{x}$ interpolated between real and fake samples.
Mode collapse, precisely
Mode collapse is the case where $p_g$ concentrates on a strict subset of the support of $p_{\text{data}}$. It does not show up in the value function. A generator covering one mode perfectly can look fine to a discriminator that has collapsed to the same region.
Measurement needs a separate instrument. On synthetic mixtures, count recovered modes and reverse KL, as Metz et al. (2017) do for unrolled GANs. On real images the standard pair is:
- FID (Heusel et al., 2017): Fréchet distance between Gaussians fitted to Inception features. Sensitive to fidelity and diversity together, which is also its weakness, and biased by sample count.
- Precision and recall for generative models, from Sajjadi et al. (2018) and Kynkäänniemi et al. (2019). This separates fidelity (precision) from coverage (recall). Mode collapse appears as high precision with low recall, which FID alone cannot reveal.
Report both. A paper reporting only FID has not shown that its model covers the data.
Where GAN training stands now
Diffusion models displaced GANs for open-domain text-to-image generation, but the stability story did not end there.
Huang et al. (2024) derive a regularised relativistic loss carrying both R1 and R2 penalties. They prove local convergence for it. That lets them discard the accumulated bag of tricks. Their paper is The GAN is dead; long live the GAN! A Modern GAN Baseline, NeurIPS 2024. Their R3GAN uses a plain ResNet-and-ConvNeXt style backbone, with no StyleGAN-specific machinery. It surpasses StyleGAN2 on FFHQ, CIFAR, ImageNet and Stacked MNIST.
The claim worth taking from that work is narrow and important. Much of the historic instability was an artefact of a badly conditioned objective. It was not inherent to adversarial training.
GANs remain the practical choice wherever a single forward pass matters. They still dominate real-time super-resolution and on-device image enhancement.
Papers
- Goodfellow et al., Generative Adversarial Networks, 2014 — arxiv.org/abs/1406.2661
- Arjovsky and Bottou, Towards Principled Methods for Training GANs, 2017 — arxiv.org/abs/1701.04862
- Arjovsky et al., Wasserstein GAN, 2017 — arxiv.org/abs/1701.07875
- Gulrajani et al., Improved Training of Wasserstein GANs, 2017 — arxiv.org/abs/1704.00028
- Heusel et al., GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium, 2017 — arxiv.org/abs/1706.08500
- Miyato et al., Spectral Normalization for Generative Adversarial Networks, 2018 — arxiv.org/abs/1802.05957
- Mescheder et al., Which Training Methods for GANs do actually Converge?, 2018 — arxiv.org/abs/1801.04406
- Kynkäänniemi et al., Improved Precision and Recall Metric for Assessing Generative Models, 2019 — arxiv.org/abs/1904.06991
- Huang et al., The GAN is dead; long live the GAN! A Modern GAN Baseline, NeurIPS 2024 — arxiv.org/abs/2501.05441
What to learn next
- StyleGAN and latent space — what a stabilised GAN can do once it trains.
- GANs — the architecture from the beginning, if this page moved too fast.
- Diffusion models — the family that replaced GANs for text-to-image.