GAN (two optimizers)¶
A non-saturating GAN on a correlated 2-D Gaussian. CPU, no downloads, ~12 s.
gatle-ignite train --config=examples/gan/configs/gan_v0.py
python examples/gan/scripts/verify.py # the proof
This example exists to answer one question: can a framework with a single optimizer_name field
train a GAN? It can, and it needs no framework change, because optimizer_name never meant "pick a
torch optimizer". It means "name a module that returns the thing I will hold".
Nothing beyond the walk: dataset · model · checkpoints · logging · running
| step | what this example changes |
|---|---|
| prep_batch | overrides train_spec()overrides train_step() |
| loss | criterion_name = examples.gan.losses.gan_loss |
| metrics | val_metrics: moments (distribution), swd (distribution)score_name = valid/swdscore_factor = -1 |
| optimizer | optimizer_name = examples.gan.optimizer.dual_optimizer |
The two-optimizer trick¶
get_optimizer returns one object holding two real optimizers. The framework only ever asks it for
param_groups, state_dict, load_state_dict, so anything with those works, and it need not
subclass torch.optim.Optimizer.
"""Two optimizers behind the framework's one-optimizer contract.
Not a torch optimizer: it implements only the surface the framework touches.
"""
import torch
class GanOptimizers:
def __init__(self, gen_opt, disc_opt):
self.gen = gen_opt
self.disc = disc_opt
@property
def param_groups(self):
# G's groups first, so train/lr_0 is the generator.
return list(self.gen.param_groups) + list(self.disc.param_groups)
def state_dict(self):
return {"gen": self.gen.state_dict(), "disc": self.disc.state_dict()}
def load_state_dict(self, state):
self.gen.load_state_dict(state["gen"])
self.disc.load_state_dict(state["disc"])
def zero_grad(self, set_to_none=True):
self.gen.zero_grad(set_to_none=set_to_none)
self.disc.zero_grad(set_to_none=set_to_none)
def step(self, *args, **kwargs):
# Driven as one optimizer, this would step G on D's gradients: refuse, don't train wrong.
raise RuntimeError(
"GanOptimizers.step() called. A GAN's two optimizers must be stepped "
"separately by the trainer's two-phase train_step (examples/gan/trainer/gan_trainer.py); "
"stepping them together would apply the discriminator's update to the "
"generator. If you see this, something is driving this object as if it were a "
"single torch optimizer, most likely BaseTrainer.backward() or a scheduler."
)
def get_optimizer(model, gen_lr=2e-4, disc_lr=2e-4, betas=(0.5, 0.999), **kwargs):
"""betas=(0.5, 0.999): the default 0.9 first moment makes an adversarial update oscillate."""
# Under DDP `model.gen` raises AttributeError: reach submodules through `.module`.
model = getattr(model, "module", model)
# ml_collections hands tuples through as lists; torch wants a 2-tuple.
betas = tuple(betas)
return GanOptimizers(
torch.optim.Adam(model.gen.parameters(), lr=gen_lr, betas=betas, **kwargs),
torch.optim.Adam(model.disc.parameters(), lr=disc_lr, betas=betas, **kwargs),
)
The alternating update is an ordinary train_step override:
"""The GAN trainer: an alternating two-optimizer train_step. The default eval_step samples G."""
import torch
from gatle_ignite import BaseTrainer, to_device
class Trainer(BaseTrainer):
def prep_batch(self, batch, split="train", **kwargs):
"""Draw z here, not in the dataset, so its size matches a short final batch."""
(real,) = batch
z = torch.randn(real.shape[0], self._z_dim(), device=real.device)
return to_device({"model_input": {"z": z}, "targets": {"real": real}})
def train_step(self, engine, batch, split="train"):
"""Alternate: D first, so G's update is scored by the freshest critic."""
engine.state.batch = None
engine.state.output = None
self.model.train()
if self.accum_steps != 1:
# This replaces the default train_step's accumulation: refuse a field that does nothing.
raise NotImplementedError(
f"accum_steps={self.accum_steps} is not supported by the GAN trainer: "
f"this train_step replaces the framework's accumulation window. Use "
f"accum_steps=1 and a larger bs."
)
gen, disc = self._nets()
x = self.prep_batch(batch, split=split)
real, z = x["targets"]["real"], x["model_input"]["z"]
# ---- phase 1: the discriminator ----
self.optimizer.disc.zero_grad(set_to_none=True)
with self._autocast():
# .detach() is load-bearing: without it D's backward trains G to help its critic.
fake = gen(z).detach()
loss_d = self.criterion.disc_loss(disc(real), disc(fake))
loss_d.backward()
self.optimizer.disc.step()
# ---- phase 2: the generator ----
self.optimizer.gen.zero_grad(set_to_none=True)
with self._autocast():
fake = gen(z)
# Dirties D's .grad, harmlessly: phase 1 zeroes it, and D is not stepped here.
loss_g = self.criterion.gen_loss(disc(fake))
loss_g.backward()
self.optimizer.gen.step()
return {
"y_pred": {"fake": fake.detach()},
"target": x,
# No "loss": D + G falls when EITHER wins. Watch loss_d/loss_g; trust only valid/swd.
"losses": {
"loss_d": loss_d.detach(),
"loss_g": loss_g.detach(),
},
}
def train_spec(self):
"""No train/loss_avg, whose D + G falls when either wins; keep loss_d_avg and loss_g_avg."""
from dataclasses import replace
return replace(super().train_spec(), total_loss=False)
def _nets(self):
"""(generator, discriminator), through DDP's `.module` if present.
The DDP path is untested: it would also need find_unused_parameters or no_sync.
"""
model = getattr(self.model, "module", self.model)
return model.gen, model.disc
def _z_dim(self):
return getattr(self.model, "module", self.model).z_dim
def _autocast(self):
return torch.autocast(
device_type=self.device_type, dtype=self.dtype, enabled=self.autocast_enabled
)
Why the loss is not the evidence¶
A GAN's total loss can fall because either adversary is winning. It is not a progress signal, and
this example's train/loss_avg should be ignored. The evidence is a sliced-Wasserstein distance
between generated and real samples, which is what selects the best checkpoint (lower is better, hence
score_factor = -1).
verify.py scores three numbers through the framework's own inference mode:
FLOOR real vs real (sampling noise) swd = 0.0387
UNTRAINED generator at init swd = 1.6609
TRAINED generator, best checkpoint swd = 0.1087
improvement over untrained: 15.3x (required: >5.0x)
The floor is the number that makes this honest: real-vs-real scores 0.0387 rather than 0, so a
trained score of 0.1087 is "close to the data", not "suspiciously perfect". The untrained control
is what a no-op training loop would score, and verify.py exits non-zero against it, so the check
can fail, which is the only reason to trust it passing.