Adversarial training has two optimizers and two different gradient destinations; a mistaken detach or frozen-state rule can train the wrong network.
GAN discriminator and generator update boundaries
Give each network one objective
The generator turns a noise vector into a candidate sample. The discriminator scores whether a sample resembles the training population. For a simple logistic setup, train the discriminator on real inputs as real and generator outputs as fake; train the generator so the discriminator calls its output real. The discriminator should return logits, not a probability already passed through sigmoid, when using a logits-based binary cross entropy. Check the input scale and generator output activation together. The logits contract prevents a silent extra sigmoid.
Cut the graph during the discriminator step
When optimizing the discriminator, fake samples are training inputs for that network. Detach them from the generator graph so the discriminator backward pass does not spend memory producing generator gradients. Zero the discriminator gradients before its backward pass, then step only its optimizer. Real and fake batches should have compatible shapes and preprocessing. Do not cache old fake samples without documenting that policy; a stale fake distribution changes what the discriminator is asked to separate. The code uses a single tiny batch to expose the update boundary.
Preserve input gradients during the generator step
For the generator update, compute fresh fake samples and pass them through the discriminator without detaching. The discriminator weights should not be updated, but the derivative of its score with respect to its input is needed to train the generator. Temporarily marking discriminator parameters nontrainable can save their gradient storage while still allowing gradients to flow into generated samples. Step only the generator optimizer, then restore discriminator parameter settings before the next discriminator update. A no-grad block around discriminator evaluation would sever the generator signal.
Track more than a loss curve
A discriminator loss near a familiar value does not prove realistic output or balanced competition. Save fixed-noise outputs across checkpoints; measure generated diversity, nearest training neighbors and downstream utility. Watch gradient norms, score distributions on real and fake inputs, and whether one network becomes an easy shortcut detector for image borders or normalization. A generator can improve its own objective while producing a tiny set of repeated samples. Coverage auditing makes that failure measurable.
Choose the adversarial variant deliberately
The example uses a non-saturating logistic generator loss because it gives a useful signal when the discriminator is initially confident. Other adversarial objectives have different score meanings and constraints; changing to a critic objective is not a one-line label swap. Keep optimizer settings, discriminator updates per generator update, batch order and model revision in the experiment manifest. Compare candidate checkpoints against a held-out real set and an applied task, not against the discriminator that co-trained with them. The project gives a bounded use case.
Implementation
import torch
from torch import nn
from torch.nn import functional as functional
torch.manual_seed(47)
generator = nn.Sequential(nn.Linear(7, 24), nn.ReLU(), nn.Linear(24, 16), nn.Tanh())
discriminator = nn.Sequential(nn.Linear(16, 18), nn.ReLU(), nn.Linear(18, 1))
generator_optimizer = torch.optim.AdamW(generator.parameters(), lr=0.0004)
discriminator_optimizer = torch.optim.AdamW(discriminator.parameters(), lr=0.0003)
real_texture = torch.rand(5, 16) * 2 - 1
latent_noise = torch.randn(5, 7)
discriminator_optimizer.zero_grad(set_to_none=True)
detached_fake = generator(latent_noise).detach()
real_scores = discriminator(real_texture)
fake_scores = discriminator(detached_fake)
discriminator_loss = (functional.binary_cross_entropy_with_logits(
real_scores, torch.ones_like(real_scores)) +
functional.binary_cross_entropy_with_logits(
fake_scores, torch.zeros_like(fake_scores))) / 2
discriminator_loss.backward()
discriminator_optimizer.step()
for parameter in discriminator.parameters():
parameter.requires_grad_(False)
discriminator_optimizer.zero_grad(set_to_none=True)
generator_optimizer.zero_grad(set_to_none=True)
fresh_fake = generator(torch.randn(5, 7))
generator_scores = discriminator(fresh_fake)
generator_loss = functional.binary_cross_entropy_with_logits(
generator_scores, torch.ones_like(generator_scores))
generator_loss.backward()
generator_optimizer.step()
for parameter in discriminator.parameters():
parameter.requires_grad_(True)
assert torch.isfinite(generator_loss + discriminator_loss)Performance and operating cost
Each paired update performs one generator forward for discriminator training and another generator forward with a backward pass for generator training, plus discriminator forwards on real and fake inputs. The exact ratio depends on the chosen update schedule. Detaching the first fake batch avoids an unneeded generator graph, while freezing discriminator parameters during the second pass reduces gradient storage without removing its input derivative. Adversarial training can require many checkpoints and evaluations; one loss number is cheap but inadequate evidence of output quality.
Common Mistakes
- Do not detach generator output during its own update.
- Do not wrap the discriminator in no-grad during the generator update.
- Do not interpret discriminator loss alone as a quality or coverage score.
