{
  "app_id": "ap-3DM0j6iU8LjOkKSd3xD5os",
  "order": [
    "baseline",
    "ours",
    "ours",
    "baseline",
    "baseline",
    "ours"
  ],
  "methods": {
    "baseline": {
      "filename": "mlpg_k4_w256_s800_b512.py",
      "function": "mlp",
      "source": "\"\"\"mlpg-k4-w256-s800-b512: 4 batched 60-256-256-10 MLPs, 800 steps of 512, one step captured in a CUDA graph.\n\nGenerated by mnist/experiments/release-cutoffs-20260925/mlp_timing/mlp_family.py. The\nsame networks, regularisation, optimiser, schedule and EMA as the eager member of the\nfamily, with two differences: every step samples its batch with replacement (one\nrandint inside the graph) instead of walking a per-epoch shuffle, and the whole step --\nforward, backward, fused AdamW, EMA -- is captured once in a CUDA graph and replayed.\nThe capture happens on the first call (the harness's untimed warm-up). Every call then\nre-initialises the weights, optimiser moments, EMA, step counter and random streams in\nplace from fixed seeds, so nothing learned is carried between calls; only the graph is.\n\"\"\"\n\nimport math\n\nimport torch\nimport torch.nn.functional as F\n\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\n\nK, WIDTH, CLASSES = 4, 256, 10\nTARGET_STEPS, BATCH = 800, 512\nLR, WEIGHT_DECAY, WARMUP = 0.004, 1e-2, 0.1\nDROPOUT, NOISE, SMOOTHING, EMA = 0.1, 0.3, 0.1, 0.995\nSEED = 20260925\nCACHE = {}\n\n\ndef lr_factor(step, total):\n    warmup = int(round(WARMUP * total))\n    if warmup and step < warmup:\n        return (step + 1) / warmup\n    return 0.5 * (1.0 + math.cos(math.pi * min(1.0, (step - warmup) / max(1, total - warmup))))\n\n\ndef forward(params, h, train):\n    last = len(params) // 2 - 1\n    for index in range(last + 1):\n        h = torch.baddbmm(params[2 * index + 1], h, params[2 * index])\n        if index < last:\n            h = F.relu(h)\n            if train:\n                h = h * (torch.rand_like(h) >= DROPOUT) / (1.0 - DROPOUT)\n    return h\n\n\nclass Trainer:\n    def __init__(self, n, d, device):\n        self.n, self.device = n, device\n        self.fans = [d, WIDTH, WIDTH]\n        self.x = torch.zeros(n, d, device=device)\n        self.targets = torch.zeros(n, CLASSES, device=device)\n        self.params = []\n        for fan_in, fan_out in ((d, WIDTH), (WIDTH, WIDTH), (WIDTH, CLASSES)):\n            self.params.append(torch.zeros(K, fan_in, fan_out, device=device, requires_grad=True))\n            self.params.append(torch.zeros(K, 1, fan_out, device=device, requires_grad=True))\n        self.lr = torch.tensor(LR, device=device)\n        self.optimizer = torch.optim.AdamW(self.params, lr=self.lr, weight_decay=WEIGHT_DECAY,\n                                           betas=(0.9, 0.999), eps=1e-8, fused=True, capturable=True)\n        self.shadow = [torch.zeros_like(p) for p in self.params]\n        self.schedule = torch.tensor([LR * lr_factor(s, TARGET_STEPS) for s in range(TARGET_STEPS)] + [0.0],\n                                     device=device)\n        self.counter = torch.zeros(1, dtype=torch.long, device=device)\n        self.initialise()\n        side = torch.cuda.Stream()\n        side.wait_stream(torch.cuda.current_stream())\n        with torch.cuda.stream(side):\n            for _ in range(3):\n                self.optimizer.zero_grad(set_to_none=True)\n                self.step()\n        torch.cuda.current_stream().wait_stream(side)\n        self.graph = torch.cuda.CUDAGraph()\n        self.optimizer.zero_grad(set_to_none=True)\n        with torch.cuda.graph(self.graph):\n            self.step()\n\n    def step(self):\n        index = torch.randint(0, self.n, (K, BATCH), device=self.device)\n        xb = self.x[index] + NOISE * torch.randn(K, BATCH, self.x.shape[1], device=self.device)\n        self.lr.copy_(self.schedule.index_select(0, self.counter).squeeze(0))\n        logits = forward(self.params, xb, True)\n        loss = -(self.targets[index] * F.log_softmax(logits, dim=-1)).sum(-1).mean() * K\n        loss.backward()\n        self.optimizer.step()\n        with torch.no_grad():\n            torch._foreach_mul_(self.shadow, EMA)\n            torch._foreach_add_(self.shadow, [p.detach() for p in self.params], alpha=1.0 - EMA)\n            self.counter.add_(1)\n\n    @torch.no_grad()\n    def initialise(self):\n        generator = torch.Generator(device=self.device)\n        generator.manual_seed(SEED)\n        for index, fan_in in enumerate(self.fans):\n            bound = 1.0 / math.sqrt(fan_in)\n            for p in self.params[2 * index:2 * index + 2]:\n                p.copy_((torch.rand(p.shape, generator=generator, device=self.device) * 2 - 1) * bound)\n        for s, p in zip(self.shadow, self.params):\n            s.copy_(p)\n        for p in self.params:\n            state = self.optimizer.state.get(p)\n            if state:\n                state[\"exp_avg\"].zero_()\n                state[\"exp_avg_sq\"].zero_()\n                state[\"step\"].zero_()\n        self.counter.zero_()\n        torch.cuda.manual_seed(SEED + 1)\n\n\ndef mlp(train_x, train_y, test_x):\n    device = train_x.device\n    x = train_x.reshape(train_x.shape[0], -1).float()\n    q = test_x.reshape(test_x.shape[0], -1).float()\n    mean, std = x.mean(0), x.std(0).clamp_min(1e-3)\n    key = (tuple(x.shape), str(device))\n    if key not in CACHE:\n        CACHE.clear()\n        CACHE[key] = Trainer(x.shape[0], x.shape[1], device)\n    trainer = CACHE[key]\n    with torch.no_grad():\n        trainer.x.copy_((x - mean) / std)\n        trainer.targets.copy_(F.one_hot(train_y.long(), CLASSES).float() * (1.0 - SMOOTHING) + SMOOTHING / CLASSES)\n    trainer.initialise()\n    start = [p.detach().clone() for p in trainer.params]\n    for _ in range(TARGET_STEPS):\n        trainer.graph.replay()\n    with torch.no_grad():\n        bias = EMA ** TARGET_STEPS\n        averaged = [(s - bias * s0) / (1.0 - bias) for s, s0 in zip(trainer.shadow, start)]\n        q = (q - mean) / std\n        probabilities = torch.zeros(q.shape[0], CLASSES, device=device)\n        for begin in range(0, q.shape[0], 2000):\n            chunk = q[begin:begin + 2000].unsqueeze(0).expand(K, -1, -1).contiguous()\n            probabilities[begin:begin + 2000] = F.softmax(forward(averaged, chunk, False), dim=-1).sum(0)\n    return probabilities.argmax(1)\n",
      "sha256": "24453851a3faec1808bb157296ac3bab21f8c88b170dea11352bd98dbd6d0ba5"
    },
    "ours": {
      "filename": "kernel_pcg.py",
      "function": "classify",
      "source": "\"\"\"Non-neural kernel ridge classifier, refitted from the supplied labels.\"\"\"\n\nimport torch\nimport triton\nimport triton.language as tl\n\ntorch.backends.cuda.matmul.allow_tf32 = False\n\nGAMMA, RIDGE, STEPS = 0.02, 0.1, 16\nRANK = 256\nKIND = \"rbf\"\nNORMALIZE = True\nUSE_TRITON = True\nMETRIC_UPDATES, METRIC_SAMPLES, METRIC_BLEND = 1, 512, 0.05\n\n\n@triton.jit\ndef _product(K, P, PART, N: tl.constexpr, SPLITS: tl.constexpr,\n             BM: tl.constexpr, BK: tl.constexpr):\n    rows = tl.program_id(0) * BM + tl.arange(0, BM)\n    split = tl.program_id(1)\n    cols = tl.arange(0, 16)\n    acc = tl.zeros((BM, 16), tl.float32)\n    for block in range(tl.cdiv(N, BK * SPLITS)):\n        inner = (block * SPLITS + split) * BK + tl.arange(0, BK)\n        a = tl.load(K + rows[:, None] * N + inner[None, :],\n                    (rows[:, None] < N) & (inner[None, :] < N), 0)\n        b = tl.load(P + inner[:, None] * 10 + cols[None, :],\n                    (inner[:, None] < N) & (cols[None, :] < 10), 0)\n        acc = tl.dot(a, b, acc, input_precision=\"tf32x3\")\n    tl.store(PART + split * N * 16 + rows[:, None] * 16 + cols[None, :],\n             acc, rows[:, None] < N)\n\n\n@triton.jit\ndef _product_reduce(PART, P, OUT, N: tl.constexpr, RIDGE: tl.constexpr,\n                    SPLITS: tl.constexpr, BLOCK: tl.constexpr):\n    index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)\n    row, col = index // 10, index % 10\n    splits = tl.arange(0, SPLITS)\n    partial = tl.load(PART + splits[:, None] * N * 16 + row[None, :] * 16 + col[None, :],\n                      row[None, :] < N, 0)\n    p = tl.load(P + index, index < N * 10, 0)\n    tl.store(OUT + index, tl.sum(partial, 0) + RIDGE * p, index < N * 10)\n\n\ndef product(k, direction, ridge):\n    if not USE_TRITON:\n        return k @ direction + ridge * direction\n    n = k.shape[0]\n    partial = torch.empty((8, n, 16), device=k.device)\n    out = torch.empty_like(direction)\n    _product[(triton.cdiv(n, 32), 8)](k, direction, partial, n, 8, 32, 64)\n    _product_reduce[(triton.cdiv(n * 10, 256),)](partial, direction, out, n, ridge, 8, 256)\n    return out\n\n\ndef kernel(x, z, gamma, kind):\n    distance = (x.square().sum(1)[:, None] + z.square().sum(1)[None, :] - 2 * (x @ z.T)).clamp_min_(0)\n    if kind == \"laplacian\":\n        distance.sqrt_()\n    return distance.mul_(-gamma).exp_()\n\n\ndef solve(k, targets, ridge, steps):\n    if RANK:\n        m = min(RANK, len(k))\n        landmark = k[:m, :m].clone()\n        landmark.diagonal().add_(1e-4)\n        factor = torch.linalg.cholesky(landmark)\n        u = torch.linalg.solve_triangular(factor, k[:, :m].T.contiguous(), upper=False).T.contiguous()\n        small = u.T @ u\n        small.diagonal().add_(ridge)\n        factor = torch.linalg.cholesky(small)\n        v = torch.linalg.solve_triangular(factor, u.T.contiguous(), upper=False).T.contiguous()\n\n    def precondition(residual):\n        if not RANK:\n            return residual\n        return (residual - v @ (v.T @ residual)) / ridge\n\n    weights = torch.zeros_like(targets)\n    residual = targets.clone()\n    direction = precondition(residual).clone()\n    norm = (residual * direction).sum(0)\n    for _ in range(steps):\n        applied = product(k, direction, ridge)\n        alpha = norm / (direction.mul(applied).sum(0).clamp_min(1e-20))\n        weights.add_(direction * alpha)\n        residual.sub_(applied * alpha)\n        preconditioned = precondition(residual)\n        next_norm = (residual * preconditioned).sum(0)\n        direction = preconditioned + direction * (next_norm / norm.clamp_min(1e-20))\n        norm = next_norm\n    return weights\n\n\ndef prepare(train_x, test_x):\n    x, q = train_x.float(), test_x.float()\n    if NORMALIZE:\n        x = x / x.norm(dim=1, keepdim=True).clamp_min(1e-6) * (x.shape[1] ** 0.5)\n        q = q / q.norm(dim=1, keepdim=True).clamp_min(1e-6) * (q.shape[1] ** 0.5)\n    return x, q\n\n\ndef gradient_metric(x, transform, weights, gamma, kind, samples, blend):\n    z = x @ transform\n    points = z[:samples]\n    factor = kernel(points, z, gamma, kind)\n    if kind == \"laplacian\":\n        distance = (points.square().sum(1)[:, None] + z.square().sum(1)[None] - 2 * (points @ z.T)).clamp_min_(0)\n        factor = factor / distance.sqrt_().clamp_min_(1e-6)\n        index = torch.arange(len(points), device=x.device)\n        factor[index, index] = 0\n    scores = factor @ weights\n    weighted_points = (weights[:, :, None] * z[:, None, :]).reshape(len(z), -1)\n    gradients = (factor @ weighted_points).reshape(len(points), 10, -1) - scores[:, :, None] * points[:, None, :]\n    gradients = gradients @ transform.T\n    gradients = gradients.reshape(-1, x.shape[1])\n    metric = gradients.T @ gradients\n    metric = metric * (x.shape[1] / metric.trace().clamp_min(1e-12))\n    metric = (1 - blend) * metric + blend * torch.eye(x.shape[1], device=x.device)\n    values, vectors = torch.linalg.eigh(metric)\n    return vectors * values.clamp_min(1e-6).sqrt()[None, :]\n\n\ndef classify(train_x, train_y, test_x):\n    with torch.no_grad():\n        x, q = prepare(train_x, test_x)\n        targets = torch.zeros((len(x), 10), device=x.device)\n        targets.scatter_(1, train_y[:, None], 1.0)\n        # ponytail: quadratic storage for feasibility; use landmarks if bandwidth dominates.\n        transform = torch.eye(x.shape[1], device=x.device)\n        for update in range(METRIC_UPDATES + 1):\n            z = x @ transform if METRIC_UPDATES else x\n            k = kernel(z, z, GAMMA, KIND)\n            k.diagonal().fill_(1.0)\n            weights = solve(k, targets, RIDGE, STEPS)\n            if update < METRIC_UPDATES:\n                transform = gradient_metric(x, transform, weights, GAMMA, KIND, METRIC_SAMPLES, METRIC_BLEND)\n        q = q @ transform if METRIC_UPDATES else q\n        return torch.cat([\n            (kernel(chunk, z, GAMMA, KIND) @ weights).argmax(1)\n            for chunk in q.split(2000)\n        ])\n\n\ndef self_check():\n    global USE_TRITON, RANK\n    original = USE_TRITON\n    torch.manual_seed(782)\n    x = torch.randn(128, 12, device=\"cuda\")\n    targets = torch.randn(128, 10, device=\"cuda\")\n    k = kernel(x, x, 0.1, \"rbf\")\n    rotation, _ = torch.linalg.qr(torch.randn(12, 12, device=\"cuda\"))\n    torch.testing.assert_close(kernel(x @ rotation, x @ rotation, 0.1, \"rbf\"), k,\n                               atol=1e-5, rtol=1e-5)\n    actual = solve(k, targets, 0.5, 64)\n    expected = torch.linalg.solve(k + 0.5 * torch.eye(128, device=\"cuda\"), targets)\n    torch.testing.assert_close(actual, expected, atol=0.003, rtol=0.003)\n    original_rank = RANK\n    try:\n        RANK = 0\n        torch.testing.assert_close(solve(k, targets, 0.5, 64), expected, atol=0.003, rtol=0.003)\n    finally:\n        RANK = original_rank\n    try:\n        USE_TRITON = True\n        for n in (128, 1003):\n            matrix = torch.randn(n, n, device=\"cuda\")\n            direction = torch.randn(n, 10, device=\"cuda\")\n            torch.testing.assert_close(product(matrix, direction, 0.1), matrix @ direction + 0.1 * direction,\n                                       atol=0.001, rtol=0.001)\n        torch.testing.assert_close(solve(k, targets, 0.5, 64), expected, atol=0.003, rtol=0.003)\n    finally:\n        USE_TRITON = original\n    # Check the supervised metric against a small autograd Jacobian, not just itself.\n    small = torch.randn(17, 4, device=\"cuda\")\n    coefficients = torch.randn(17, 10, device=\"cuda\")\n    transform = torch.randn(4, 4, device=\"cuda\") / 2\n    actual = gradient_metric(small, transform, coefficients, 0.2, \"rbf\", 5, 0.05)\n    points = small[:5].clone().requires_grad_()\n    scores = kernel(points @ transform, small @ transform, 0.2, \"rbf\") @ coefficients\n    gradients = torch.stack([torch.autograd.grad(scores[:, c].sum(), points, retain_graph=True)[0]\n                             for c in range(10)], 1).reshape(-1, 4)\n    expected = gradients.T @ gradients\n    expected = 0.95 * expected * (4 / expected.trace()) + 0.05 * torch.eye(4, device=\"cuda\")\n    torch.testing.assert_close(actual @ actual.T, expected, atol=1e-4, rtol=1e-4)\n\n\nif __name__ == \"__main__\":\n    self_check()\n",
      "sha256": "fc0ddef0d3dba8f00fc44f9cd576548c2a5184c5cb88741767d33d6f10b38012"
    }
  },
  "scorer_sha256": "1bc6d8d96a776565bc70f3ccdb91be7de57187d4495b2b0437fac98f5e120064",
  "same_inputs": false,
  "design": "one board; three fresh-process stock-scored runs per method"
}
