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.

On this page 9
  1. The short answer
  2. The analogy
  3. The two players
  4. Why this is fragile
  5. The four ways it goes wrong
  6. How you would notice
  7. Where you have seen this
  8. Remember this
  9. What to learn next

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 has

Why 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

Developer — Code and libraries.

Setup

bash
pip install torch==2.5.1

Everything 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.

gan_failures.py
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})")
Output
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

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

MethodMechanismReference
Non-saturating lossRescales the generator gradientGoodfellow et al., 2014
Minibatch discriminationLets $D$ see batch statistics, so repeats become detectableSalimans et al., 2016
WGANReplaces JSD with the Wasserstein-1 distanceArjovsky et al., 2017
WGAN-GPEnforces the Lipschitz constraint by a gradient penaltyGulrajani et al., 2017
Spectral normalisationDivides each weight by its spectral normMiyato et al., 2018
TTURSeparate learning rates for $D$ and $G$Heusel et al., 2017
R1 / R2Zero-centred gradient penalty, proven local convergenceMescheder et al., 2018
ADADifferentiable augmentation with an adaptive rateKarras 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

What to learn next