#!POPCORN leaderboard mnist-a100-2
#!POPCORN gpu A100
#!POPCORN function custom_kernel
# Submission by @yaroslavvb, based on the 16-member MLP ensemble by @yaroslavvb.
# Faster variant: 8 ensemble members; all other hyperparameters unchanged.
"""mlp-k8-w1024-s400-b512: 8 batched 60-1024-1024-10 MLPs, 400 steps of 512, averaged.

Generated by mnist/experiments/release-cutoffs-20260925/mlp_timing/mlp_family.py
from popcorn3/submissions/mlp_ensemble.py: dropout 0.1, input noise 0.3, label
smoothing 0.1, AdamW (weight decay 0.01, learning rate 2e-3 x sqrt(batch/128)),
10% warm-up then cosine to zero over exactly TARGET_STEPS steps, EMA weights
(decay 1 - 4/TARGET_STEPS, bias-corrected), members' softmax averaged.
Re-initialised from fixed seeds on every call; nothing is carried between calls.
"""

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 = 8, 1024, 10
TARGET_STEPS, BATCH = 400, 512
LR, WEIGHT_DECAY, WARMUP = 0.004, 1e-2, 0.1
DROPOUT, NOISE, SMOOTHING, EMA = 0.1, 0.3, 0.1, 0.99
SEED = 20260925


def init_stack(fan_in, fan_out, generator, device):
    bound = 1.0 / math.sqrt(fan_in)
    weight = (torch.rand(K, fan_in, fan_out, generator=generator, device=device) * 2 - 1) * bound
    bias = (torch.rand(K, 1, fan_out, generator=generator, device=device) * 2 - 1) * bound
    return weight.requires_grad_(), bias.requires_grad_()


def forward(params, h, generator=None):
    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 generator is not None:
                keep = torch.rand(h.shape, device=h.device, generator=generator) >= DROPOUT
                h = h * keep / (1.0 - DROPOUT)
    return h


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 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)
    x, q = (x - mean) / std, (q - mean) / std
    n, width = x.shape
    y = train_y.long()

    generator = torch.Generator(device=device)
    generator.manual_seed(SEED)
    params = []
    for fan_in, fan_out in ((width, WIDTH), (WIDTH, WIDTH), (WIDTH, CLASSES)):
        params.extend(init_stack(fan_in, fan_out, generator, device))
    optimizer = torch.optim.AdamW(params, lr=LR, weight_decay=WEIGHT_DECAY, betas=(0.9, 0.999),
                                  eps=1e-8, fused=device.type == "cuda")
    shadow = [p.detach().clone() for p in params]
    start = [p.detach().clone() for p in params]
    targets = F.one_hot(y, CLASSES).float() * (1.0 - SMOOTHING) + SMOOTHING / CLASSES

    total, step = TARGET_STEPS, 0
    while step < total:
        order = torch.rand(K, n, device=device, generator=generator).argsort(1)
        for begin in range(0, n, BATCH):
            if step == total:
                break
            index = order[:, begin:begin + BATCH]
            xb = x[index] + NOISE * torch.randn(K, index.shape[1], width, device=device, generator=generator)
            for group in optimizer.param_groups:
                group["lr"] = LR * lr_factor(step, total)
            logits = forward(params, xb, generator)
            loss = -(targets[index] * F.log_softmax(logits, dim=-1)).sum(-1).mean() * K
            optimizer.zero_grad(set_to_none=True)
            loss.backward()
            optimizer.step()
            with torch.no_grad():
                torch._foreach_mul_(shadow, EMA)
                torch._foreach_add_(shadow, [p.detach() for p in params], alpha=1.0 - EMA)
            step += 1

    with torch.no_grad():
        bias = EMA ** step
        averaged = [(s - bias * s0) / (1.0 - bias) for s, s0 in zip(shadow, start)]
        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), dim=-1).sum(0)
    return probabilities.argmax(1)


def custom_kernel(train_x, train_y, test_x):
    return mlp(train_x, train_y, test_x)
