{
  "rows": [
    {
      "variant": "reference",
      "n": 1000,
      "width": 1000,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 95.77471733093262
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 77.92640209197998
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 4,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 56,
      "n_spills": 0,
      "shared_bytes": 16,
      "cuda_us": 131.59423828125
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 138.99776458740234
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 16,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 142.39744186401367
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 4,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 122,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 89.12896156311035
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 16,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 77.63967990875244
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 116,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 86.15936279296875
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 63.52896213531494
    },
    {
      "variant": "reference",
      "n": 1000,
      "width": 500,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 40.4582405090332
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 40.427517890930176
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 4,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 56,
      "n_spills": 0,
      "shared_bytes": 16,
      "cuda_us": 64.13311958312988
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 67.43040084838867
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 16,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 72.4787187576294
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 4,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 122,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 40.263681411743164
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 16,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 41.349120140075684
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 116,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 29.224960803985596
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 27.115519046783447
    },
    {
      "variant": "reference",
      "n": 1000,
      "width": 250,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 25.518081188201904
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 25.487360954284668
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 4,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 56,
      "n_spills": 0,
      "shared_bytes": 16,
      "cuda_us": 40.908799171447754
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 41.51296138763428
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 16,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 42.72128105163574
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 4,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 122,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 25.589759349822998
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 16,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 27.22815990447998
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 116,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 19.42528009414673
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 16.773120164871216
    },
    {
      "variant": "reference",
      "n": 1000,
      "width": 60,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 14.510079622268677
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 14.510079622268677
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 4,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 56,
      "n_spills": 0,
      "shared_bytes": 16,
      "cuda_us": 18.053120374679565
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 17.26464033126831
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 16,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 16.599040031433105
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 4,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 122,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 16.74239993095398
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 16,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 14.57152009010315
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 116,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 17.879040241241455
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 14.714879989624023
    },
    {
      "variant": "reference",
      "n": 1000,
      "width": 10,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 12.584960460662842
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 12.584960460662842
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 4,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 56,
      "n_spills": 0,
      "shared_bytes": 16,
      "cuda_us": 12.646399736404419
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 10.321919918060303
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 16,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 32,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 11.263999938964844
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 4,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 126,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 14.05951976776123
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 16,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 62,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 10.94655990600586
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 128,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 11.612160205841064
    },
    {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 12.615679502487183
    },
    {
      "variant": "reference",
      "n": 257,
      "width": 1000,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 16.076799631118774
    },
    {
      "variant": "candidate",
      "n": 257,
      "width": 1000,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 16.076799631118774
    },
    {
      "variant": "reference",
      "n": 257,
      "width": 500,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 12.974079847335815
    },
    {
      "variant": "candidate",
      "n": 257,
      "width": 500,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 12.974079847335815
    },
    {
      "variant": "reference",
      "n": 257,
      "width": 250,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 11.41759991645813
    },
    {
      "variant": "candidate",
      "n": 257,
      "width": 250,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 11.41759991645813
    },
    {
      "variant": "reference",
      "n": 257,
      "width": 60,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 10.444799661636353
    },
    {
      "variant": "candidate",
      "n": 257,
      "width": 60,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 10.43455958366394
    },
    {
      "variant": "reference",
      "n": 257,
      "width": 10,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 10.4038405418396
    },
    {
      "variant": "candidate",
      "n": 257,
      "width": 10,
      "warps": 8,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 128,
      "cuda_us": 10.4038405418396
    }
  ],
  "selected": {
    "1000": {
      "variant": "candidate",
      "n": 1000,
      "width": 1000,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 63.52896213531494
    },
    "500": {
      "variant": "candidate",
      "n": 1000,
      "width": 500,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 27.115519046783447
    },
    "250": {
      "variant": "candidate",
      "n": 1000,
      "width": 250,
      "warps": 16,
      "columns": 4,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 256,
      "cuda_us": 16.773120164871216
    },
    "60": {
      "variant": "candidate",
      "n": 1000,
      "width": 60,
      "warps": 8,
      "columns": 2,
      "passed": true,
      "error": null,
      "n_regs": 64,
      "n_spills": 0,
      "shared_bytes": 64,
      "cuda_us": 14.510079622268677
    },
    "10": {
      "variant": "candidate",
      "n": 1000,
      "width": 10,
      "warps": 8,
      "columns": 1,
      "passed": true,
      "error": null,
      "n_regs": 40,
      "n_spills": 0,
      "shared_bytes": 32,
      "cuda_us": 10.321919918060303
    }
  },
  "energy": {
    "baseline": {
      "above_idle_j_per_kernel_sequence": 0.032478745081445666,
      "calls": 17090,
      "idle_before_w": 68.91076320668893,
      "idle_after_w": 84.523200495456
    },
    "selected": {
      "above_idle_j_per_kernel_sequence": 0.024043014319111957,
      "calls": 21852,
      "idle_before_w": 69.25115789681739,
      "idle_after_w": 86.1640471577071
    }
  },
  "kernel_sequence_widths": [
    1000,
    500,
    250,
    250,
    250,
    10
  ],
  "gpu": "NVIDIA A100-SXM4-80GB",
  "device": {
    "name": "NVIDIA A100-SXM4-80GB",
    "uuid": "GPU-ccf053cd-b578-b752-53e0-d260d2587357",
    "vbios": "92.00.9E.00.02",
    "driver": "580.95.05",
    "power_limit_w": 500.0
  },
  "elapsed_s": 22.055326562999998,
  "official": false,
  "source_sha256": "2176514de32020ff9924896583f2f178a076c95e45c9b1cc7a71699be0053b47",
  "source": "\"\"\"Derived from SethTS, ladder-sampled-20260929; experiment: reduced training steps.\n\nLadder network for mnist-a100: batch 1,000, TF32, fused Triton kernels, labelled rows sampled by their loss.\n\nModel and objective: `ladder-xlong-s11` of mnist/experiments/release-cutoffs-20260925 (the fully supervised\nAMLP[2,2] Ladder of Pezeshki et al., ICML 2016): 60-1000-500-250-250-250-10, noise 0.3 at the input and every\nlayer, inputs scaled by 0.6, input reconstruction weight 2000, labelled and unlabelled streams normalised\nseparately, unlabelled rows drawn from train_x and test_x together (transductive), BatchNorm calibrated on the\ntraining rows at the end. Retuned as in ../ladder-fast-20260929: 1,000 labelled and 1,000 unlabelled rows per\nstep, Adam peak learning rate 0.008 with the recipe's 100/150 schedule shape, TF32 matmuls, fewer steps.\n\nNew here (submissions/ladder-sampled-20260929/README.md): each step draws its labelled rows with\np = (1 - mix) / n + mix * loss / sum(loss), where loss is each row's most recent training cross-entropy,\nrecorded by the step itself. mix is 0 for the first 10% of steps, MIX until 60%, then falls linearly to 0 at\n90%, so training ends on uniform sampling. Unlabelled rows keep the per-epoch shuffle.\n\nAs in ../ladder-triton-20260928: fused Triton kernels per layer (encoder: per-stream batch norm, in-kernel\nnoise, bias, activation; decoder: batch norm and the unit-wise combinator) with exact batch statistics; the\nstep captured in a CUDA graph on the warm-up and replayed, all learned state re-initialised every call.\n\"\"\"\nimport math\nimport os\n\nimport torch\nimport torch.nn.functional as F\nimport triton\nimport triton.language as tl\n\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = False\n\nSTEPS, BATCH = 800, 1000\nDIMS = (60, 1000, 500, 250, 250, 250, 10)\nNOISE, INPUT_SCALE, RECONSTRUCTION = 0.3, 0.6, 2000.0\nLR, EPOCHS_CONFIG, DECAY_START_CONFIG = 0.008, 150, 100\nEPS, COMBINATOR_STD = 1e-10, 0.025\nSEED = 11\nMIX = 0.9  # peak weight of loss-proportional sampling\nWARPS, TILE = 8, 4096  # warps per Triton program; each program holds all rows and TILE // rows columns\nCACHE = {}\n# Combinator parameter rows: W1[i, o] = 2i + o, b1 6-7, W2[j, o] = 8 + 2j + o, b2 12-13, W3[j] = 14 + j, b3 16.\nWEIGHT_ROWS = (0, 1, 2, 3, 4, 5, 8, 9, 10, 11, 14, 15)\n\n\n@triton.jit\ndef _leaky(x):\n    return tl.where(x > 0, x, 0.1 * x)\n\n\n@triton.jit\ndef _column_norm(u, rows_ok, n, eps):\n    \"\"\"Biased batch statistics over the rows of a [rows, columns] tile; returns (normalised, rstd).\"\"\"\n    mean = tl.sum(u, 0) / n\n    centred = tl.where(rows_ok, u - mean[None, :], 0.0)\n    rstd = 1.0 / tl.sqrt(tl.sum(centred * centred, 0) / n + eps)\n    return centred * rstd[None, :], rstd\n\n\n@triton.jit\ndef _decoder_forward(u_ptr, lat_ptr, p_ptr, out_ptr, n, s, eps, ROWS: tl.constexpr, COLS: tl.constexpr):\n    cols = tl.program_id(0) * COLS + tl.arange(0, COLS)\n    rows = tl.arange(0, ROWS)\n    rows_ok = (rows < n)[:, None]\n    ok = rows_ok & (cols < s)[None, :]\n    index = rows[:, None] * s + cols[None, :]\n    v, _ = _column_norm(tl.load(u_ptr + index, mask=ok, other=0.0), rows_ok, n, eps)\n    l = tl.load(lat_ptr + index, mask=ok, other=0.0)\n    col_ok = cols < s\n    p = p_ptr + cols\n    x2 = v * l\n    a0 = v * tl.load(p, col_ok)[None, :] + l * tl.load(p + 2 * s, col_ok)[None, :] + x2 * tl.load(p + 4 * s, col_ok)[None, :] + tl.load(p + 6 * s, col_ok)[None, :]\n    a1 = v * tl.load(p + s, col_ok)[None, :] + l * tl.load(p + 3 * s, col_ok)[None, :] + x2 * tl.load(p + 5 * s, col_ok)[None, :] + tl.load(p + 7 * s, col_ok)[None, :]\n    h0, h1 = _leaky(a0), _leaky(a1)\n    c0 = h0 * tl.load(p + 8 * s, col_ok)[None, :] + h1 * tl.load(p + 10 * s, col_ok)[None, :] + tl.load(p + 12 * s, col_ok)[None, :]\n    c1 = h0 * tl.load(p + 9 * s, col_ok)[None, :] + h1 * tl.load(p + 11 * s, col_ok)[None, :] + tl.load(p + 13 * s, col_ok)[None, :]\n    out = _leaky(c0) * tl.load(p + 14 * s, col_ok)[None, :] + _leaky(c1) * tl.load(p + 15 * s, col_ok)[None, :] + tl.load(p + 16 * s, col_ok)[None, :]\n    tl.store(out_ptr + index, out, mask=ok)\n\n\n@triton.jit\ndef _decoder_backward(u_ptr, lat_ptr, p_ptr, g_ptr, du_ptr, dlat_ptr, dp_ptr, n, s, eps,\n                      ROWS: tl.constexpr, COLS: tl.constexpr):\n    cols = tl.program_id(0) * COLS + tl.arange(0, COLS)\n    rows = tl.arange(0, ROWS)\n    rows_ok = (rows < n)[:, None]\n    col_ok = cols < s\n    ok = rows_ok & col_ok[None, :]\n    index = rows[:, None] * s + cols[None, :]\n    v, rstd = _column_norm(tl.load(u_ptr + index, mask=ok, other=0.0), rows_ok, n, eps)\n    l = tl.load(lat_ptr + index, mask=ok, other=0.0)\n    g = tl.load(g_ptr + index, mask=ok, other=0.0)\n    p = p_ptr + cols\n    w00, w01 = tl.load(p, col_ok)[None, :], tl.load(p + s, col_ok)[None, :]\n    w10, w11 = tl.load(p + 2 * s, col_ok)[None, :], tl.load(p + 3 * s, col_ok)[None, :]\n    w20, w21 = tl.load(p + 4 * s, col_ok)[None, :], tl.load(p + 5 * s, col_ok)[None, :]\n    u00, u01 = tl.load(p + 8 * s, col_ok)[None, :], tl.load(p + 9 * s, col_ok)[None, :]\n    u10, u11 = tl.load(p + 10 * s, col_ok)[None, :], tl.load(p + 11 * s, col_ok)[None, :]\n    z0, z1 = tl.load(p + 14 * s, col_ok)[None, :], tl.load(p + 15 * s, col_ok)[None, :]\n    x2 = v * l\n    a0 = v * w00 + l * w10 + x2 * w20 + tl.load(p + 6 * s, col_ok)[None, :]\n    a1 = v * w01 + l * w11 + x2 * w21 + tl.load(p + 7 * s, col_ok)[None, :]\n    h0, h1 = _leaky(a0), _leaky(a1)\n    c0 = h0 * u00 + h1 * u10 + tl.load(p + 12 * s, col_ok)[None, :]\n    c1 = h0 * u01 + h1 * u11 + tl.load(p + 13 * s, col_ok)[None, :]\n    q = dp_ptr + cols\n    tl.store(q + 14 * s, tl.sum(g * _leaky(c0), 0), col_ok)\n    tl.store(q + 15 * s, tl.sum(g * _leaky(c1), 0), col_ok)\n    tl.store(q + 16 * s, tl.sum(g, 0), col_ok)\n    dc0 = g * z0 * tl.where(c0 > 0, 1.0, 0.1)\n    dc1 = g * z1 * tl.where(c1 > 0, 1.0, 0.1)\n    tl.store(q + 8 * s, tl.sum(dc0 * h0, 0), col_ok)\n    tl.store(q + 9 * s, tl.sum(dc1 * h0, 0), col_ok)\n    tl.store(q + 10 * s, tl.sum(dc0 * h1, 0), col_ok)\n    tl.store(q + 11 * s, tl.sum(dc1 * h1, 0), col_ok)\n    tl.store(q + 12 * s, tl.sum(dc0, 0), col_ok)\n    tl.store(q + 13 * s, tl.sum(dc1, 0), col_ok)\n    da0 = (dc0 * u00 + dc1 * u01) * tl.where(a0 > 0, 1.0, 0.1)\n    da1 = (dc0 * u10 + dc1 * u11) * tl.where(a1 > 0, 1.0, 0.1)\n    tl.store(q, tl.sum(da0 * v, 0), col_ok)\n    tl.store(q + s, tl.sum(da1 * v, 0), col_ok)\n    tl.store(q + 2 * s, tl.sum(da0 * l, 0), col_ok)\n    tl.store(q + 3 * s, tl.sum(da1 * l, 0), col_ok)\n    tl.store(q + 4 * s, tl.sum(da0 * x2, 0), col_ok)\n    tl.store(q + 5 * s, tl.sum(da1 * x2, 0), col_ok)\n    tl.store(q + 6 * s, tl.sum(da0, 0), col_ok)\n    tl.store(q + 7 * s, tl.sum(da1, 0), col_ok)\n    dx2 = da0 * w20 + da1 * w21\n    dv = da0 * w00 + da1 * w01 + dx2 * l\n    tl.store(dlat_ptr + index, da0 * w10 + da1 * w11 + dx2 * v, mask=ok)\n    dv = tl.where(ok, dv, 0.0)\n    du = rstd[None, :] * (dv - (tl.sum(dv, 0) / n)[None, :] - v * (tl.sum(dv * v, 0) / n)[None, :])\n    tl.store(du_ptr + index, du, mask=ok)\n\n\n\n@triton.jit\ndef _encoder_forward(raw_ptr, beta_ptr, gamma_ptr, counter_ptr, z_ptr, h_ptr, n, d, eps, noise, layer,\n                     TOP: tl.constexpr, ROWS: tl.constexpr, COLS: tl.constexpr):\n    \"\"\"Both streams (labelled rows 0..n-1, unlabelled n..2n-1), normalised separately, plus noise and activation.\"\"\"\n    cols = tl.program_id(0) * COLS + tl.arange(0, COLS)\n    rows = tl.arange(0, ROWS)\n    rows_ok = (rows < n)[:, None]\n    col_ok = cols < d\n    ok = rows_ok & col_ok[None, :]\n    beta = tl.load(beta_ptr + cols, col_ok)[None, :]\n    seed = tl.load(counter_ptr) * 16 + layer + 1234\n    for stream in range(2):\n        index = (stream * n + rows[:, None]) * d + cols[None, :]\n        xhat, _ = _column_norm(tl.load(raw_ptr + index, mask=ok, other=0.0), rows_ok, n, eps)\n        z = xhat + noise * tl.randn(seed, index)\n        tl.store(z_ptr + index, z, mask=ok)\n        if TOP:\n            h = (z + beta) * tl.load(gamma_ptr + cols, col_ok)[None, :]\n        else:\n            h = tl.maximum(z + beta, 0.0)\n        tl.store(h_ptr + index, h, mask=ok)\n\n\n@triton.jit\ndef _encoder_backward(raw_ptr, z_ptr, beta_ptr, gamma_ptr, dz_ptr, dh_ptr, draw_ptr, dbeta_ptr, dgamma_ptr,\n                      n, d, eps, TOP: tl.constexpr, ROWS: tl.constexpr, COLS: tl.constexpr):\n    cols = tl.program_id(0) * COLS + tl.arange(0, COLS)\n    rows = tl.arange(0, ROWS)\n    rows_ok = (rows < n)[:, None]\n    col_ok = cols < d\n    ok = rows_ok & col_ok[None, :]\n    beta = tl.load(beta_ptr + cols, col_ok)[None, :]\n    dbeta = tl.zeros((COLS,), tl.float32)\n    dgamma = tl.zeros((COLS,), tl.float32)\n    for stream in range(2):\n        index = (stream * n + rows[:, None]) * d + cols[None, :]\n        xhat, rstd = _column_norm(tl.load(raw_ptr + index, mask=ok, other=0.0), rows_ok, n, eps)\n        pre = tl.load(z_ptr + index, mask=ok, other=0.0) + beta\n        dh = tl.load(dh_ptr + index, mask=ok, other=0.0)\n        if TOP:\n            dpre = dh * tl.load(gamma_ptr + cols, col_ok)[None, :]\n            dgamma += tl.sum(dh * pre, 0)\n        else:\n            dpre = tl.where(pre > 0, dh, 0.0)\n        dbeta += tl.sum(dpre, 0)\n        dz = tl.where(ok, dpre + tl.load(dz_ptr + index, mask=ok, other=0.0), 0.0)\n        draw = rstd[None, :] * (dz - (tl.sum(dz, 0) / n)[None, :] - xhat * (tl.sum(dz * xhat, 0) / n)[None, :])\n        tl.store(draw_ptr + index, draw, mask=ok)\n    tl.store(dbeta_ptr + cols, dbeta, col_ok)\n    if TOP:\n        tl.store(dgamma_ptr + cols, dgamma, col_ok)\n\n\ndef _columns(n, forward=False):\n    return max(1, (TILE if forward else 2048) // triton.next_power_of_2(n))\n\n\ndef _launch(width, n, forward=False):\n    return (triton.cdiv(width, _columns(n, forward)),)\n\n\nclass Decode(torch.autograd.Function):\n    \"\"\"out = combinator(lateral, normalize(u)), unit-wise, for u and lateral of shape (rows, units).\"\"\"\n\n    @staticmethod\n    def forward(ctx, u, lateral, p):\n        n, s = u.shape\n        out = torch.empty_like(u)\n        _decoder_forward[_launch(s, n, True)](u, lateral, p, out, n, s, EPS, ROWS=triton.next_power_of_2(n), COLS=_columns(n, True), num_warps=WARPS)\n        ctx.save_for_backward(u, lateral, p)\n        return out\n\n    @staticmethod\n    def backward(ctx, g):\n        u, lateral, p = ctx.saved_tensors\n        n, s = u.shape\n        du, dlat, dp = torch.empty_like(u), torch.empty_like(lateral), torch.empty_like(p)\n        cols = (4 if s >= 250 else 2 if s >= 60 else 1) if n == BATCH else _columns(n)\n        _decoder_backward[(triton.cdiv(s, cols),)](u, lateral, p, g.contiguous(), du, dlat, dp, n, s, EPS,\n                                      ROWS=triton.next_power_of_2(n), COLS=cols, num_warps=WARPS)\n        return du, dlat, dp\n\n\nclass Encode(torch.autograd.Function):\n    \"\"\"(z, h) for raw of shape (2, rows, units): z = normalize(raw) per stream + noise, h its activation.\"\"\"\n\n    @staticmethod\n    def forward(ctx, raw, beta, gamma, counter, layer):\n        _, n, d = raw.shape\n        z, h = torch.empty_like(raw), torch.empty_like(raw)\n        top = gamma is not None\n        _encoder_forward[_launch(d, n, True)](raw, beta, gamma if top else beta, counter, z, h, n, d, EPS, NOISE, layer,\n                                     TOP=top, ROWS=triton.next_power_of_2(n), COLS=_columns(n, True), num_warps=WARPS)\n        ctx.top = top\n        ctx.save_for_backward(raw, z, beta, gamma if top else beta)\n        return z, h\n\n    @staticmethod\n    def backward(ctx, dz, dh):\n        raw, z, beta, gamma = ctx.saved_tensors\n        _, n, d = raw.shape\n        dz = torch.zeros_like(raw) if dz is None else dz.contiguous()\n        dh = torch.zeros_like(raw) if dh is None else dh.contiguous()\n        draw, dbeta, dgamma = torch.empty_like(raw), torch.empty_like(beta), torch.empty_like(gamma)\n        _encoder_backward[_launch(d, n)](raw, z, beta, gamma, dz, dh, draw, dbeta, dgamma, n, d, EPS,\n                                      TOP=ctx.top, ROWS=triton.next_power_of_2(n), COLS=_columns(n), num_warps=WARPS)\n        return draw, dbeta, dgamma if ctx.top else None, None, None\n\n\ndef loss(params, counter, x_labelled, y_labelled, x_unlabelled):\n    h = torch.stack((x_labelled, x_unlabelled))  # the two streams, normalised separately\n    h = h + NOISE * torch.randn_like(h)\n    lateral = [h[1]]\n    last = len(DIMS) - 2\n    for layer, weight in enumerate(params[\"encoder\"]):\n        z, h = Encode.apply(h @ weight, params[\"beta\"][layer], params[\"gamma\"] if layer == last else None,\n                            counter, layer)\n        lateral.append(z[1])\n    reconstruction = Decode.apply(F.softmax(h[1], dim=-1), lateral[-1], params[\"combinators\"][-1])\n    for layer in range(last, -1, -1):\n        reconstruction = Decode.apply(reconstruction @ params[\"decoder\"][layer], lateral[layer],\n                                      params[\"combinators\"][layer])\n    ce = F.cross_entropy(h[0], y_labelled, reduction=\"none\")\n    return ce.mean() + RECONSTRUCTION * F.mse_loss(reconstruction, x_unlabelled), ce\n\n\ndef encode(params, layer, z):\n    \"\"\"The clean encoder's activation after normalisation, for calibration and prediction.\"\"\"\n    beta = params[\"beta\"][layer]\n    if layer == len(DIMS) - 2:\n        return (z + beta) * params[\"gamma\"]\n    return F.relu(z + beta)\n\n\nclass Trainer:\n    def __init__(self, n, pool_n, device):\n        self.device = device\n        self.steps_per_epoch = math.ceil(n / BATCH)\n        self.epochs = max(1, math.ceil(STEPS / self.steps_per_epoch))\n        self.total = self.epochs * self.steps_per_epoch\n        decay_start = min(self.epochs, max(1, round(self.epochs * DECAY_START_CONFIG / EPOCHS_CONFIG)))\n        factor = [1.0 if self.epochs <= decay_start else\n                  max(0.0, min(1.0, (self.epochs - e) / (self.epochs - decay_start))) for e in range(self.epochs)]\n        self.schedule = (LR * torch.tensor(factor, device=device)).repeat_interleave(self.steps_per_epoch)\n        self.n, self.pool_n = n, pool_n\n        self.x = torch.zeros(n, DIMS[0], device=device)\n        self.y = torch.zeros(n, dtype=torch.long, device=device)\n        self.pool = torch.zeros(pool_n, DIMS[0], device=device)\n        self.unlabelled = torch.zeros(self.total, BATCH, dtype=torch.long, device=device)\n        self.counter = torch.zeros(1, dtype=torch.long, device=device)\n        # Loss-proportional sampling of labelled rows: p = (1 - mix_t) / n + mix_t * score / sum(score),\n        # scores = each row's most recent training cross-entropy, mix_t a per-step table.\n        self.scores = torch.ones(n, device=device)\n        f = torch.arange(self.total, device=device) / self.total\n        self.mix = MIX * torch.where(f < 0.1, 0.0, ((0.9 - f) / 0.3).clamp(0.0, 1.0))\n        zeros = lambda *shape: torch.zeros(*shape, device=device, requires_grad=True)\n        self.params = {\n            \"encoder\": [zeros(a, b) for a, b in zip(DIMS[:-1], DIMS[1:])],\n            \"beta\": [zeros(b) for b in DIMS[1:]],\n            \"gamma\": zeros(DIMS[-1]),\n            \"decoder\": [zeros(b, a) for a, b in zip(DIMS[:-1], DIMS[1:])],\n            \"combinators\": [zeros(17, s) for s in DIMS],\n        }\n        self.flat = (self.params[\"encoder\"] + self.params[\"beta\"] + [self.params[\"gamma\"]] + self.params[\"decoder\"]\n                     + self.params[\"combinators\"])\n        self.lr = torch.tensor(LR, device=device)\n        self.optimizer = torch.optim.Adam(self.flat, lr=self.lr, betas=(0.9, 0.999), eps=1e-8,\n                                          fused=True, capturable=True)\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        lam = self.mix.index_select(0, self.counter).squeeze(0)\n        cdf = torch.cumsum((1.0 - lam) / self.n + lam * self.scores / self.scores.sum(), 0)\n        labelled = torch.searchsorted(cdf, torch.rand(BATCH, device=self.device) * cdf[-1]).clamp_(max=self.n - 1)\n        unlabelled = self.unlabelled.index_select(0, self.counter).squeeze(0)\n        self.lr.copy_(self.schedule.index_select(0, self.counter).squeeze(0))\n        total, ce = loss(self.params, self.counter, self.x[labelled], self.y[labelled], self.pool[unlabelled])\n        total.backward()\n        self.optimizer.step()\n        with torch.no_grad():\n            self.scores.scatter_(0, labelled, ce.detach())\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        normal = lambda p, std: p.copy_(torch.randn(p.shape, generator=generator, device=self.device) * std)\n        for weight in self.params[\"encoder\"] + self.params[\"decoder\"]:\n            normal(weight, weight.shape[0] ** -0.5)\n        for p in self.params[\"combinators\"]:\n            p.zero_()\n            p[list(WEIGHT_ROWS)] = torch.randn(len(WEIGHT_ROWS), p.shape[1], generator=generator,\n                                               device=self.device) * COMBINATOR_STD\n        for p in self.params[\"beta\"]:\n            p.zero_()\n        self.params[\"gamma\"].fill_(1.0)\n        for p in self.flat:\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        self.scores.fill_(1.0)\n        # Unlabelled rows: each epoch the first steps_per_epoch * BATCH of a shuffle of train and test rows.\n        epochs, spe = self.epochs, self.steps_per_epoch\n        recon = torch.rand(epochs, self.pool_n, generator=generator, device=self.device).argsort(1)\n        self.unlabelled.copy_(recon[:, :spe * BATCH].reshape(self.total, BATCH))\n        torch.cuda.manual_seed(SEED + 1)\n\n    @torch.no_grad()\n    def predict(self, q):\n        \"\"\"Calibrate BatchNorm on the training rows (minibatch means and unbiased variances, averaged over\n        one pass, deeper layers fed minibatch-normalised activations), then run the clean encoder on q.\"\"\"\n        params = self.params\n        batches = self.n // BATCH\n        generator = torch.Generator(device=self.device)\n        generator.manual_seed(1)\n        rows = torch.randperm(self.n, generator=generator, device=self.device)[:batches * BATCH]\n        h = self.x[rows].reshape(batches, BATCH, DIMS[0])\n        stats = []\n        for layer, weight in enumerate(params[\"encoder\"]):\n            raw = h @ weight\n            variance, mean = torch.var_mean(raw, dim=1, unbiased=False, keepdim=True)\n            stats.append((mean.mean(0), variance.mean(0) * BATCH / (BATCH - 1)))\n            h = encode(params, layer, (raw - mean) * torch.rsqrt(variance + EPS))\n        h = q\n        for layer, weight in enumerate(params[\"encoder\"]):\n            mean, variance = stats[layer]\n            h = encode(params, layer, (h @ weight - mean) * torch.rsqrt(variance + EPS))\n        return h\n\n\ndef classify(train_x, train_y, test_x):\n    torch.backends.cuda.preferred_blas_library(\"default\")\n    # Triton compiles on first use and caches what it builds; the sandboxed worker's home may not be writable.\n    os.environ.setdefault(\"TRITON_CACHE_DIR\", \"/tmp/triton-cache\")\n    device = train_x.device\n    x = train_x.reshape(train_x.shape[0], -1).float() * INPUT_SCALE\n    q = test_x.reshape(test_x.shape[0], -1).float() * INPUT_SCALE\n    key = (tuple(x.shape), tuple(q.shape), str(device), STEPS, BATCH)\n    if key not in CACHE:\n        CACHE.clear()\n        CACHE[key] = Trainer(x.shape[0], x.shape[0] + q.shape[0], device)\n    trainer = CACHE[key]\n    with torch.no_grad():\n        trainer.x.copy_(x)\n        trainer.y.copy_(train_y.long())\n        trainer.pool.copy_(torch.cat((x, q)))\n    trainer.initialise()\n    for _ in range(trainer.total):\n        trainer.graph.replay()\n    return trainer.predict(q).argmax(1)\n",
  "reference_source_sha256": "9d99a55a3c700d0804a86a53e5772c619acbacf5fcc4658bdb967af0446d7b52",
  "app_id": "ap-mBaGrxzY23Bs5IDUqvmGVD",
  "launch_elapsed_s": 44.04689254099503
}
