#!POPCORN leaderboard mnist-medium-3p40pct
#!POPCORN gpu A100

"""mlpg-k1-w256-s200-b512: 1 batched 60-256-256-10 MLPs, 200 steps of 512, one step captured in a CUDA graph.

Generated by mnist/experiments/release-cutoffs-20260925/mlp_timing/mlp_family.py. The
same networks, regularisation, optimiser, schedule and EMA as the eager member of the
family, with two differences: every step samples its batch with replacement (one
randint inside the graph) instead of walking a per-epoch shuffle, and the whole step --
forward, backward, fused AdamW, EMA -- is captured once in a CUDA graph and replayed.
The capture happens on the first call (the harness's untimed warm-up). Every call then
re-initialises the weights, optimiser moments, EMA, step counter and random streams in
place from fixed seeds, so nothing learned is carried between calls; only the graph is.
"""

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 custom_kernel(data):
    train_x, train_y, test_x = data
    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)
