{
  "configs": [
    {
      "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"
    },
    {
      "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"
    }
  ],
  "checks": [
    {
      "name": "decoder-tuned-reference",
      "n": 1000,
      "width": 1000,
      "passed": true,
      "max_relative_l2": 1.1030241608978031e-07
    },
    {
      "name": "decoder-tuned-reference",
      "n": 1000,
      "width": 500,
      "passed": true,
      "max_relative_l2": 1.0635559277716311e-07
    },
    {
      "name": "decoder-tuned-reference",
      "n": 1000,
      "width": 250,
      "passed": true,
      "max_relative_l2": 1.0848045661759897e-07
    },
    {
      "name": "decoder-tuned-reference",
      "n": 1000,
      "width": 60,
      "passed": true,
      "max_relative_l2": 9.542122114680751e-08
    },
    {
      "name": "decoder-tuned-reference",
      "n": 1000,
      "width": 10,
      "passed": true,
      "max_relative_l2": 1.440240140482274e-07
    },
    {
      "name": "decoder-tuned-reference",
      "n": 257,
      "width": 1000,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "decoder-tuned-reference",
      "n": 257,
      "width": 500,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "decoder-tuned-reference",
      "n": 257,
      "width": 250,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "decoder-tuned-reference",
      "n": 257,
      "width": 60,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "decoder-tuned-reference",
      "n": 257,
      "width": 10,
      "passed": true,
      "max_relative_l2": 0.0
    },
    {
      "name": "encoder-tuned",
      "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": "decoder-tuned-reference",
      "type": "full_loss_all_parameter_gradients",
      "relative_l2": [
        0.0,
        0.0,
        0.0001517666387371719,
        9.079906158149242e-05,
        4.147819709032774e-05,
        1.535824776510708e-05,
        5.7101960919681005e-06,
        1.497974039921246e-06,
        0.0001094975377782248,
        5.243884152150713e-05,
        1.9222066839574836e-05,
        8.350662028533407e-06,
        2.2257458454078005e-07,
        1.1820296919040629e-07,
        5.006213754654709e-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": "encoder-tuned",
      "split": 17403,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.31,
      "cuda_ms": 1159.5172119140625,
      "wall_ms": 1159.5958439999984
    },
    {
      "name": "encoder-tuned",
      "split": 17403,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.39,
      "cuda_ms": 1159.759033203125,
      "wall_ms": 1159.8543559999966
    },
    {
      "name": "encoder-tuned",
      "split": 17403,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.41,
      "cuda_ms": 1159.53369140625,
      "wall_ms": 1159.591356
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.56,
      "cuda_ms": 1192.9241943359375,
      "wall_ms": 1193.074523
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.46,
      "cuda_ms": 1193.6981201171875,
      "wall_ms": 1193.8291290000009
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.4,
      "cuda_ms": 1193.6019287109375,
      "wall_ms": 1193.7086830000005
    },
    {
      "name": "encoder-tuned",
      "split": 17403,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.56,
      "cuda_ms": 1159.5089111328125,
      "wall_ms": 1159.5865250000018
    },
    {
      "name": "encoder-tuned",
      "split": 17403,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.64,
      "cuda_ms": 1159.797607421875,
      "wall_ms": 1159.8910149999995
    },
    {
      "name": "encoder-tuned",
      "split": 17403,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.54,
      "cuda_ms": 1160.0693359375,
      "wall_ms": 1160.195787999996
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.73,
      "cuda_ms": 1193.4219970703125,
      "wall_ms": 1193.5166129999998
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.69,
      "cuda_ms": 1193.3187255859375,
      "wall_ms": 1193.4121480000001
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17403,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.64,
      "cuda_ms": 1193.703125,
      "wall_ms": 1193.8084999999958
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.69,
      "cuda_ms": 1159.666015625,
      "wall_ms": 1159.745848
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.62,
      "cuda_ms": 1159.835693359375,
      "wall_ms": 1159.933009999996
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.62,
      "cuda_ms": 1159.856201171875,
      "wall_ms": 1159.9588289999972
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 0,
      "repeat": 0,
      "error_pct": 2.4800000000000004,
      "cuda_ms": 1193.58740234375,
      "wall_ms": 1193.7085550000006
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 0,
      "repeat": 1,
      "error_pct": 2.62,
      "cuda_ms": 1193.55859375,
      "wall_ms": 1193.7082489999966
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 0,
      "repeat": 2,
      "error_pct": 2.5900000000000003,
      "cuda_ms": 1193.4632568359375,
      "wall_ms": 1193.574955999999
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.5700000000000003,
      "cuda_ms": 1159.5064697265625,
      "wall_ms": 1159.5795369999992
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.6100000000000003,
      "cuda_ms": 1160.076904296875,
      "wall_ms": 1160.1853859999949
    },
    {
      "name": "encoder-tuned",
      "split": 17404,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.63,
      "cuda_ms": 1159.9698486328125,
      "wall_ms": 1160.078771000002
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 1,
      "repeat": 0,
      "error_pct": 2.67,
      "cuda_ms": 1193.6083984375,
      "wall_ms": 1193.6890930000031
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 1,
      "repeat": 1,
      "error_pct": 2.67,
      "cuda_ms": 1193.5628662109375,
      "wall_ms": 1193.6506190000032
    },
    {
      "name": "decoder-tuned-reference",
      "split": 17404,
      "draw": 1,
      "repeat": 2,
      "error_pct": 2.7100000000000004,
      "cuda_ms": 1193.72412109375,
      "wall_ms": 1193.818598
    }
  ],
  "profiles": [
    {
      "name": "encoder-tuned",
      "initialise": {
        "cuda_ms": 1.9445760250091553,
        "wall_ms": 1.9900980000002733
      },
      "training": {
        "cuda_ms": 1154.9716796875,
        "wall_ms": 1155.0160379999993
      },
      "predict": {
        "cuda_ms": 2.5661439895629883,
        "wall_ms": 2.599766000003001
      },
      "full_energy": {
        "above_idle_j": 143.87509502345603,
        "calls": 18,
        "idle_before_w": 73.67634025842221,
        "idle_after_w": 76.27127493321554
      },
      "five_step_kernel_trace": [
        {
          "name": "_encoder_forward",
          "calls": 30,
          "total_us": 934.1399999999996
        },
        {
          "name": "_encoder_backward",
          "calls": 30,
          "total_us": 908.7730000000015
        },
        {
          "name": "_decoder_backward",
          "calls": 35,
          "total_us": 875.6760000000017
        },
        {
          "name": "_decoder_forward",
          "calls": 35,
          "total_us": 318.7349999999983
        },
        {
          "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": 301.9150000000004
        },
        {
          "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": 266.8119999999981
        },
        {
          "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.16799999999807
        },
        {
          "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.76000000000022
        },
        {
          "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": 211.3939999999975
        },
        {
          "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": 205.23999999999955
        },
        {
          "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.0209999999995
        },
        {
          "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.05399999999872
        },
        {
          "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": 171.60400000000163
        },
        {
          "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": 169.42700000000013
        },
        {
          "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": 164.39899999999966
        },
        {
          "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.99800000000073
        },
        {
          "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.6929999999993
        },
        {
          "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": 135.59999999999968
        },
        {
          "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.43299999999931
        },
        {
          "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.33699999999953
        },
        {
          "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.91000000000167
        },
        {
          "name": "memcpy32_post",
          "calls": 45,
          "total_us": 93.82999999999674
        },
        {
          "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.80400000000031
        },
        {
          "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": 59.42200000000071
        },
        {
          "name": "memcpy128",
          "calls": 10,
          "total_us": 51.34700000000157
        },
        {
          "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": 51.253999999999905
        },
        {
          "name": "ampere_sgemm_32x32_sliced1x4_nn",
          "calls": 5,
          "total_us": 50.26199999999858
        },
        {
          "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.63500000000067
        },
        {
          "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": 45.264000000000124
        },
        {
          "name": "ampere_sgemm_64x32_sliced1x4_nt",
          "calls": 5,
          "total_us": 40.42699999999968
        },
        {
          "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.92800000000011
        },
        {
          "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": 35.0120000000004
        },
        {
          "name": "void at::native::vectorized_gather_kernel<16, long>(char*, char*, long*, int, long, long, long, long, bool)",
          "calls": 10,
          "total_us": 34.8199999999988
        },
        {
          "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.489000000000942
        },
        {
          "name": "ampere_sgemm_32x128_tn",
          "calls": 5,
          "total_us": 30.88200000000097
        },
        {
          "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.879999999999654
        },
        {
          "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.24699999999939
        },
        {
          "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.245999999999412
        },
        {
          "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.73599999999874
        },
        {
          "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.67600000000016
        },
        {
          "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.204000000000406
        },
        {
          "name": "ampere_sgemm_128x32_nn",
          "calls": 5,
          "total_us": 25.532000000000835
        },
        {
          "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.83899999999835
        },
        {
          "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": 21.077999999999747
        },
        {
          "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.131999999999493
        },
        {
          "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.07499999999959
        },
        {
          "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.254999999999882
        },
        {
          "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.517999999999802
        },
        {
          "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.198999999999842
        },
        {
          "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.910000000000537
        },
        {
          "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.652000000000498
        },
        {
          "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.36600000000044
        },
        {
          "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.598000000000411
        },
        {
          "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.596999999999298
        },
        {
          "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.338999999999942
        },
        {
          "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.860000000000127
        },
        {
          "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.635999999999513
        },
        {
          "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.540000000000873
        },
        {
          "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.476999999999862
        },
        {
          "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.250999999999522
        },
        {
          "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.833999999999833
        },
        {
          "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.770999999999049
        },
        {
          "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.544999999999618
        },
        {
          "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.418999999999414
        },
        {
          "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.193000000000438
        },
        {
          "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.13000000000011
        },
        {
          "name": "Memset (Unknown)",
          "calls": 5,
          "total_us": 7.015000000000327
        },
        {
          "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.503999999998996
        }
      ],
      "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": 147.346100433878,
        "calls": 4,
        "idle_before_w": 68.39858362021342,
        "idle_after_w": 76.34419517674121
      },
      "predict_energy": {
        "above_idle_j": 0.3874642921652523,
        "calls": 1652,
        "idle_before_w": 76.00029317177601,
        "idle_after_w": 76.59400411908123
      }
    },
    {
      "name": "decoder-tuned-reference",
      "initialise": {
        "cuda_ms": 2.618016004562378,
        "wall_ms": 2.6768019999963144
      },
      "training": {
        "cuda_ms": 1191.2685546875,
        "wall_ms": 1191.3155769999976
      },
      "predict": {
        "cuda_ms": 2.8412160873413086,
        "wall_ms": 2.8823749999986603
      },
      "full_energy": {
        "above_idle_j": 152.34589403126398,
        "calls": 17,
        "idle_before_w": 76.46579464578255,
        "idle_after_w": 77.53735400862134
      },
      "five_step_kernel_trace": [
        {
          "name": "_encoder_backward",
          "calls": 30,
          "total_us": 1125.1740000000007
        },
        {
          "name": "_encoder_forward",
          "calls": 30,
          "total_us": 939.6409999999985
        },
        {
          "name": "_decoder_backward",
          "calls": 35,
          "total_us": 870.3140000000017
        },
        {
          "name": "_decoder_forward",
          "calls": 35,
          "total_us": 313.81300000000147
        },
        {
          "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": 302.8639999999991
        },
        {
          "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": 265.73900000000003
        },
        {
          "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.38099999999918
        },
        {
          "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": 212.06700000000365
        },
        {
          "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.28399999999988
        },
        {
          "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.42300000000296
        },
        {
          "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.58400000000051
        },
        {
          "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.54000000000133
        },
        {
          "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.9779999999996
        },
        {
          "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": 169.11600000000044
        },
        {
          "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": 164.3460000000009
        },
        {
          "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.50599999999918
        },
        {
          "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": 141.24
        },
        {
          "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.5999999999999
        },
        {
          "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.50599999999918
        },
        {
          "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.51999999999907
        },
        {
          "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.66300000000047
        },
        {
          "name": "memcpy32_post",
          "calls": 45,
          "total_us": 93.77799999999525
        },
        {
          "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": 65.54899999999975
        },
        {
          "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.602000000000544
        },
        {
          "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.600000000000364
        },
        {
          "name": "ampere_sgemm_32x32_sliced1x4_nn",
          "calls": 5,
          "total_us": 50.534000000000106
        },
        {
          "name": "memcpy128",
          "calls": 10,
          "total_us": 50.02499999999918
        },
        {
          "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.74500000000012
        },
        {
          "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.10400000000118
        },
        {
          "name": "ampere_sgemm_64x32_sliced1x4_nt",
          "calls": 5,
          "total_us": 39.42900000000009
        },
        {
          "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.54300000000012
        },
        {
          "name": "void at::native::vectorized_gather_kernel<16, long>(char*, char*, long*, int, long, long, long, long, bool)",
          "calls": 10,
          "total_us": 34.81900000000087
        },
        {
          "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.37399999999934
        },
        {
          "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.62099999999964
        },
        {
          "name": "ampere_sgemm_32x128_tn",
          "calls": 5,
          "total_us": 31.04400000000078
        },
        {
          "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.885999999999967
        },
        {
          "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.18800000000101
        },
        {
          "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": 28.805000000000064
        },
        {
          "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.582000000000335
        },
        {
          "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.205000000000837
        },
        {
          "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.370000000000573
        },
        {
          "name": "ampere_sgemm_128x32_nn",
          "calls": 5,
          "total_us": 25.541000000001077
        },
        {
          "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.981000000000677
        },
        {
          "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.99400000000037
        },
        {
          "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.531999999999698
        },
        {
          "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.123000000001184
        },
        {
          "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.242999999999938
        },
        {
          "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.731999999998834
        },
        {
          "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.410999999999376
        },
        {
          "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.83199999999988
        },
        {
          "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.419000000000551
        },
        {
          "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": 12.067000000000235
        },
        {
          "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": 11.904999999999745
        },
        {
          "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.586999999999534
        },
        {
          "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.521000000001322
        },
        {
          "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.881999999998698
        },
        {
          "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.720999999999776
        },
        {
          "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.528999999999996
        },
        {
          "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.497999999998683
        },
        {
          "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.274000000000342
        },
        {
          "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.697999999999183
        },
        {
          "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.695999999998321
        },
        {
          "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.440999999999804
        },
        {
          "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.121999999999161
        },
        {
          "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.119999999998981
        },
        {
          "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.480999999998858
        },
        {
          "name": "Memset (Unknown)",
          "calls": 5,
          "total_us": 7.298000000001139
        },
        {
          "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.528000000000475
        }
      ],
      "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": 147.26208365832807,
        "calls": 4,
        "idle_before_w": 77.4070124015452,
        "idle_after_w": 77.4929564668283
      },
      "predict_energy": {
        "above_idle_j": 0.3876781654154905,
        "calls": 1634,
        "idle_before_w": 77.17602637119197,
        "idle_after_w": 77.33842251891704
      }
    }
  ],
  "official": false,
  "full_energy_window_s": 20.0,
  "gpu": "NVIDIA A100 80GB PCIe",
  "device": {
    "name": "NVIDIA A100 80GB PCIe",
    "uuid": "GPU-6717620b-78b9-4a0d-ab49-2484c1fb4b0a",
    "vbios": "92.00.68.00.01",
    "driver": "580.95.05",
    "power_limit_w": 300.0
  },
  "elapsed_s": 127.891042811,
  "app_id": "ap-M0nmO0ShNllwXKhTIUchJC",
  "launch_elapsed_s": 139.29774066701066
}
