{
  "configs": [
    {
      "name": "decoder-tuned-reference",
      "warps": 8,
      "tile": 4096,
      "blas": "default",
      "encoder_cols": 2,
      "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    # 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",
      "sha256": "6a937f3bf0e7d77f658312c223b3fe5d8602487141f3d7731503f3c9a5333c96"
    },
    {
      "name": "encoder-tuned",
      "warps": 8,
      "tile": 4096,
      "blas": "default",
      "encoder_cols": 4,
      "encoder_warps": 16,
      "source": "\"\"\"Derived from SethTS, ladder-sampled-20260929; experiment: reduced steps and tuned kernels.\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        cols = (4 if d >= 250 else 1) if n == BATCH else _columns(n)\n        warps = 16 if n == BATCH and d >= 250 else WARPS\n        _encoder_backward[(triton.cdiv(d, cols),)](raw, z, beta, gamma, dz, dh, draw, dbeta, dgamma, n, d, EPS,\n                                      TOP=ctx.top, ROWS=triton.next_power_of_2(n), COLS=cols, 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    # 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",
      "sha256": "5052f86e9cdf6086783bd64c90a31fcb86b3d19d1d0f27ef8f21457923c713e1"
    }
  ],
  "checks": [
    {
      "name": "encoder-tuned",
      "n": 1000,
      "width": 1000,
      "passed": true,
      "max_relative_l2": 1.1030241608978031e-07
    },
    {
      "name": "encoder-tuned",
      "n": 1000,
      "width": 500,
      "passed": true,
      "max_relative_l2": 1.0635559277716311e-07
    },
    {
      "name": "encoder-tuned",
      "n": 1000,
      "width": 250,
      "passed": true,
      "max_relative_l2": 1.0848046372302633e-07
    },
    {
      "name": "encoder-tuned",
      "n": 1000,
      "width": 60,
      "passed": true,
      "max_relative_l2": 9.542122114680751e-08
    },
    {
      "name": "encoder-tuned",
      "n": 1000,
      "width": 10,
      "passed": true,
      "max_relative_l2": 1.440240140482274e-07
    },
    {
      "name": "encoder-tuned",
      "n": 257,
      "width": 1000,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "encoder-tuned",
      "n": 257,
      "width": 500,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "encoder-tuned",
      "n": 257,
      "width": 250,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "encoder-tuned",
      "n": 257,
      "width": 60,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "encoder-tuned",
      "n": 257,
      "width": 10,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "decoder-tuned-reference",
      "type": "full_loss_all_parameter_gradients",
      "relative_l2": [
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0
      ],
      "passed": true
    },
    {
      "name": "encoder-tuned",
      "type": "full_loss_all_parameter_gradients",
      "relative_l2": [
        0.0,
        0.0,
        0.0001618572132429108,
        0.00010399276652606204,
        4.983359031029977e-05,
        1.9263330614194274e-05,
        6.1850428210163955e-06,
        2.04582511287299e-06,
        0.00012046255869790912,
        5.561411671806127e-05,
        2.1420248231152073e-05,
        7.41238955015433e-06,
        1.9763723457799642e-07,
        2.0019083990518993e-07,
        4.571763767557968e-08,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0,
        0.0
      ],
      "passed": true
    }
  ],
  "rows": [
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.16,
      "cuda_ms": 1183.898681640625,
      "wall_ms": 1184.0345190000007
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.21,
      "cuda_ms": 1184.0693359375,
      "wall_ms": 1184.1894979999986
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.22,
      "cuda_ms": 1183.8040771484375,
      "wall_ms": 1183.9332460000023
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.22,
      "cuda_ms": 1151.3057861328125,
      "wall_ms": 1151.5260880000042
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.23,
      "cuda_ms": 1150.8475341796875,
      "wall_ms": 1150.9517739999992
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.31,
      "cuda_ms": 1150.9100341796875,
      "wall_ms": 1151.0364309999943
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.41,
      "cuda_ms": 1184.21533203125,
      "wall_ms": 1184.3463379999976
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.44,
      "cuda_ms": 1184.0230712890625,
      "wall_ms": 1184.2061780000038
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17401,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.3800000000000003,
      "cuda_ms": 1184.206787109375,
      "wall_ms": 1184.3221740000017
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.3200000000000003,
      "cuda_ms": 1150.955078125,
      "wall_ms": 1151.0794229999988
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.45,
      "cuda_ms": 1151.0284423828125,
      "wall_ms": 1151.2003160000006
    },
    {
      "name": "encoder-tuned",
      "split": 17401,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.4800000000000004,
      "cuda_ms": 1150.7489013671875,
      "wall_ms": 1150.8538450000003
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.39,
      "cuda_ms": 1184.267822265625,
      "wall_ms": 1184.3602359999963
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.33,
      "cuda_ms": 1184.0059814453125,
      "wall_ms": 1184.2223129999993
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.34,
      "cuda_ms": 1183.6514892578125,
      "wall_ms": 1183.7448929999964
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.2500000000000004,
      "cuda_ms": 1150.9208984375,
      "wall_ms": 1151.0809809999998
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.3,
      "cuda_ms": 1151.0584716796875,
      "wall_ms": 1151.2153190000022
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.2600000000000002,
      "cuda_ms": 1150.793212890625,
      "wall_ms": 1150.960767000001
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.56,
      "cuda_ms": 1183.6414794921875,
      "wall_ms": 1183.7562449999978
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.4800000000000004,
      "cuda_ms": 1183.9749755859375,
      "wall_ms": 1184.0997049999942
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17402,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.45,
      "cuda_ms": 1184.07958984375,
      "wall_ms": 1184.2131169999989
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.53,
      "cuda_ms": 1150.8994140625,
      "wall_ms": 1151.0053790000043
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.33,
      "cuda_ms": 1151.0535888671875,
      "wall_ms": 1151.2052849999961
    },
    {
      "name": "encoder-tuned",
      "split": 17402,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.43,
      "cuda_ms": 1150.9810791015625,
      "wall_ms": 1151.101605000001
    }
  ],
  "profiles": [
    {
      "name": "decoder-tuned-reference",
      "initialise": {
        "cuda_ms": 2.1483519077301025,
        "wall_ms": 2.2034930000032205
      },
      "training": {
        "cuda_ms": 1178.935302734375,
        "wall_ms": 1178.994361000001
      },
      "predict": {
        "cuda_ms": 2.9265921115875244,
        "wall_ms": 3.0041809999943325
      },
      "full_energy": {
        "above_idle_j": 146.91062120719351,
        "calls": 17,
        "idle_before_w": 86.22012666392395,
        "idle_after_w": 86.30528533817056
      },
      "five_step_kernel_trace": [
        {
          "name": "_encoder_backward",
          "calls": 30,
          "total_us": 1110.7850000000008
        },
        {
          "name": "_encoder_forward",
          "calls": 30,
          "total_us": 900.9589999999985
        },
        {
          "name": "_decoder_backward",
          "calls": 35,
          "total_us": 876.2990000000007
        },
        {
          "name": "_decoder_forward",
          "calls": 35,
          "total_us": 313.89599999999973
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align1>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align1::Params)",
          "calls": 20,
          "total_us": 300.032999999999
        },
        {
          "name": "void at::native::(anonymous namespace)::multi_tensor_apply_kernel<at::native::(anonymous namespace)::FusedOptimizerTensorListMetadata<4>, at::native::(anonymous namespace)::FusedAdamMathFunctor<float, 4, (at::native::ADAM_MODE)0, false>, float const*, double, double, double, double, double, bool, float const*, float const*>(at::native::(anonymous namespace)::FusedOptimizerTensorListMetadata<4>, at::native::(anonymous namespace)::FusedAdamMathFunctor<float, 4, (at::native::ADAM_MODE)0, false>, float const*, double, double, double, double, double, bool, float const*, float const*)",
          "calls": 5,
          "total_us": 263.5269999999987
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_nn_align1>(cutlass_80_tensorop_s1688gemm_128x64_16x6_nn_align1::Params)",
          "calls": 15,
          "total_us": 224.14199999999732
        },
        {
          "name": "void cublasLt::splitKreduce_kernel<32, 16, int, float, float, float, float, false, float, float, float, true, false, false, false>(cublasLt::cublasSplitKParams<float>, float const*, float const*, float*, float*, float const*, float const*, float const*, float const*, float*, void*, long, float*, int*, float*, float*, float const*, float const*, float const*, float const*, float const*)",
          "calls": 65,
          "total_us": 213.25800000000208
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align1::Params)",
          "calls": 15,
          "total_us": 210.65999999999985
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_tn_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_tn_align1::Params)",
          "calls": 15,
          "total_us": 203.49099999999726
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x64_16x4_nt_align4>(cutlass_80_tensorop_s1688gemm_256x64_16x4_nt_align4::Params)",
          "calls": 5,
          "total_us": 200.7679999999998
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align1>(cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align1::Params)",
          "calls": 10,
          "total_us": 171.0540000000001
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nt_align4>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nt_align4::Params)",
          "calls": 5,
          "total_us": 169.26000000000113
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x128_32x3_tn_align4>(cutlass_80_tensorop_s1688gemm_256x128_32x3_tn_align4::Params)",
          "calls": 5,
          "total_us": 168.9390000000003
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align4>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align4::Params)",
          "calls": 5,
          "total_us": 162.56899999999928
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align1>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align1::Params)",
          "calls": 5,
          "total_us": 156.83600000000024
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align4::Params)",
          "calls": 5,
          "total_us": 140.3140000000003
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nn_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nn_align4::Params)",
          "calls": 5,
          "total_us": 133.65199999999936
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::FillFunctor<float>, std::array<char*, 1ul> >(int, at::native::FillFunctor<float>, std::array<char*, 1ul>)",
          "calls": 60,
          "total_us": 122.79499999999757
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align1::Params)",
          "calls": 10,
          "total_us": 112.58200000000033
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_tn_align1>(cutlass_80_tensorop_s1688gemm_128x128_32x3_tn_align1::Params)",
          "calls": 5,
          "total_us": 95.93300000000136
        },
        {
          "name": "memcpy32_post",
          "calls": 45,
          "total_us": 93.2469999999987
        },
        {
          "name": "void at::native::reduce_kernel<512, 1, at::native::ReduceOp<float, at::native::MeanOps<float, float, float, float>, unsigned int, float, 4, 4> >(at::native::ReduceOp<float, at::native::MeanOps<float, float, float, float>, unsigned int, float, 4, 4>)",
          "calls": 10,
          "total_us": 64.74500000000057
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x128_16x3_nn_align4>(cutlass_80_tensorop_s1688gemm_256x128_16x3_nn_align4::Params)",
          "calls": 5,
          "total_us": 58.88400000000024
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align4>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align4::Params)",
          "calls": 5,
          "total_us": 50.720000000000255
        },
        {
          "name": "ampere_sgemm_32x32_sliced1x4_nn",
          "calls": 5,
          "total_us": 50.08100000000036
        },
        {
          "name": "memcpy128",
          "calls": 10,
          "total_us": 49.1170000000011
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align4>(cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align4::Params)",
          "calls": 5,
          "total_us": 47.64699999999948
        },
        {
          "name": "void at::native::_scatter_gather_elementwise_kernel<128, 8, at::native::_cuda_scatter_gather_internal_kernel<true, at::native::OpaqueType<4>, long>::operator()<at::native::TensorAssign>(at::TensorIterator&, long, long, long, at::native::TensorAssign const&)::{lambda(int)#1}>(int, at::native::_cuda_scatter_gather_internal_kernel<true, at::native::OpaqueType<4>, long>::operator()<at::native::TensorAssign>(at::TensorIterator&, long, long, long, at::native::TensorAssign const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 44.475000000000364
        },
        {
          "name": "ampere_sgemm_64x32_sliced1x4_nt",
          "calls": 5,
          "total_us": 40.50500000000011
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align4>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align4::Params)",
          "calls": 5,
          "total_us": 37.914000000000215
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align4::Params)",
          "calls": 5,
          "total_us": 34.580999999999904
        },
        {
          "name": "void at::native::vectorized_gather_kernel<16, long>(char*, char*, long*, int, long, long, long, long, bool)",
          "calls": 10,
          "total_us": 33.87699999999745
        },
        {
          "name": "void at::native::index_elementwise_kernel<128, 4, at::native::gpu_index_kernel<at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1}>(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>, at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1} const&, bool)::{lambda(int)#1}>(long, at::native::gpu_index_kernel<at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1}>(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>, at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1} const&, bool)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 31.15600000000086
        },
        {
          "name": "void at::native::(anonymous namespace)::searchsorted_cuda_kernel<float, long>(long*, float const*, float const*, long const*, long, long, long, bool, bool)",
          "calls": 5,
          "total_us": 30.963999999999942
        },
        {
          "name": "ampere_sgemm_32x128_tn",
          "calls": 5,
          "total_us": 30.80299999999852
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align1>(cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align1::Params)",
          "calls": 5,
          "total_us": 29.20100000000184
        },
        {
          "name": "void at::native::(anonymous namespace)::indexSelectSmallIndex<float, long, unsigned int, 1, 1, -2>(at::cuda::detail::TensorInfo<float, unsigned int>, at::cuda::detail::TensorInfo<float const, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, int, int, unsigned int, long)",
          "calls": 10,
          "total_us": 29.140000000000327
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::CUDAFunctor_add<float>, std::array<char*, 3ul> >(int, at::native::CUDAFunctor_add<float>, std::array<char*, 3ul>)",
          "calls": 15,
          "total_us": 28.947000000001253
        },
        {
          "name": "void at::native::(anonymous namespace)::distribution_elementwise_grid_stride_kernel<float, 4, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1}>(long, at::PhiloxCudaState, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1})",
          "calls": 5,
          "total_us": 27.471999999998616
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::AUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul> >(int, at::native::AUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul>)",
          "calls": 15,
          "total_us": 26.161999999999125
        },
        {
          "name": "ampere_sgemm_128x32_nn",
          "calls": 5,
          "total_us": 25.80699999999956
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 10,
          "total_us": 22.95799999999963
        },
        {
          "name": "void at::native::reduce_kernel<512, 1, at::native::ReduceOp<float, at::native::func_wrapper_t<float, at::native::sum_functor<float, float, float>::operator()(at::TensorIterator&)::{lambda(float, float)#1}>, unsigned int, float, 4, 4> >(at::native::ReduceOp<float, at::native::func_wrapper_t<float, at::native::sum_functor<float, float, float>::operator()(at::TensorIterator&)::{lambda(float, float)#1}>, unsigned int, float, 4, 4>)",
          "calls": 5,
          "total_us": 20.971000000000686
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::FillFunctor<long>, std::array<char*, 1ul> >(int, at::native::FillFunctor<long>, std::array<char*, 1ul>)",
          "calls": 10,
          "total_us": 17.770000000001573
        },
        {
          "name": "void at_cuda_detail::cub::detail::scan::DeviceScanKernel<at_cuda_detail::cub::detail::scan::policy_hub<float, float, float, unsigned int, std::plus<float> >::Policy1000, float const*, float*, at_cuda_detail::cub::ScanTileState<float, true>, std::plus<float>, at_cuda_detail::cub::NullType, unsigned int, float, false, at_cuda_detail::cub::NullType>(float const*, float*, at_cuda_detail::cub::ScanTileState<float, true>, int, std::plus<float>, at_cuda_detail::cub::NullType, unsigned int)",
          "calls": 5,
          "total_us": 16.715000000000146
        },
        {
          "name": "void at::native::(anonymous namespace)::nll_loss_backward_no_reduce_cuda_kernel<float, long>(int, long const*, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float const, 0ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<1ul, long>, float const, 1ul, torch::headeronly::DefaultPtrTraits, long>, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float, 1ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<2ul, long>, float, 2ul, torch::headeronly::DefaultPtrTraits, long>, float const*, long, long)",
          "calls": 5,
          "total_us": 14.24999999999909
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1}>(at::TensorIteratorBase&, at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1} const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1}>(at::TensorIteratorBase&, at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1} const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 13.578000000000202
        },
        {
          "name": "void at::native::(anonymous namespace)::nll_loss_forward_no_reduce_cuda_kernel<float, long>(long, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float, 1ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<2ul, long>, float, 2ul, torch::headeronly::DefaultPtrTraits, long>, long const*, float*, float const*, long, long)",
          "calls": 5,
          "total_us": 13.096999999998616
        },
        {
          "name": "void at::native::(anonymous namespace)::indexSelectSmallIndex<long, long, unsigned int, 2, 2, -2>(at::cuda::detail::TensorInfo<long, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, int, int, unsigned int, long)",
          "calls": 5,
          "total_us": 12.80799999999931
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 12.393000000000029
        },
        {
          "name": "void at::native::(anonymous namespace)::distribution_elementwise_grid_stride_kernel<float, 4, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1}>(long, at::PhiloxCudaState, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1})",
          "calls": 5,
          "total_us": 12.23400000000106
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_forward<float, float, float, 4, true, false, 32>(float*, float const*, int, int, int, bool const*, int, bool)",
          "calls": 5,
          "total_us": 11.847000000000662
        },
        {
          "name": "void at::native::(anonymous namespace)::multi_tensor_apply_kernel<at::native::(anonymous namespace)::TensorListMetadata<1>, at::native::(anonymous namespace)::BinaryOpScalarFunctor<float, 1, 1, 0>, std::plus<float>, float>(at::native::(anonymous namespace)::TensorListMetadata<1>, at::native::(anonymous namespace)::BinaryOpScalarFunctor<float, 1, 1, 0>, std::plus<float>, float)",
          "calls": 5,
          "total_us": 11.6550000000002
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::mse_kernel_cuda(at::TensorIteratorBase&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float)#1}, std::array<char*, 3ul> >(int, at::native::mse_kernel_cuda(at::TensorIteratorBase&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float)#1}, std::array<char*, 3ul>)",
          "calls": 5,
          "total_us": 11.078000000000884
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::CUDAFunctor_add<float> >(at::TensorIteratorBase&, at::native::CUDAFunctor_add<float> const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::CUDAFunctor_add<float> >(at::TensorIteratorBase&, at::native::CUDAFunctor_add<float> const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 11.044999999999845
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::CUDAFunctorOnOther_add<float>, std::array<char*, 2ul> >(int, at::native::CUDAFunctorOnOther_add<float>, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 10.663000000001148
        },
        {
          "name": "void at::native::(anonymous namespace)::CatArrayBatchedCopy_vectorized<at::native::(anonymous namespace)::OpaqueType<4u>, unsigned int, 1, 128, 1, 16, 4>(char*, at::native::(anonymous namespace)::CatArrInputTensorMetadata<at::native::(anonymous namespace)::OpaqueType<4u>, unsigned int, 128, 1>, at::native::(anonymous namespace)::TensorSizeStride<unsigned int, 4u>, int, unsigned int)",
          "calls": 5,
          "total_us": 10.632000000000517
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul> >(int, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 10.566999999999098
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_forward<float, float, float, 4, false, false, 32>(float*, float const*, int, int, int, bool const*, int, bool)",
          "calls": 5,
          "total_us": 10.278000000000247
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::(anonymous namespace)::launch_clamp_scalar(at::TensorIteratorBase&, c10::Scalar, c10::Scalar, at::native::detail::ClampLimits)::{lambda()#1}::operator()() const::{lambda()#4}::operator()() const::{lambda(long)#1}, std::array<char*, 2ul> >(int, at::native::(anonymous namespace)::launch_clamp_scalar(at::TensorIteratorBase&, c10::Scalar, c10::Scalar, at::native::detail::ClampLimits)::{lambda()#1}::operator()() const::{lambda()#4}::operator()() const::{lambda(long)#1}, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 9.830999999999904
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_backward<float, float, float, 4, true, false, 32>(float*, float const*, float const*, int, int, int, bool const*)",
          "calls": 5,
          "total_us": 9.766000000000531
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 9.510000000000673
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::CUDAFunctorOnSelf_add<long>, std::array<char*, 2ul> >(int, at::native::CUDAFunctorOnSelf_add<long>, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 9.350000000000819
        },
        {
          "name": "Memset (Unknown)",
          "calls": 5,
          "total_us": 9.222000000000662
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 3ul> >(int, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 3ul>)",
          "calls": 5,
          "total_us": 9.190000000001419
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_backward<float, float, float, 4, false, false, 32>(float*, float const*, float const*, int, int, int, bool const*)",
          "calls": 5,
          "total_us": 9.128000000001066
        },
        {
          "name": "void at_cuda_detail::cub::detail::scan::DeviceScanInitKernel<at_cuda_detail::cub::ScanTileState<float, true> >(at_cuda_detail::cub::ScanTileState<float, true>, int)",
          "calls": 5,
          "total_us": 6.4370000000001255
        }
      ],
      "encoder_backward_resources": [
        {
          "width": 1000,
          "columns": 2,
          "registers": 64,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 64
        },
        {
          "width": 500,
          "columns": 2,
          "registers": 64,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 64
        },
        {
          "width": 250,
          "columns": 2,
          "registers": 64,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 64
        },
        {
          "width": 60,
          "columns": 2,
          "registers": 64,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 64
        },
        {
          "width": 10,
          "columns": 2,
          "registers": 64,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 64
        }
      ],
      "training_plus_initialise_energy": {
        "above_idle_j": 152.85636798946842,
        "calls": 4,
        "idle_before_w": 70.01432942318006,
        "idle_after_w": 85.80134737506764
      },
      "predict_energy": {
        "above_idle_j": 0.38409427568056725,
        "calls": 1548,
        "idle_before_w": 69.65352955109667,
        "idle_after_w": 86.01021217671051
      }
    },
    {
      "name": "encoder-tuned",
      "initialise": {
        "cuda_ms": 2.8650879859924316,
        "wall_ms": 2.9229790000044886
      },
      "training": {
        "cuda_ms": 1149.2357177734375,
        "wall_ms": 1149.3142369999987
      },
      "predict": {
        "cuda_ms": 3.4410879611968994,
        "wall_ms": 3.5211579999980813
      },
      "full_energy": {
        "above_idle_j": 139.33628769719067,
        "calls": 18,
        "idle_before_w": 85.76382416700307,
        "idle_after_w": 88.76967704520915
      },
      "five_step_kernel_trace": [
        {
          "name": "_encoder_backward",
          "calls": 30,
          "total_us": 905.2989999999982
        },
        {
          "name": "_encoder_forward",
          "calls": 30,
          "total_us": 899.5939999999966
        },
        {
          "name": "_decoder_backward",
          "calls": 35,
          "total_us": 876.2799999999975
        },
        {
          "name": "_decoder_forward",
          "calls": 35,
          "total_us": 311.776000000001
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align1>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align1::Params)",
          "calls": 20,
          "total_us": 300.7600000000002
        },
        {
          "name": "void at::native::(anonymous namespace)::multi_tensor_apply_kernel<at::native::(anonymous namespace)::FusedOptimizerTensorListMetadata<4>, at::native::(anonymous namespace)::FusedAdamMathFunctor<float, 4, (at::native::ADAM_MODE)0, false>, float const*, double, double, double, double, double, bool, float const*, float const*>(at::native::(anonymous namespace)::FusedOptimizerTensorListMetadata<4>, at::native::(anonymous namespace)::FusedAdamMathFunctor<float, 4, (at::native::ADAM_MODE)0, false>, float const*, double, double, double, double, double, bool, float const*, float const*)",
          "calls": 5,
          "total_us": 262.0450000000001
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_nn_align1>(cutlass_80_tensorop_s1688gemm_128x64_16x6_nn_align1::Params)",
          "calls": 15,
          "total_us": 225.2179999999978
        },
        {
          "name": "void cublasLt::splitKreduce_kernel<32, 16, int, float, float, float, float, false, float, float, float, true, false, false, false>(cublasLt::cublasSplitKParams<float>, float const*, float const*, float*, float*, float const*, float const*, float const*, float const*, float*, void*, long, float*, int*, float*, float*, float const*, float const*, float const*, float const*, float const*)",
          "calls": 65,
          "total_us": 211.76600000000212
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align1::Params)",
          "calls": 15,
          "total_us": 209.74799999999777
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_tn_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_tn_align1::Params)",
          "calls": 15,
          "total_us": 203.73199999999997
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x64_16x4_nt_align4>(cutlass_80_tensorop_s1688gemm_256x64_16x4_nt_align4::Params)",
          "calls": 5,
          "total_us": 199.8229999999994
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align1>(cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align1::Params)",
          "calls": 10,
          "total_us": 172.44400000000314
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nt_align4>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nt_align4::Params)",
          "calls": 5,
          "total_us": 169.78700000000208
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x128_32x3_tn_align4>(cutlass_80_tensorop_s1688gemm_256x128_32x3_tn_align4::Params)",
          "calls": 5,
          "total_us": 168.82399999999825
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align4>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align4::Params)",
          "calls": 5,
          "total_us": 162.99700000000053
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align1>(cutlass_80_tensorop_s1688gemm_128x128_32x3_nn_align1::Params)",
          "calls": 5,
          "total_us": 157.17000000000007
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align4::Params)",
          "calls": 5,
          "total_us": 140.7100000000005
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nn_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nn_align4::Params)",
          "calls": 5,
          "total_us": 132.6710000000012
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::FillFunctor<float>, std::array<char*, 1ul> >(int, at::native::FillFunctor<float>, std::array<char*, 1ul>)",
          "calls": 60,
          "total_us": 123.63799999999765
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align1>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align1::Params)",
          "calls": 10,
          "total_us": 113.83999999999946
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x128_32x3_tn_align1>(cutlass_80_tensorop_s1688gemm_128x128_32x3_tn_align1::Params)",
          "calls": 5,
          "total_us": 96.54999999999973
        },
        {
          "name": "memcpy32_post",
          "calls": 45,
          "total_us": 93.89300000000117
        },
        {
          "name": "void at::native::reduce_kernel<512, 1, at::native::ReduceOp<float, at::native::MeanOps<float, float, float, float>, unsigned int, float, 4, 4> >(at::native::ReduceOp<float, at::native::MeanOps<float, float, float, float>, unsigned int, float, 4, 4>)",
          "calls": 10,
          "total_us": 66.03099999999995
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_256x128_16x3_nn_align4>(cutlass_80_tensorop_s1688gemm_256x128_16x3_nn_align4::Params)",
          "calls": 5,
          "total_us": 58.21600000000035
        },
        {
          "name": "ampere_sgemm_32x32_sliced1x4_nn",
          "calls": 5,
          "total_us": 50.81899999999905
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align4>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nt_align4::Params)",
          "calls": 5,
          "total_us": 50.788000000000466
        },
        {
          "name": "memcpy128",
          "calls": 10,
          "total_us": 48.57999999999811
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align4>(cutlass_80_tensorop_s1688gemm_128x64_16x6_tn_align4::Params)",
          "calls": 5,
          "total_us": 48.32299999999941
        },
        {
          "name": "void at::native::_scatter_gather_elementwise_kernel<128, 8, at::native::_cuda_scatter_gather_internal_kernel<true, at::native::OpaqueType<4>, long>::operator()<at::native::TensorAssign>(at::TensorIterator&, long, long, long, at::native::TensorAssign const&)::{lambda(int)#1}>(int, at::native::_cuda_scatter_gather_internal_kernel<true, at::native::OpaqueType<4>, long>::operator()<at::native::TensorAssign>(at::TensorIterator&, long, long, long, at::native::TensorAssign const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 43.80699999999888
        },
        {
          "name": "ampere_sgemm_64x32_sliced1x4_nt",
          "calls": 5,
          "total_us": 39.74100000000135
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align4>(cutlass_80_tensorop_s1688gemm_64x64_32x6_nn_align4::Params)",
          "calls": 5,
          "total_us": 37.53100000000063
        },
        {
          "name": "void at::native::vectorized_gather_kernel<16, long>(char*, char*, long*, int, long, long, long, long, bool)",
          "calls": 10,
          "total_us": 34.582000000001244
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align4>(cutlass_80_tensorop_s1688gemm_64x64_16x6_nt_align4::Params)",
          "calls": 5,
          "total_us": 34.392999999999574
        },
        {
          "name": "void at::native::index_elementwise_kernel<128, 4, at::native::gpu_index_kernel<at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1}>(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>, at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1} const&, bool)::{lambda(int)#1}>(long, at::native::gpu_index_kernel<at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1}>(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>, at::native::index_kernel_impl<at::native::OpaqueType<8> >(at::TensorIteratorBase&, c10::ArrayRef<long>, c10::ArrayRef<long>)::{lambda(char*, char const*, long)#1} const&, bool)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 31.864000000000715
        },
        {
          "name": "void at::native::(anonymous namespace)::searchsorted_cuda_kernel<float, long>(long*, float const*, float const*, long const*, long, long, long, bool, bool)",
          "calls": 5,
          "total_us": 30.87100000000055
        },
        {
          "name": "ampere_sgemm_32x128_tn",
          "calls": 5,
          "total_us": 30.837999999999738
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::CUDAFunctor_add<float>, std::array<char*, 3ul> >(int, at::native::CUDAFunctor_add<float>, std::array<char*, 3ul>)",
          "calls": 15,
          "total_us": 29.08100000000013
        },
        {
          "name": "void cutlass::Kernel2<cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align1>(cutlass_80_tensorop_s1688gemm_64x64_16x6_tn_align1::Params)",
          "calls": 5,
          "total_us": 29.07799999999952
        },
        {
          "name": "void at::native::(anonymous namespace)::indexSelectSmallIndex<float, long, unsigned int, 1, 1, -2>(at::cuda::detail::TensorInfo<float, unsigned int>, at::cuda::detail::TensorInfo<float const, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, int, int, unsigned int, long)",
          "calls": 10,
          "total_us": 28.372999999998683
        },
        {
          "name": "void at::native::(anonymous namespace)::distribution_elementwise_grid_stride_kernel<float, 4, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1}>(long, at::PhiloxCudaState, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::normal_kernel<at::CUDAGeneratorImpl*>(at::TensorBase const&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1})",
          "calls": 5,
          "total_us": 27.059000000000424
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::AUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul> >(int, at::native::AUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul>)",
          "calls": 15,
          "total_us": 26.353000000000748
        },
        {
          "name": "ampere_sgemm_128x32_nn",
          "calls": 5,
          "total_us": 25.33100000000036
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 10,
          "total_us": 23.184999999999263
        },
        {
          "name": "void at::native::reduce_kernel<512, 1, at::native::ReduceOp<float, at::native::func_wrapper_t<float, at::native::sum_functor<float, float, float>::operator()(at::TensorIterator&)::{lambda(float, float)#1}>, unsigned int, float, 4, 4> >(at::native::ReduceOp<float, at::native::func_wrapper_t<float, at::native::sum_functor<float, float, float>::operator()(at::TensorIterator&)::{lambda(float, float)#1}>, unsigned int, float, 4, 4>)",
          "calls": 5,
          "total_us": 20.75
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::FillFunctor<long>, std::array<char*, 1ul> >(int, at::native::FillFunctor<long>, std::array<char*, 1ul>)",
          "calls": 10,
          "total_us": 18.38099999999963
        },
        {
          "name": "void at_cuda_detail::cub::detail::scan::DeviceScanKernel<at_cuda_detail::cub::detail::scan::policy_hub<float, float, float, unsigned int, std::plus<float> >::Policy1000, float const*, float*, at_cuda_detail::cub::ScanTileState<float, true>, std::plus<float>, at_cuda_detail::cub::NullType, unsigned int, float, false, at_cuda_detail::cub::NullType>(float const*, float*, at_cuda_detail::cub::ScanTileState<float, true>, int, std::plus<float>, at_cuda_detail::cub::NullType, unsigned int)",
          "calls": 5,
          "total_us": 17.033999999999878
        },
        {
          "name": "void at::native::(anonymous namespace)::nll_loss_backward_no_reduce_cuda_kernel<float, long>(int, long const*, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float const, 0ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<1ul, long>, float const, 1ul, torch::headeronly::DefaultPtrTraits, long>, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float, 1ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<2ul, long>, float, 2ul, torch::headeronly::DefaultPtrTraits, long>, float const*, long, long)",
          "calls": 5,
          "total_us": 14.249999999997954
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1}>(at::TensorIteratorBase&, at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1} const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1}>(at::TensorIteratorBase&, at::native::mse_backward_cuda_kernel(at::TensorIterator&, c10::Scalar const&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float, float)#1} const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 13.80200000000059
        },
        {
          "name": "void at::native::(anonymous namespace)::nll_loss_forward_no_reduce_cuda_kernel<float, long>(long, torch::headeronly::detail::GenericPackedTensorAccessor<torch::headeronly::detail::TensorAccessor<c10::ArrayRef<long>, float, 1ul, torch::headeronly::DefaultPtrTraits, long>, at::detail::IndexBoundsCheck<2ul, long>, float, 2ul, torch::headeronly::DefaultPtrTraits, long>, long const*, float*, float const*, long, long)",
          "calls": 5,
          "total_us": 13.38599999999974
        },
        {
          "name": "void at::native::(anonymous namespace)::indexSelectSmallIndex<long, long, unsigned int, 2, 2, -2>(at::cuda::detail::TensorInfo<long, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, at::cuda::detail::TensorInfo<long const, unsigned int>, int, int, unsigned int, long)",
          "calls": 5,
          "total_us": 12.809999999999718
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > >(at::TensorIteratorBase&, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::DivFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 12.423000000000002
        },
        {
          "name": "void at::native::(anonymous namespace)::distribution_elementwise_grid_stride_kernel<float, 4, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1}>(long, at::PhiloxCudaState, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::(anonymous namespace)::distribution_nullary_kernel<float, float, float4, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2}, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_and_transform<float, float, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1}>(at::TensorIteratorBase&, at::CUDAGeneratorImpl*, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(curandStatePhilox4_32_10*)#2} const&, at::native::templates::cuda::uniform_kernel<at::CUDAGeneratorImpl*>(at::TensorIteratorBase&, double, double, at::CUDAGeneratorImpl*)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float)#1})::{lambda(int, float)#1})",
          "calls": 5,
          "total_us": 12.103999999999587
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_forward<float, float, float, 4, true, false, 32>(float*, float const*, int, int, int, bool const*, int, bool)",
          "calls": 5,
          "total_us": 11.942999999999756
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::mse_kernel_cuda(at::TensorIteratorBase&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float)#1}, std::array<char*, 3ul> >(int, at::native::mse_kernel_cuda(at::TensorIteratorBase&)::{lambda()#1}::operator()() const::{lambda()#2}::operator()() const::{lambda(float, float)#1}, std::array<char*, 3ul>)",
          "calls": 5,
          "total_us": 11.59299999999962
        },
        {
          "name": "void at::native::(anonymous namespace)::multi_tensor_apply_kernel<at::native::(anonymous namespace)::TensorListMetadata<1>, at::native::(anonymous namespace)::BinaryOpScalarFunctor<float, 1, 1, 0>, std::plus<float>, float>(at::native::(anonymous namespace)::TensorListMetadata<1>, at::native::(anonymous namespace)::BinaryOpScalarFunctor<float, 1, 1, 0>, std::plus<float>, float)",
          "calls": 5,
          "total_us": 11.400999999998476
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul> >(int, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 11.111000000000786
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::CUDAFunctor_add<float> >(at::TensorIteratorBase&, at::native::CUDAFunctor_add<float> const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::CUDAFunctor_add<float> >(at::TensorIteratorBase&, at::native::CUDAFunctor_add<float> const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 10.697000000000799
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_forward<float, float, float, 4, false, false, 32>(float*, float const*, int, int, int, bool const*, int, bool)",
          "calls": 5,
          "total_us": 10.534999999999627
        },
        {
          "name": "void at::native::(anonymous namespace)::CatArrayBatchedCopy_vectorized<at::native::(anonymous namespace)::OpaqueType<4u>, unsigned int, 1, 128, 1, 16, 4>(char*, at::native::(anonymous namespace)::CatArrInputTensorMetadata<at::native::(anonymous namespace)::OpaqueType<4u>, unsigned int, 128, 1>, at::native::(anonymous namespace)::TensorSizeStride<unsigned int, 4u>, int, unsigned int)",
          "calls": 5,
          "total_us": 10.506000000001222
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::CUDAFunctorOnOther_add<float>, std::array<char*, 2ul> >(int, at::native::CUDAFunctorOnOther_add<float>, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 10.503000000000156
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_backward<float, float, float, 4, true, false, 32>(float*, float const*, float const*, int, int, int, bool const*)",
          "calls": 5,
          "total_us": 9.734000000000151
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::(anonymous namespace)::launch_clamp_scalar(at::TensorIteratorBase&, c10::Scalar, c10::Scalar, at::native::detail::ClampLimits)::{lambda()#1}::operator()() const::{lambda()#4}::operator()() const::{lambda(long)#1}, std::array<char*, 2ul> >(int, at::native::(anonymous namespace)::launch_clamp_scalar(at::TensorIteratorBase&, c10::Scalar, c10::Scalar, at::native::detail::ClampLimits)::{lambda()#1}::operator()() const::{lambda()#4}::operator()() const::{lambda(long)#1}, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 9.605999999999995
        },
        {
          "name": "void at::native::elementwise_kernel<128, 2, at::native::gpu_kernel_impl_nocast<at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1}>(int, at::native::gpu_kernel_impl_nocast<at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > >(at::TensorIteratorBase&, at::native::BUnaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> > const&)::{lambda(int)#1})",
          "calls": 5,
          "total_us": 9.544000000000551
        },
        {
          "name": "Memset (Unknown)",
          "calls": 5,
          "total_us": 9.44800000000032
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<4, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 3ul> >(int, at::native::BinaryFunctor<float, float, float, at::native::binary_internal::MulFunctor<float> >, std::array<char*, 3ul>)",
          "calls": 5,
          "total_us": 9.221000000000004
        },
        {
          "name": "void (anonymous namespace)::softmax_warp_backward<float, float, float, 4, false, false, 32>(float*, float const*, float const*, int, int, int, bool const*)",
          "calls": 5,
          "total_us": 9.15900000000056
        },
        {
          "name": "void at::native::vectorized_elementwise_kernel<2, at::native::CUDAFunctorOnSelf_add<long>, std::array<char*, 2ul> >(int, at::native::CUDAFunctorOnSelf_add<long>, std::array<char*, 2ul>)",
          "calls": 5,
          "total_us": 8.389999999998963
        },
        {
          "name": "void at_cuda_detail::cub::detail::scan::DeviceScanInitKernel<at_cuda_detail::cub::ScanTileState<float, true> >(at_cuda_detail::cub::ScanTileState<float, true>, int)",
          "calls": 5,
          "total_us": 6.501000000000204
        }
      ],
      "encoder_backward_resources": [
        {
          "width": 1000,
          "columns": 4,
          "registers": 64,
          "warps": 16,
          "spills": 0,
          "shared_bytes": 256
        },
        {
          "width": 500,
          "columns": 4,
          "registers": 64,
          "warps": 16,
          "spills": 0,
          "shared_bytes": 256
        },
        {
          "width": 250,
          "columns": 4,
          "registers": 64,
          "warps": 16,
          "spills": 0,
          "shared_bytes": 256
        },
        {
          "width": 60,
          "columns": 1,
          "registers": 32,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 32
        },
        {
          "width": 10,
          "columns": 1,
          "registers": 40,
          "warps": 8,
          "spills": 0,
          "shared_bytes": 32
        }
      ],
      "training_plus_initialise_energy": {
        "above_idle_j": 146.2363137031025,
        "calls": 4,
        "idle_before_w": 70.6127775204419,
        "idle_after_w": 87.85201245399712
      },
      "predict_energy": {
        "above_idle_j": 0.3890338393743506,
        "calls": 1557,
        "idle_before_w": 70.61061105313587,
        "idle_after_w": 87.87095609849298
      }
    }
  ],
  "official": false,
  "full_energy_window_s": 20.0,
  "gpu": "NVIDIA A100-SXM4-80GB",
  "device": {
    "name": "NVIDIA A100-SXM4-80GB",
    "uuid": "GPU-f320757f-acf6-8d35-712d-1cce7855872c",
    "vbios": "92.00.9E.00.01",
    "driver": "580.95.05",
    "power_limit_w": 400.0
  },
  "elapsed_s": 130.85690439200002,
  "app_id": "ap-YpBnBwFUE6jBkp6nKr3Q4W",
  "launch_elapsed_s": 142.70001966698328
}
