"""Original two-dimensional GAN training lesson, Goodfellow et al. (2014).

Tiny MLP generator/discriminator. BCE uses logits for stability. Discriminator
updates use detached generated samples; generator updates freeze discriminator
parameters while differentiating through it. No trained image generator.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F


class Generator(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden = nn.Linear(4, 8)
        self.relu = nn.ReLU()
        self.project = nn.Linear(8, 2)

    def forward(self, z):
        return self.project(self.relu(self.hidden(z)))


class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden = nn.Linear(2, 8)
        self.relu = nn.ReLU()
        self.logit = nn.Linear(8, 1)

    def forward(self, x):
        return self.logit(self.relu(self.hidden(x)))


class TinyGAN(nn.Module):
    def __init__(self):
        super().__init__()
        self.generator = Generator()
        self.discriminator = Discriminator()

    def forward(self, noise, real):
        fake = self.generator(noise)
        real_logits = self.discriminator(real)
        fake_logits = self.discriminator(fake)
        return torch.stack([real_logits, fake_logits])


def setup():
    torch.manual_seed(29)
    model = TinyGAN().double()
    noise = torch.randn(4, 4, dtype=torch.float64)
    real = torch.tensor(
        [[0.8, 0.9], [1.1, 1.0], [0.9, 1.2], [1.2, 0.8]], dtype=torch.float64
    )
    return model, noise, real


def discriminator_step(model, noise, real):
    optimizer = torch.optim.SGD(model.discriminator.parameters(), lr=0.1)
    optimizer.zero_grad()
    fake = model.generator(noise).detach()
    loss = (
        F.softplus(-model.discriminator(real)).mean()
        + F.softplus(model.discriminator(fake)).mean()
    )
    loss.backward()
    optimizer.step()
    return loss.item()


def generator_step(model, noise):
    for parameter in model.discriminator.parameters():
        parameter.requires_grad_(False)
    optimizer = torch.optim.SGD(model.generator.parameters(), lr=0.1)
    optimizer.zero_grad()
    # Practical non-saturating objective from the original paper.
    loss = F.softplus(-model.discriminator(model.generator(noise))).mean()
    loss.backward()
    optimizer.step()
    for parameter in model.discriminator.parameters():
        parameter.requires_grad_(True)
    return loss.item()


def flat(module):
    return torch.cat([p.detach().flatten() for p in module.parameters()])


def verify():
    model, noise, real = setup()
    g0 = flat(model.generator).clone()
    d0 = flat(model.discriminator).clone()
    discriminator_step(model, noise, real)
    torch.testing.assert_close(flat(model.generator), g0, rtol=0, atol=0)
    assert not torch.allclose(flat(model.discriminator), d0)
    assert all(p.grad is None for p in model.generator.parameters())
    d1 = flat(model.discriminator).clone()
    generator_step(model, noise)
    torch.testing.assert_close(flat(model.discriminator), d1, rtol=0, atol=0)
    assert not torch.allclose(flat(model.generator), g0)
    logits = torch.tensor([-6.0, -2.0, 2.0], dtype=torch.float64, requires_grad=True)
    minimax = -F.softplus(logits)
    nonsaturating = F.softplus(-logits)
    gm = torch.autograd.grad(minimax.sum(), logits, retain_graph=True)[0]
    gn = torch.autograd.grad(nonsaturating.sum(), logits)[0]
    torch.testing.assert_close(gm, -logits.sigmoid())
    torch.testing.assert_close(gn, logits.sigmoid() - 1)
    assert abs(gn[0]) > 100 * abs(gm[0])
    return [
        "A discriminator update changes only D; detached samples give the generator no gradient",
        "A generator update changes only G while retaining a gradient path through frozen D",
        "Stable logit losses match minimax and non-saturating analytic gradients",
        "The non-saturating loss gives a stronger gradient in the confident-fake example",
    ]


def experiment():
    model, noise, real = setup()
    g0 = flat(model.generator).clone()
    d0 = flat(model.discriminator).clone()
    cases = []
    for identity, label in [
        ("initial", "Before either update"),
        ("discriminator", "After one discriminator update"),
        ("generator", "After one generator update"),
    ]:
        if identity == "discriminator":
            discriminator_step(model, noise, real)
        if identity == "generator":
            generator_step(model, noise)
        with torch.no_grad():
            fake = model.generator(noise)
            dreal = model.discriminator(real)
            dfake = model.discriminator(fake)
            cases.append(
                {
                    "id": identity,
                    "label": label,
                    "target": "generator"
                    if identity == "generator"
                    else "discriminator",
                    "note": {
                        "initial": "Start from the same random initialization and fixed noise samples. The discriminator assigns probabilities to real and generated points. Two separate optimizers will update one network at a time.",
                        "discriminator": "Update D to raise scores on real points and lower scores on detached generated points. G receives no gradient, so its generated coordinates are unchanged. This is one discriminator step, not convergence.",
                        "generator": "Keep D fixed and update G using the non-saturating loss −log D(G(z)). Gradients pass through D to G, while D's parameters remain unchanged from the preceding step. Reported parameter deltas are cumulative from initialization.",
                    }[identity],
                    "scatter": [
                        {"label": "Real training points", "points": real.tolist()},
                        {
                            "label": "Generated points · fixed noise",
                            "points": fake.tolist(),
                        },
                    ],
                    "vectors": [
                        {
                            "label": "D(real) probabilities",
                            "values": dreal.sigmoid().flatten().tolist(),
                        },
                        {
                            "label": "D(G(z)) probabilities",
                            "values": dfake.sigmoid().flatten().tolist(),
                        },
                    ],
                    "metrics": [
                        {
                            "label": "Generator parameter change · norm since initialization",
                            "value": (flat(model.generator) - g0).norm().item(),
                        },
                        {
                            "label": "Discriminator parameter change · norm since initialization",
                            "value": (flat(model.discriminator) - d0).norm().item(),
                        },
                        {
                            "label": "Discriminator loss",
                            "value": (
                                F.softplus(-dreal).mean() + F.softplus(dfake).mean()
                            ).item(),
                        },
                        {
                            "label": "Non-saturating generator loss",
                            "value": F.softplus(-dfake).mean().item(),
                        },
                    ],
                }
            )
    logits = torch.tensor([-6.0, -2.0, 2.0], dtype=torch.float64)
    cases.append(
        {
            "id": "gradients",
            "label": "Why use the non-saturating generator loss?",
            "target": "generator",
            "scatterAxes": {
                "xLabel": "Discriminator logit",
                "yLabel": "Loss derivative",
                "equalScale": False,
            },
            "note": "The original minimax game minimizes log(1−D(G(z))) for G. When D confidently rejects a fake, this gradient is small. The original paper also recommends maximizing log D(G(z)), the non-saturating alternative used in these updates. These analytic derivatives are with respect to the discriminator logit, not a measured training outcome.",
            "scatter": [
                {
                    "label": "Minimax gradient · x = logit",
                    "points": list(zip(logits.tolist(), (-logits.sigmoid()).tolist())),
                },
                {
                    "label": "Non-saturating gradient · x = logit",
                    "points": list(
                        zip(logits.tolist(), (logits.sigmoid() - 1).tolist())
                    ),
                },
            ],
            "vectors": [
                {"label": "Discriminator logits", "values": logits.tolist()},
                {"label": "Minimax derivative", "values": (-logits.sigmoid()).tolist()},
                {
                    "label": "Non-saturating derivative",
                    "values": (logits.sigmoid() - 1).tolist(),
                },
            ],
            "metrics": [
                {
                    "label": "Gradient magnitude ratio at logit −6",
                    "value": ((1 - logits[0].sigmoid()) / logits[0].sigmoid()).item(),
                }
            ],
        }
    )
    return {
        "kind": "scatter",
        "title": "Update one network at a time.",
        "description": "Recorded two-dimensional GAN updates on four toy points. Choose the discriminator and generator stages to inspect which parameters and generated samples change. One alternating update does not establish a learned distribution or stable GAN training.",
        "controlLabel": "Training stage",
        "cases": cases,
    }
