Skip to content

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/swd
score_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.

examples/gan/optimizer/dual_optimizer.py
"""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:

examples/gan/trainer/gan_trainer.py
"""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.