"""A 60-256-256-10 MLP learned from scratch, with one training step in a CUDA graph.

Ported from the repository's mlpg-k1-w256-s200-b512 cutoff experiment, which
qualified at 95.07% MNIST accuracy and 61.1 ms in the popcorn3 A100 harness.
A warm-up builds the graph. Every invocation resets weights, Adam moments,
EMA weights, step count and random streams before learning the supplied draw.
No dataset examples or pretrained weights are included.

Official check from mnist-a100/:
    python run_modal.py /path/to/fast_mlp.py:fast_mlp --difficulty 1 --runs 3
"""

import math

import torch
import torch.nn.functional as F

torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True

K, WIDTH, CLASSES = 1, 256, 10
TARGET_STEPS, BATCH = 200, 512
LR, WEIGHT_DECAY, WARMUP = 0.004, 1e-2, 0.1
DROPOUT, NOISE, SMOOTHING, EMA = 0.1, 0.3, 0.1, 0.98
SEED = 20260925
CACHE = {}


def lr_factor(step, total):
    warmup = int(round(WARMUP * total))
    if warmup and step < warmup:
        return (step + 1) / warmup
    return 0.5 * (1.0 + math.cos(math.pi * min(1.0, (step - warmup) / max(1, total - warmup))))


def forward(params, h, train):
    last = len(params) // 2 - 1
    for index in range(last + 1):
        h = torch.baddbmm(params[2 * index + 1], h, params[2 * index])
        if index < last:
            h = F.relu(h)
            if train:
                h = h * (torch.rand_like(h) >= DROPOUT) / (1.0 - DROPOUT)
    return h


class Trainer:
    def __init__(self, n, d, device):
        self.n, self.device = n, device
        self.fans = [d, WIDTH, WIDTH]
        self.x = torch.zeros(n, d, device=device)
        self.targets = torch.zeros(n, CLASSES, device=device)
        self.params = []
        for fan_in, fan_out in ((d, WIDTH), (WIDTH, WIDTH), (WIDTH, CLASSES)):
            self.params.append(torch.zeros(K, fan_in, fan_out, device=device, requires_grad=True))
            self.params.append(torch.zeros(K, 1, fan_out, device=device, requires_grad=True))
        self.lr = torch.tensor(LR, device=device)
        self.optimizer = torch.optim.AdamW(self.params, lr=self.lr, weight_decay=WEIGHT_DECAY,
                                           betas=(0.9, 0.999), eps=1e-8, fused=True, capturable=True)
        self.shadow = [torch.zeros_like(p) for p in self.params]
        self.schedule = torch.tensor([LR * lr_factor(s, TARGET_STEPS) for s in range(TARGET_STEPS)] + [0.0],
                                     device=device)
        self.counter = torch.zeros(1, dtype=torch.long, device=device)
        self.initialise()
        side = torch.cuda.Stream()
        side.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(side):
            for _ in range(3):
                self.optimizer.zero_grad(set_to_none=True)
                self.step()
        torch.cuda.current_stream().wait_stream(side)
        self.graph = torch.cuda.CUDAGraph()
        self.optimizer.zero_grad(set_to_none=True)
        with torch.cuda.graph(self.graph):
            self.step()

    def step(self):
        index = torch.randint(0, self.n, (K, BATCH), device=self.device)
        xb = self.x[index] + NOISE * torch.randn(K, BATCH, self.x.shape[1], device=self.device)
        self.lr.copy_(self.schedule.index_select(0, self.counter).squeeze(0))
        logits = forward(self.params, xb, True)
        loss = -(self.targets[index] * F.log_softmax(logits, dim=-1)).sum(-1).mean() * K
        loss.backward()
        self.optimizer.step()
        with torch.no_grad():
            torch._foreach_mul_(self.shadow, EMA)
            torch._foreach_add_(self.shadow, [p.detach() for p in self.params], alpha=1.0 - EMA)
            self.counter.add_(1)

    @torch.no_grad()
    def initialise(self):
        generator = torch.Generator(device=self.device)
        generator.manual_seed(SEED)
        for index, fan_in in enumerate(self.fans):
            bound = 1.0 / math.sqrt(fan_in)
            for p in self.params[2 * index:2 * index + 2]:
                p.copy_((torch.rand(p.shape, generator=generator, device=self.device) * 2 - 1) * bound)
        for s, p in zip(self.shadow, self.params):
            s.copy_(p)
        for p in self.params:
            state = self.optimizer.state.get(p)
            if state:
                state["exp_avg"].zero_()
                state["exp_avg_sq"].zero_()
                state["step"].zero_()
        self.counter.zero_()
        torch.cuda.manual_seed(SEED + 1)


def fast_mlp(train_x, train_y, test_x):
    device = train_x.device
    x = train_x.reshape(train_x.shape[0], -1).float()
    q = test_x.reshape(test_x.shape[0], -1).float()
    mean, std = x.mean(0), x.std(0).clamp_min(1e-3)
    key = (tuple(x.shape), str(device))
    if key not in CACHE:
        CACHE.clear()
        CACHE[key] = Trainer(x.shape[0], x.shape[1], device)
    trainer = CACHE[key]
    with torch.no_grad():
        trainer.x.copy_((x - mean) / std)
        trainer.targets.copy_(F.one_hot(train_y.long(), CLASSES).float() * (1.0 - SMOOTHING) + SMOOTHING / CLASSES)
    trainer.initialise()
    start = [p.detach().clone() for p in trainer.params]
    for _ in range(TARGET_STEPS):
        trainer.graph.replay()
    with torch.no_grad():
        bias = EMA ** TARGET_STEPS
        averaged = [(s - bias * s0) / (1.0 - bias) for s, s0 in zip(trainer.shadow, start)]
        q = (q - mean) / std
        probabilities = torch.zeros(q.shape[0], CLASSES, device=device)
        for begin in range(0, q.shape[0], 2000):
            chunk = q[begin:begin + 2000].unsqueeze(0).expand(K, -1, -1).contiguous()
            probabilities[begin:begin + 2000] = F.softmax(forward(averaged, chunk, False), dim=-1).sum(0)
    return probabilities.argmax(1)
