"""A compact, runnable MNIST DDPM used by the Diffusion Field Notes tutorial.

Smoke test:
  python mini-ddpm-mnist.py --epochs 1 --limit-batches 20 --steps 200

This verifies the pipeline; it is not enough training for good samples.
"""
from __future__ import annotations

import argparse
import math
import random
from pathlib import Path

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms, utils


def seed_everything(seed: int) -> None:
    random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


def extract(values: torch.Tensor, t: torch.Tensor, shape: torch.Size) -> torch.Tensor:
    return values.gather(0, t).reshape(t.shape[0], *((1,) * (len(shape) - 1)))


class Diffusion(nn.Module):
    def __init__(self, steps: int = 1000) -> None:
        super().__init__()
        betas = torch.linspace(1e-4, 2e-2, steps)
        alphas = 1.0 - betas
        alpha_bars = alphas.cumprod(0)
        alpha_bars_prev = F.pad(alpha_bars[:-1], (1, 0), value=1.0)
        self.steps = steps
        self.register_buffer("betas", betas)
        self.register_buffer("alphas", alphas)
        self.register_buffer("alpha_bars", alpha_bars)
        self.register_buffer(
            "posterior_variance",
            betas * (1.0 - alpha_bars_prev) / (1.0 - alpha_bars),
        )

    def q_sample(
        self, x0: torch.Tensor, t: torch.Tensor, noise: torch.Tensor | None = None
    ) -> tuple[torch.Tensor, torch.Tensor]:
        noise = torch.randn_like(x0) if noise is None else noise
        alpha_bar = extract(self.alpha_bars, t, x0.shape)
        xt = alpha_bar.sqrt() * x0 + (1.0 - alpha_bar).sqrt() * noise
        return xt, noise


class TimeEmbedding(nn.Module):
    def __init__(self, dim: int) -> None:
        super().__init__()
        self.dim = dim

    def forward(self, t: torch.Tensor) -> torch.Tensor:
        half = self.dim // 2
        frequencies = torch.exp(
            -math.log(10000) * torch.arange(half, device=t.device) / max(half - 1, 1)
        )
        angles = t.float()[:, None] * frequencies[None, :]
        return torch.cat((angles.sin(), angles.cos()), dim=1)


class ResBlock(nn.Module):
    def __init__(self, cin: int, cout: int, time_dim: int) -> None:
        super().__init__()
        self.norm1 = nn.GroupNorm(8, cin)
        self.conv1 = nn.Conv2d(cin, cout, 3, padding=1)
        self.norm2 = nn.GroupNorm(8, cout)
        self.conv2 = nn.Conv2d(cout, cout, 3, padding=1)
        self.time = nn.Linear(time_dim, cout)
        self.skip = nn.Conv2d(cin, cout, 1) if cin != cout else nn.Identity()

    def forward(self, x: torch.Tensor, temb: torch.Tensor) -> torch.Tensor:
        h = self.conv1(F.silu(self.norm1(x)))
        h = h + self.time(F.silu(temb))[:, :, None, None]
        h = self.conv2(F.silu(self.norm2(h)))
        return h + self.skip(x)


class TinyUNet(nn.Module):
    def __init__(self, base: int = 32, time_dim: int = 128) -> None:
        super().__init__()
        self.time_mlp = nn.Sequential(
            TimeEmbedding(base), nn.Linear(base, time_dim), nn.SiLU(),
            nn.Linear(time_dim, time_dim),
        )
        self.input = nn.Conv2d(1, base, 3, padding=1)
        self.down1 = ResBlock(base, base * 2, time_dim)
        self.pool1 = nn.Conv2d(base * 2, base * 2, 4, stride=2, padding=1)
        self.down2 = ResBlock(base * 2, base * 4, time_dim)
        self.pool2 = nn.Conv2d(base * 4, base * 4, 4, stride=2, padding=1)
        self.middle = ResBlock(base * 4, base * 4, time_dim)
        self.up2 = nn.ConvTranspose2d(base * 4, base * 4, 4, stride=2, padding=1)
        self.decode2 = ResBlock(base * 8, base * 2, time_dim)
        self.up1 = nn.ConvTranspose2d(base * 2, base * 2, 4, stride=2, padding=1)
        self.decode1 = ResBlock(base * 4, base, time_dim)
        self.output = nn.Sequential(
            nn.GroupNorm(8, base), nn.SiLU(), nn.Conv2d(base, 1, 3, padding=1)
        )

    def forward(self, x: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
        temb = self.time_mlp(t)
        h0 = self.input(x)                     # [B, 32, 28, 28]
        h1 = self.down1(h0, temb)              # [B, 64, 28, 28]
        h2 = self.down2(self.pool1(h1), temb)  # [B,128, 14, 14]
        h = self.middle(self.pool2(h2), temb)  # [B,128,  7,  7]
        h = self.decode2(torch.cat((self.up2(h), h2), dim=1), temb)
        h = self.decode1(torch.cat((self.up1(h), h1), dim=1), temb)
        return self.output(h)                  # unconstrained epsilon prediction


def diffusion_loss(model: nn.Module, diffusion: Diffusion, x0: torch.Tensor) -> torch.Tensor:
    t = torch.randint(0, diffusion.steps, (x0.shape[0],), device=x0.device)
    xt, noise = diffusion.q_sample(x0, t)
    return F.mse_loss(model(xt, t), noise)


@torch.no_grad()
def p_sample(
    model: nn.Module, diffusion: Diffusion, xt: torch.Tensor, t: torch.Tensor
) -> torch.Tensor:
    beta = extract(diffusion.betas, t, xt.shape)
    alpha = extract(diffusion.alphas, t, xt.shape)
    alpha_bar = extract(diffusion.alpha_bars, t, xt.shape)
    eps = model(xt, t)
    mean = (xt - beta / (1.0 - alpha_bar).sqrt() * eps) / alpha.sqrt()
    variance = extract(diffusion.posterior_variance, t, xt.shape)
    noise = torch.randn_like(xt)
    nonzero = (t != 0).float().reshape(-1, 1, 1, 1)
    return mean + nonzero * variance.clamp_min(1e-20).sqrt() * noise


@torch.no_grad()
def sample(
    model: nn.Module, diffusion: Diffusion, count: int, device: torch.device
) -> torch.Tensor:
    model.eval()
    x = torch.randn(count, 1, 28, 28, device=device)
    for step in reversed(range(diffusion.steps)):
        t = torch.full((count,), step, device=device, dtype=torch.long)
        x = p_sample(model, diffusion, x, t)
    return x.clamp(-1, 1)


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--epochs", type=int, default=10)
    parser.add_argument("--batch-size", type=int, default=128)
    parser.add_argument("--steps", type=int, default=1000)
    parser.add_argument("--lr", type=float, default=2e-4)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--limit-batches", type=int, default=0)
    parser.add_argument("--output", type=Path, default=Path("runs/mini-ddpm"))
    args = parser.parse_args()

    seed_everything(args.seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    args.output.mkdir(parents=True, exist_ok=True)
    data = datasets.MNIST(
        "data", train=True, download=True,
        transform=transforms.Compose((transforms.ToTensor(), transforms.Lambda(lambda x: x * 2 - 1))),
    )
    loader = DataLoader(data, batch_size=args.batch_size, shuffle=True, num_workers=0)
    model = TinyUNet().to(device)
    diffusion = Diffusion(args.steps).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr)

    for epoch in range(args.epochs):
        model.train()
        running = 0.0
        for batch_index, (x0, _) in enumerate(loader):
            x0 = x0.to(device)
            optimizer.zero_grad(set_to_none=True)
            loss = diffusion_loss(model, diffusion, x0)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            running += loss.item()
            if args.limit_batches and batch_index + 1 >= args.limit_batches:
                break
        batches = min(len(loader), args.limit_batches or len(loader))
        print(f"epoch={epoch + 1} mse={running / batches:.5f}")
        torch.save(
            {"model": model.state_dict(), "optimizer": optimizer.state_dict(),
             "epoch": epoch, "steps": args.steps, "seed": args.seed},
            args.output / "latest.pt",
        )
        images = (sample(model, diffusion, 16, device) + 1) / 2
        utils.save_image(images, args.output / f"samples-{epoch + 1:03d}.png", nrow=4)


if __name__ == "__main__":
    main()
