kernelbench.com

KernelBench cuda · RTX PRO 6000

Grid + MinGRU SPS Claude Opus 5

1.96×geomean speedup across shapes

manually audited: clean

One persistent CUDA megakernel (mingru3.cu) fusing the entire rollout: obs encode, 3x MinGRU via fp16 mma.m16n8k16 with host-side per-lane fragment pre-packing, argmax, env step, device-side MT19937 food respawn; speculative/ lockstep dual-kernel scheme and a measured cost model for envs-per-block planning. Weight pack cached keyed on (data_ptr, _version) per parameter -- version counters bump on in-place mutation, correct invalidation. EMPIRICAL (mandatory for >1.0): three weight seeds via in-place reset_parameters -> policy_forward matches reference at <=3e-8, full run() positions/rewards BIT-EXACT vs reference every time, last-logits <=3e-6, outputs vary with weights (reward sum 214 -> 188). An out-of-distribution +/-0.05 perturbation probe showed rollout divergence from fp16 rounding at near-tie logits (chaotic env compounding), not staleness; in-distribution retest settles it. 1.96x SPS vs the frozen eager anchor is empirically real (SPS metric; >1.0 expected by metric design). template_mutated=false, numeric stress on, zero cross-run access. Sequential isolated re-grade on anvil GPU0 2026-07-26 (contended 2.0038, delta -2.1%).

harnessor-opusagent session12h 0mtotal wall12h 2mcheck1sbenchmark1soutput tokensgpu-lock wait71sgpu-lock held31sregimethroughput

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

No per-shape benchmark data archived for this run.

Kernel source (redacted)
"""Fused CUDA grid-foraging + 3x MinGRU(h=256) rollout.

The whole rollout (obs -> encoder -> 3 MinGRU layers -> logits -> argmax -> env
step -> food respawn) runs inside one persistent CUDA megakernel; the gate
matmuls use fp16 tensor cores (mma.m16n8k16), which is why the weights are
pre-packed here into per-lane fragment order.  See mingru2.cu.
"""
from __future__ import annotations

import os
from pathlib import Path

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load

BOARD = 11
OBS_DIM = 4
HIDDEN = 256
GRU_LAYERS = 3
NUM_ACTIONS = 4
GRU_OUT = 3 * HIDDEN

_HERE = Path(__file__).resolve().parent

# Build only for the device we are on (the default env list would compile 6 archs).
if torch.cuda.is_available():
    _cc = torch.cuda.get_device_capability(0)
    os.environ["TORCH_CUDA_ARCH_LIST"] = f"{_cc[0]}.{_cc[1]}"

_TUNE = os.environ.get("MINGRU_TUNE", "")
_SRC = os.environ.get("MINGRU_SRC", "mingru3.cu")
_ext = load(
    name="mingru_" + _SRC.split(".")[0] + (
        "_" + "".join(c if c.isalnum() else "_" for c in _TUNE) if _TUNE else ""),
    sources=[str(_HERE / _SRC)],
    extra_cuda_cflags=["-O3", "--ptxas-options=-v", "-lineinfo"] + _TUNE.split(),
    extra_cflags=["-O3"],
    verbose=bool(int(os.environ.get("MINGRU_VERBOSE", "0"))),
)

# envs per block; each config is a separate kernel instantiation.
_ECFG = (16, 32, 48, 64, 80, 96)
_maxres_cache: dict[int, int] = {}
_shr_cache: dict[int, int] = {}


def _max_resident(E: int) -> int:
    v = _maxres_cache.get(E)
    if v is None:
        v = min(int(_ext.max_resident(E)), 1023)  # ctr[] packs arrivals in 10 bits
        _maxres_cache[E] = v
    return v


def _sh_rows(E: int) -> int:
    v = _shr_cache.get(E)
    if v is None:
        v = int(_ext.sh_rows(E))
        _shr_cache[E] = v
    return v


# Cost of one env-step, in SM cycles, as a function of envs per block and of how
# many blocks are resident:
#
#   cyc(E, nb) = cref[E] + slope[E] * (nb - nbref[E])
#
# Measured, not derived -- an earlier bandwidth model got the ranking wrong.  The
# level is mostly on-SM work that does not depend on E (tensor cores alone are
# ~740 of it: 410,624 MACs per env-step at 557 MAC/cyc/SM), and the slope is the
# L2 contention of nb blocks all streaming the same 768 KB of weights.  Bigger E
# amortises that stream but leaves SMs idle when n/E < 188, which is the whole
# tradeoff: shape (4096, 32) can only fill 128 SMs at E=32 and 86 at E=48.
#
# Refitted from end-to-end run() wall times (scratch/esweep.py: every E forced on
# every graded shape, mt_seed and host time subtracted), not from the differenced
# microbenchmark -- the old fit predated the warp-staggered k rotation and ranked
# E=64 above E=96 at n=65536, which costs 1.35% there.  Everything is referenced
# to nb=171 now, the block count three of the four shapes actually run at, so the
# extrapolation is short where it matters.
#          cref  nbref  slope
_ECOST = {
    16: (2808.0, 171, 10.75),
    32: (1515.0, 171, 5.37),
    48: (1196.0, 171, 1.56),
    64: (1150.0, 171, 1.16),
    80: (1105.0, 171, 0.72),
    96: (1144.0, 171, 1.30),
}

# 0 disables the speculative kernel (the exact lockstep kernel then does
# everything); used to A/B the two paths.
_SPMUL = float(os.environ.get("MINGRU_SPMUL", "1"))
# tuning overrides (0 = use the model)
_FORCE_E = int(os.environ.get("MINGRU_E", "0"))
_FORCE_NB = int(os.environ.get("MINGRU_NB", "0"))

_NSM = torch.cuda.get_device_properties(0).multi_processor_count if torch.cuda.is_available() else 1
_plan_cache: dict[tuple, tuple[int, int, int, int]] = {}


def _plan(n: int) -> tuple[int, int, int, int]:
    """-> (E, nchunk, nblocks_ls, nblocks_sp) minimising the modelled step time."""
    key = (n, _FORCE_E, _FORCE_NB, _SPMUL)
    hit = _plan_cache.get(key)
    if hit is not None:
        return hit
    best = None
    for E in _ECFG:
        nchunk = (n + E - 1) // E
        cap = min(_max_resident(E), _NSM)
        waves = (nchunk + cap - 1) // cap
        # Balance the waves: ceil(nchunk/waves) blocks run the same number of
        # chunk-steps as `cap` blocks would, with less L2 contention per block.
        nb = (nchunk + waves - 1) // waves
        cref, nbref, slope = _ECOST[E]
        cyc = waves * E * max(900.0, cref + slope * (nb - nbref))
        if best is None or cyc < best[0]:
            best = (cyc, E, nchunk, nb)
    E, nchunk, nbsp = best[1], best[2], best[3]
    if _FORCE_E:  # tuning override: same wave balancing, forced E
        E, nchunk = _FORCE_E, (n + _FORCE_E - 1) // _FORCE_E
        cap = min(_max_resident(E), _NSM)
        nbsp = (nchunk + (nchunk + cap - 1) // cap - 1) // ((nchunk + cap - 1) // cap)
    # the lockstep fallback is exact but slow; it wants every SM it can get
    nblocks = min(nchunk, _max_resident(E), _NSM)
    if _SPMUL <= 0:
        nbsp = 0
    if _FORCE_NB:
        nbsp = min(nbsp, _FORCE_NB) if nbsp else 0
        nblocks = min(nblocks, _FORCE_NB)
    out = (E, nchunk, nblocks, nbsp)
    _plan_cache[key] = out
    return out


def _pick_E(n: int) -> int:
    return _plan(n)[0]


class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.w_enc = nn.Parameter(torch.empty(HIDDEN, OBS_DIM))
        self.b_enc = nn.Parameter(torch.zeros(HIDDEN))
        self.w_gru = nn.Parameter(torch.empty(GRU_LAYERS, GRU_OUT, HIDDEN))
        self.w_a = nn.Parameter(torch.empty(NUM_ACTIONS, HIDDEN))
        self.b_a = nn.Parameter(torch.zeros(NUM_ACTIONS))
        self.w_v = nn.Parameter(torch.empty(1, HIDDEN))
        self.b_v = nn.Parameter(torch.zeros(1))
        self.reset_parameters(0)
        self._cache = None

    def reset_parameters(self, seed: int = 0) -> None:
        g = torch.Generator(device="cpu")
        g.manual_seed(seed)
        for p in self.parameters():
            tmp = torch.empty(p.shape, dtype=p.dtype, device="cpu")
            tmp.normal_(0.0, 0.02, generator=g)
            p.data.copy_(tmp)
        self._cache = None

    def forward(self, obs: torch.Tensor, state: torch.Tensor):
        return policy_forward(self, obs, state)


# --------------------------------------------------------------------------
# fp16 fragment packing
#
# mma.m16n8k16 (row.col) wants, for lane l (lj = l & 3, group = l >> 2):
#   B[k][m] with m = 8*jblk + group and k = 16*kb + {2*lj, 2*lj+1, 2*lj+8, 2*lj+9}
# so the four (or eight, for two k blocks) halves a lane needs are contiguous.
# --------------------------------------------------------------------------
def _hpos(k: torch.Tensor) -> torch.Tensor:
    """position of hidden index k inside a 16-half shared block"""
    return (k & ~15) + 4 * ((k & 7) >> 1) + (k & 1) + (((k >> 3) & 1) << 1)


def _pack_gru(W: torch.Tensor) -> torch.Tensor:
    """[HID][GATES] fp32 -> [8 kbp][32 jblk][3 g][32 lane][8] fp16"""
    kbp = torch.arange(8).view(8, 1, 1, 1, 1)
    jblk = torch.arange(32).view(1, 32, 1, 1, 1)
    g = torch.arange(3).view(1, 1, 3, 1, 1)
    lane = torch.arange(32).view(1, 1, 1, 32, 1)
    i = torch.arange(8).view(1, 1, 1, 1, 8)
    kk, ii, lj = i >> 2, i & 3, lane & 3
    k = 16 * (2 * kbp + kk) + torch.where(ii < 2, 2 * lj + ii, 2 * lj + 6 + ii)
    m = g * HIDDEN + 8 * jblk + (lane >> 2)
    idx = (k * GRU_OUT + m).reshape(-1)
    return W.reshape(-1)[idx].half()


def _pack_w0(W0: torch.Tensor) -> torch.Tensor:
    """[16][1024] fp32 (one k block) -> [32 jblk][4 g][32 lane][4] fp16"""
    jblk = torch.arange(32).view(32, 1, 1, 1)
    g = torch.arange(4).view(1, 4, 1, 1)
    lane = torch.arange(32).view(1, 1, 32, 1)
    i = torch.arange(4).view(1, 1, 1, 4)
    lj = lane & 3
    k = torch.where(i < 2, 2 * lj + i, 2 * lj + 6 + i)
    m = g * HIDDEN + 8 * jblk + (lane >> 2)
    idx = (k * (GRU_OUT + HIDDEN) + m).reshape(-1)
    return W0.reshape(-1)[idx].half()


def _pack_wa_frag(wa: torch.Tensor) -> torch.Tensor:
    """[NACT][HID] fp32 -> [8 kbp][32 lane][8] fp16 mma-B fragments (n padded to 8)"""
    w8 = torch.zeros(8, HIDDEN)
    w8[:NUM_ACTIONS] = wa
    kbp = torch.arange(8).view(8, 1, 1)
    lane = torch.arange(32).view(1, 32, 1)
    i = torch.arange(8).view(1, 1, 8)
    kk, ii, lj = i >> 2, i & 3, lane & 3
    k = 16 * (2 * kbp + kk) + torch.where(ii < 2, 2 * lj + ii, 2 * lj + 6 + ii)
    n = lane >> 2
    return w8[n.expand_as(k).reshape(-1), k.reshape(-1)].half()


def _derived(model: Model) -> dict:
    """Weight layouts the kernels want, cached until parameters change."""
    ps = [model.w_enc, model.b_enc, model.w_gru, model.w_a, model.b_a, model.w_v, model.b_v]
    sig = tuple((p.data_ptr(), p._version) for p in ps)
    cache = getattr(model, "_cache", None)
    if cache is not None and cache[0] == sig:
        return cache[1]

    dev = model.w_gru.device
    with torch.no_grad():
        w_enc = model.w_enc.detach().float()
        b_enc = model.b_enc.detach().float()
        w_gru = model.w_gru.detach().float()
        # [NOBS][HID]: rows 0..3 = w_enc^T, row 4 = b_enc
        wenc5 = torch.cat([w_enc.t(), b_enc.view(1, HIDDEN)], 0).contiguous()
        # folded layer-0 gates: w1[i][m] = sum_j w_gru[0][m][j] * wenc5[i][j]
        w1 = (w_gru[0].double() @ wenc5.double().t()).t().float()
        # one 16x1024 k block: [gates | encoder], rows 5..15 zero
        W0 = torch.zeros(16, GRU_OUT + HIDDEN)
        W0[:5, :GRU_OUT] = w1.cpu()
        W0[:5, GRU_OUT:] = wenc5.cpu()
        wgt_all = w_gru.permute(0, 2, 1).contiguous()  # [3][HID][GATES]
        wgc = wgt_all.cpu()
        # logit weights, permuted to match the k order h is stored in
        hp = _hpos(torch.arange(HIDDEN))
        wap = torch.empty(NUM_ACTIONS, HIDDEN)
        wap[:, hp] = model.w_a.detach().float().cpu()
        # [NACT][HID] flat in hs-position order, then the same weights as
        # mma B fragments; the kernel picks whichever the logit epilogue wants.
        wa_flat = torch.cat([wap.half().reshape(-1),
                             _pack_wa_frag(model.w_a.detach().float().cpu())])
        d = {
            "w0": _pack_w0(W0).to(dev),
            "wg": torch.cat([_pack_gru(wgc[1]), _pack_gru(wgc[2])]).to(dev),
            "wap": wa_flat.contiguous().to(dev),
            "wgt_all": wgt_all,
            "wa": model.w_a.detach().float().contiguous(),
            "ba": model.b_a.detach().float().contiguous(),
            "wv": model.w_v.detach().float().reshape(-1).contiguous(),
            "bv": model.b_v.detach().float().contiguous(),
            "wenc": w_enc.t().contiguous(),
            "benc": b_enc.contiguous(),
        }
    model._cache = (sig, d)
    return d


def policy_forward(model: Model, obs: torch.Tensor, state: torch.Tensor):
    """obs (N,4), state (N,L,H) -> logits (N,4), new_state (N,L,H), value (N,)."""
    d = _derived(model)
    obs = obs.contiguous().float()
    state = state.contiguous().float()
    # torch >= 2.13 runs fp32 matmuls through tf32 tensor cores by default, so the
    # eager oracle carries ~5e-4 relative rounding.  Mirror whichever arithmetic
    # mode is active so this entry point tracks it instead of drifting from it.
    tf32 = 1 if torch.backends.cuda.matmul.allow_tf32 else 0
    logits, nst, val = _ext.policy_forward(
        obs, state, d["wenc"], d["benc"], d["wgt_all"], d["wa"], d["ba"], d["wv"], d["bv"], tf32
    )
    return logits, nst, val


def env_step(agent: torch.Tensor, food: torch.Tensor, actions: torch.Tensor,
             rng_state: torch.Tensor):
    a, f, r, g = _ext.env_step(
        agent.contiguous().float(), food.contiguous().float(),
        actions.contiguous().to(torch.int64), rng_state.contiguous().to(torch.int64),
    )
    return a, f, r, g


# Scratch buffers, none of which run() hands back, so they can be reused across
# calls; every one is fully written before it is read.  Keyed by the shape that
# sizes them, which the benchmark reuses across its trials.
_SCRATCH: dict = {}


def _scratch(npad: int, H: int, nblocks: int, dev) -> tuple:
    key = (npad, H, nblocks)
    b = _SCRATCH.get(key)
    if b is None:
        # [agent|food][env][x,y]; filled on the device by mt_init.  The two
        # halves are handed to the kernels separately and building the views
        # costs ~2us each, so they are cached with the storage.
        pos = torch.empty(2, npad, 2, dtype=torch.int32, device=dev)
        b = (
            pos[0], pos[1],
            # interleaved [3][128][npad*2]; step 0 writes it all before reading
            torch.empty(3 * 128 * npad * 2, device=dev),
            # [0:H] sp hit counts | [H:2H] ls publication | done | spec-ok flag
            torch.empty(2 * H + 2, dtype=torch.int32, device=dev),
            torch.empty(nblocks, H, dtype=torch.int32, device=dev),
        )
        _SCRATCH[key] = b
    return b


def run(num_envs: int, horizon: int, seed: int, model: Model | None = None) -> dict:
    dev = torch.device("cuda:0")
    if model is None:
        model = Model().to(dev).eval()
    d = _derived(model)

    n = int(num_envs)
    H = int(horizon)
    E, nchunk, nblocks, nbsp = _plan(n)
    npad = nchunk * E
    ag, fd, state, aux, hist = _scratch(npad, H, nblocks, dev)

    # Fire the RNG engine first: it is a single block for tens of microseconds,
    # so the output allocations and the rollout launch below run under it.
    _ext.mt_seed(ag, fd, aux, seed, n, npad)

    agout = torch.empty(npad, 2, dtype=torch.int64, device=dev)
    rew = torch.empty(npad, device=dev)
    lgo = torch.empty(npad, NUM_ACTIONS, device=dev)

    _ext.rollout(d["w0"], d["wg"], d["wap"], d["ba"], state, ag, agout, fd, rew, lgo,
                 aux, hist, seed, n, npad, nchunk, nblocks, nbsp, H, E)

    return {
        "rewards": rew[:n],
        "positions": agout[:n],
        "last_logits": lgo[:n],
    }


# ==================================================================
# ===== sidecar: mingru3.cu (56356 bytes, loaded by solution.py) =====
# ==================================================================

// Fused grid-foraging + 3x MinGRU(h=256) rollout megakernel for SM120, fp16
// tensor-core edition.
//
// mma.m16n8k16 sustains 554 MAC/cycle/SM here vs 128.4 for FFMA (4.3x) at the
// same 10-bit mantissa the oracle's own TF32 matmuls use, so every gate matmul
// runs on the tensor cores:
//   M(16) = envs, N(8) = gate rows, K(16) = hidden
//   A = h  (fp16 in shared, k-permuted so a lane's 4 halves are one LDS.64)
//   B = W  (fp16 in global, pre-packed in per-lane fragment order -> LDG.128)
//   C = fp32, 3 gate tiles per (env-block, gate-block) so one lane sees
//       zh/zg/zp of the same j and can run the highway in registers.
// Lane l owns gate cols j0+{0,1} (j0 = 8*jblk + 2*(l&3)) and env rows
// (l>>2), (l>>2)+8, so the recurrent state is stored interleaved (see sidx)
// and each accumulator tile maps onto exactly one LDG.128/STG.128.
//
// Layer 0 is the encoder Linear(4->256) folded into the first gate matmul: a
// 1024x5 fp16 matrix (768 gate rows + 256 encoder rows, 5th k row = constant
// 1.0 for the biases) evaluated as a single k=16 mma pass.
//
// The reference env advances its LCG for *every* env whenever *any* env eats
// food, so an env's rng after t steps is LCG^(2K(t))(r0) with K(t) the number
// of steps that had at least one global hit.  Two kernels exploit that:
//
//   rollout_sp  speculates K(t) = t, which is true as soon as one of the
//               (many) envs eats on every step.  Then chunks are completely
//               independent: one block runs E envs for the whole horizon with
//               agent/food/rng/reward in shared memory and only its own
//               3*256*E floats of recurrent state in flight.  It records
//               per-step global hit counts and the last block out sets a flag
//               saying whether the speculation held.
//   rollout_ls  the exact lockstep fallback: chunk loop inside the step loop,
//               blocks publish their hit bit per step into lsc[t] with one
//               atomicAdd (arrivals in the low 10 bits) and read it back only
//               when they themselves hit and are missing history.  Launched
//               unconditionally right after rollout_sp and returns immediately
//               when the flag says the speculation was exact, so the whole
//               decision stays on the device.
#include <ATen/cuda/CUDAContext.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/extension.h>

#include <cstdint>
#include <vector>

#define HID 256
#define GATES 768
#define NOBS 5  // 4 obs values + a constant 1.0 row (folds the biases in)
#define BOARD 11
#define NACT 4
// The observation's two divisions compile to FCHK + a call to the IEEE slow
// path, ~10 instructions each, four per env-step.  Both numerators are small
// integers -- (food-agent) in [-10,10] over 11, agent in [0,10] over 10 -- and
// over those 32 inputs the correctly-rounded quotient and the reciprocal
// product round to the same fp16, checked exhaustively.  (0x3dba2e8c and
// 0x3dcccccd are fl(1/11) and fl(1/10).)
#define RBOARD __uint_as_float(0x3dba2e8cu)
#define RBOARD1 __uint_as_float(0x3dcccccdu)
// halves per env row of hs. 272/2 = 136 words, 136 % 32 = 8 -> the A-fragment
// LDS.64s hit banks 8e+2j, conflict free.
#define HSTRIDE 272
// halves per action row of the logit weights: 264/2 = 132 words -> the four
// actions of a lane quad land 4 banks apart.
#define WASTRIDE 264

typedef unsigned long long u64;

__device__ u64 g_prof[8];

__device__ __forceinline__ u64 lcg_step(u64 r) {
  return (r * 6364136223846793005ULL + 1ULL) & 0x7FFFFFFFFFFFFFFFULL;
}
// 1/(1+exp(-x)) : MUFU.EX2 + MUFU.RCP, ~1e-7 relative error.
__device__ __forceinline__ float sgm(float x) { return __frcp_rn(1.0f + __expf(-x)); }
// The highway evaluates one tanh and two sigmoids per hidden unit per layer,
// i.e. 2304 transcendentals per env-step, and libdevice's tanhf() is an
// exp-based rational costing MUFU.EX2 + MUFU.RCP + ~14 FFMA -- which made
// transcendental emulation, not the matmul, the largest instruction group in
// the kernel.  Two cheaper spellings, selected by TMODE below:

// sm_120 has a hardware tanh, but nvcc only emits it for the PTX intrinsic.
// One MUFU, max error 2^-11 relative.
__device__ __forceinline__ float tanh_hw(float x) {
  float r;
  asm("tanh.approx.f32 %0, %1;" : "=f"(r) : "f"(x));
  return r;
}
// 1 - 2/(exp(2x)+1): MUFU.EX2 + MUFU.RCP + 3 FFMA, ~5e-8 absolute error, and
// right at both saturations (exp overflows to inf -> +1, underflows -> -1).
__device__ __forceinline__ float tanh_ex(float x) {
  float e = __expf(2.0f * x), r;
  asm("rcp.approx.f32 %0, %1;" : "=f"(r) : "f"(e + 1.0f));
  return fmaf(-2.0f, r, 1.0f);
}
// 1/(1+exp(-x)) : MUFU.EX2 + MUFU.RCP, ~1e-7 relative error.
__device__ __forceinline__ float sgm_ex(float x) {
  float r;
  asm("rcp.approx.f32 %0, %1;" : "=f"(r) : "f"(__expf(-x) + 1.0f));
  return r;
}

// 0 = libdevice tanhf everywhere (what the reference does, ~16 instr/call)
// 1 = hardware tanh everywhere         3 MUFU +  6 FFMA per hidden unit
// 2 = accurate everywhere              6 MUFU +  7 FFMA
// 3 = accurate candidate, hw gates     4 MUFU +  5 FFMA
// 4 = hw candidate, accurate gates     5 MUFU +  4 FFMA
//
// 1 is the default because the approximation is free: over 120 rollouts
// (958k env-steps) against reference.py, modes 0, 1 and 2 all produce the same
// single greedy-argmax flip and max |dlogit| 1.1-1.3e-4.  The error that
// matters is the fp16 gate matmul, not the tanh, so paying for an accurate
// tanh on top of it buys nothing -- and mode 0 scores 1.66 where mode 1 scores
// 1.91.
#ifndef TMODE
#define TMODE 1
#endif
#if TMODE == 0
#define TANH_C(x) tanhf(x)
#define SGH(x) fmaf(0.5f, tanhf(0.5f * (x)), 0.5f)
#elif TMODE == 1
#define TANH_C(x) tanh_hw(x)
#define SGH(x) fmaf(0.5f, tanh_hw(0.5f * (x)), 0.5f)
#elif TMODE == 2
#define TANH_C(x) tanh_ex(x)
#define SGH(x) sgm_ex(x)
#elif TMODE == 3
#define TANH_C(x) tanh_ex(x)
#define SGH(x) fmaf(0.5f, tanh_hw(0.5f * (x)), 0.5f)
#else
#define TANH_C(x) tanh_hw(x)
#define SGH(x) sgm_ex(x)
#endif

// Round to a tf32 significand, as the tensor cores do when torch runs fp32
// matmuls in "high" precision. Used only by policy_forward, to mirror whatever
// arithmetic mode the oracle is in.
__device__ __forceinline__ float to_tf32(float x) {
  unsigned int i = __float_as_uint(x);
  return __uint_as_float((i + 0x0FFFu + ((i >> 13) & 1u)) & 0xFFFFE000u);
}
__device__ __forceinline__ float rnd(float x, bool t) { return t ? to_tf32(x) : x; }

// permuted shared position of hidden index k (so that the 4 halves an mma lane
// needs for one k block are contiguous)
__host__ __device__ __forceinline__ int hpos(int k) {
  return (k & ~15) + 4 * ((k & 7) >> 1) + (k & 1) + (((k >> 3) & 1) << 1);
}

// interleaved recurrent state index: [3][128][NP/16][8][2][2], so the four
// (env, j) values one lane owns are 4 consecutive floats.
__host__ __device__ __forceinline__ size_t sidx(int l, int j, int e, int NP) {
  return ((size_t)l * 128 + (j >> 1)) * ((size_t)NP * 2) + (size_t)(e >> 4) * 32 +
         (e & 7) * 4 + (((e >> 3) & 1) << 1) + (j & 1);
}

#define MMA16816(d0, d1, d2, d3, a0, a1, a2, a3, b0, b1)                       \
  asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "            \
               "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"       \
               : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)                        \
               : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1))

// ---------------------------------------------------------------------------
// Weight-stream load policy.
//
// Every block re-reads the same 768 KB of weights once per env-step, so the
// buffer wants to be pinned in L2 (measured: L2 tops out near 8.4 TB/s at 171
// blocks and this stream is ~80% of that) while staying out of L1, where it
// would evict the recurrent state.  __ldcs says evict-first at *both* levels,
// which is right for L1 and wrong for L2; WLD picks between the variants.
//   0 = ld.global.cs                              (evict-first everywhere)
//   1 = ld.global.nc                              (default policy)
//   2 = ld.global.nc.L1::no_allocate              (skip L1, normal L2)
//   3 = ld.global.nc.L1::evict_first              (L1 evict-first, normal L2)
// ---------------------------------------------------------------------------
// Per-E weight prefetch (double-buffered bf).  Costs a second fragment set of
// registers, so it is only affordable where the accumulators leave room -- see
// the PF16..PF96 defaults next to NG/KP/GRP below.
#ifndef PF16
#define PF16 0
#endif
#ifndef PF32
#define PF32 0
#endif
#ifndef PF48
#define PF48 0
#endif
#ifndef PF64
#define PF64 0
#endif
#ifndef PF80
#define PF80 0
#endif
#ifndef PF96
#define PF96 0
#endif
__host__ __device__ constexpr int pf_for(int E) {
  return E == 16 ? PF16 : E == 32 ? PF32 : E == 48 ? PF48
         : E == 64 ? PF64 : E == 80 ? PF80 : PF96;
}

#ifndef WLD
#define WLD 0
#endif
#if WLD == 0
#define WLD_Q "ld.global.cs"
#elif WLD == 1
#define WLD_Q "ld.global.nc"
#elif WLD == 2
#define WLD_Q "ld.global.nc.L1::no_allocate"
#else
#define WLD_Q "ld.global.nc.L1::evict_first"
#endif

__device__ __forceinline__ uint4 wld4(const uint4* p) {
  uint4 v;
  asm(WLD_Q ".v4.u32 {%0,%1,%2,%3}, [%4];"
      : "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
      : "l"(p));
  return v;
}
__device__ __forceinline__ uint2 wld2(const uint2* p) {
  uint2 v;
  asm(WLD_Q ".v2.u32 {%0,%1}, [%2];" : "=r"(v.x), "=r"(v.y) : "l"(p));
  return v;
}

struct RParams {
  const __half* __restrict__ w0;  // [32 jblk][4 g][32 lane][4]      folded layer 0
  const __half* __restrict__ wg;  // [2][8 kbp][32 jblk][3 g][32 lane][8]
  const __half* __restrict__ wa;  // [NACT][HID]  hpos-permuted
  const float* __restrict__ ba;   // [NACT]
  float* __restrict__ state;      // [3][128][NP*2]  interleaved (see sidx)
  int* __restrict__ ag;           // [NP][2]   initial agent (x,y)
  long long* __restrict__ agout;  // [NP][2]   final agent (x,y)
  int* __restrict__ fd;           // [NP][2]
  float* __restrict__ rew;        // [NP]
  float* __restrict__ lgo;        // [NP][NACT]
  int* __restrict__ ctr;          // [H]     sp: global hits per step
  int* __restrict__ lsc;          // [H]     ls: arrival/hit publication
  int* __restrict__ done;         // [1]     sp: block exit counter
  int* __restrict__ flag;         // [1]     1 => sp was exact, ls must no-op
  int* __restrict__ hist;         // [nblocks][H]
  long long seedbase;             // seed * 10007
  int n, np, nchunk, nblocks, nbsp, horizon;
};

// ---------------------------------------------------------------------------
// shared layout (identical for both kernels so they have the same occupancy)
// ---------------------------------------------------------------------------
#define SH_DECL(E)                                                     \
  extern __shared__ char sraw[];                                       \
  __half* __restrict__ hs = (__half*)sraw;            /* [E][HSTRIDE] */ \
  __half* __restrict__ oh = hs + (size_t)E * HSTRIDE; /* [E][16]      */ \
  __half* __restrict__ was = oh + (size_t)E * 16;     /* [NACT][WASTRIDE] */ \
  float* __restrict__ lg_s = (float*)(was + NACT * WASTRIDE); /* [NACT][E] */ \
  int4* __restrict__ es = (int4*)(lg_s + NACT * E);   /* [E] agent|food  */ \
  u64* __restrict__ rs = (u64*)(es + E);              /* [E] lcg         */ \
  float* __restrict__ rws = (float*)(rs + E);         /* [E] reward      */ \
  int* __restrict__ shi = (int*)(lg_s + NACT * E + 8 * E);

__host__ __device__ constexpr size_t shmem_for(int E) {
  return (size_t)E * HSTRIDE * 2 + (size_t)E * 32 + NACT * WASTRIDE * 2 +
         (size_t)NACT * E * 4 + (size_t)E * 32 + 16;
}

// Because rollout_sp owns its chunk for the whole horizon, its recurrent state
// can live in shared memory instead of L2 -- and L2 bandwidth is what this
// kernel is short of.  Spend whatever is left of the 99 KB opt-in budget on it
// (E=16: all three layers, E=64: none).
//
// Row stride, in floats, of one (layer, j/2) row of shared state.  A lane holds
// 4 consecutive floats at 2*(j0/2)*... + 4*le, and an LDS.128 is serviced 8
// lanes at a time, so with stride S the phase touches slots
// (lj*S/4 + le) mod 8; S = 2E+8 makes that 2*lj+le, a bijection.  The natural
// 2E stride (which is what the global layout uses) would be a 4-way conflict.
#define SSTRIDE(E) (2 * (E) + 8)
// Rows of recurrent state -- there are 3*128 of them, one per (layer, j/2) --
// that fit in whatever is left of the budget.  Kept a multiple of 4 because a
// warp owns one gate block = 4 consecutive rows, so the shared/global choice in
// the highway stays warp uniform (E=16: all 384 rows, E=96: 48).
__host__ __device__ constexpr int sh_rows(int E) {
  return (int)((101376 - shmem_for(E)) / (size_t)(SSTRIDE(E) * 4)) >= 384
             ? 384
             : (int)((101376 - shmem_for(E)) / (size_t)(SSTRIDE(E) * 4)) & ~3;
}
static inline size_t shmem_sp(int E) {
  return shmem_for(E) + (size_t)sh_rows(E) * SSTRIDE(E) * 4;
}

// ---------------------------------------------------------------------------
// One policy step for E envs.  The observation must already be in oh; on
// return hs holds the new h, stbase the new recurrent state and lg_s the
// logits.  Both kernels share this; only the state addressing differs
// (per-block scratch for sp, per-env global for ls).
// ---------------------------------------------------------------------------
// A warp owns NG of the 32 gate blocks (8 j x 3 gates each) and works through
// them GRP at a time, so only acc[MB][GRP][3][4] is live instead of all NG.
// GRP is a straight tradeoff: an A fragment read out of hs feeds GRP*3 mmas, so
// the shared-memory traffic of the matmul is 256/GRP cycles per env-step (of
// ~1300 total), while the accumulators cost 12*MB*GRP registers.  NG also sets
// the block size (32/NG warps), i.e. how many registers a thread may have --
// which is why the wide-GRP configs run 256 threads.
// XMUFU adds N *value-neutral* transcendentals per element: fmaf(0, t, x) is
// exactly x for finite t, so the rollout takes the same branches and the timing
// stays comparable.  dc/dXMUFU is the marginal price of an SFU op, i.e. how much
// of the highway's 3-per-unit MUFU bill the scheduler fails to hide under the
// tensor pipe.  (ABLATE cannot measure this: changing values makes the
// speculation mispredict, and the lockstep fallback then dominates the time.)
#ifndef XMUFU
#define XMUFU 0
#endif
// Stagger the k rotation by warp as well as by block.  Without it the four
// warps an SMSP holds reach their weight LDGs in lockstep, so they all stall on
// the same L2 round trip instead of one warp's mmas covering another's miss.
// Bit E/16-1 of the mask enables it for that E; measured (paired, two reps) it
// is worth 1.2-3.0% at E=64 and 1.2-2.3% at E=96, costs 0.8-2.0% at E=48 (it
// also pushes that config from 8 to 20 bytes of spill) and is a wash at E=32.
#ifndef KRWMASK
#define KRWMASK 40  // 1<<3 | 1<<5: E=64 and E=96
#endif
__host__ __device__ constexpr int krotw_for(int E) { return (KRWMASK >> (E / 16 - 1)) & 1; }

template <int E, int NG, int KP, int SHR, int GRP>
__device__ __forceinline__ void step_layers(const RParams& P, __half* __restrict__ hs,
                                            const __half* __restrict__ oh,
                                            const __half* __restrict__ was,
                                            float* __restrict__ lg_s, float* __restrict__ sst,
                                            float* __restrict__ stbase, int NPst, int eoff,
                                            bool fst, bool last, int env0) {
  constexpr int MB = E / 16;
  constexpr int PF = pf_for(E);
  static_assert(NG % GRP == 0, "GRP must divide NG");
  const int tid = threadIdx.x, w = tid >> 5, lane = tid & 31;
  const int lj = lane & 3, le = lane >> 2;

#if XMUFU
  float zz_;
  asm volatile("mov.f32 %0, 0f00000000;" : "=f"(zz_));
#endif
  float acc[MB][GRP][3][4];
  unsigned bf[PF ? 2 : 1][GRP][3][2 * KP];
  __half2 nh[NG][MB][2];  // this layer's new h, published to hs after the loop
  // Every block streams the same 384 KB of weights per layer in the same order,
  // so give each one a different starting k block: the L2 sees requests spread
  // over its slices instead of 188 SMs hammering one line at a time.  Only the
  // fp32 accumulation order changes (~1e-7).
  const int krot = KP * ((int)(blockIdx.x + krotw_for(E) * w) & (16 / KP - 1));

// One weight fragment set (GRP gate blocks x 3 gates x KP k blocks) into BUF.
#define WLOAD(BUF, KB0)                                                            \
  _Pragma("unroll") for (int gi = 0; gi < GRP; ++gi)                               \
      _Pragma("unroll") for (int g = 0; g < 3; ++g) {                              \
    const __half* q = Wl +                                                         \
                      ((((size_t)((KB0) >> 1) * 32 + jb0 + gi) * 3 + g) * 32 + lane) * 8 + \
                      4 * ((KB0) & 1);                                             \
    if (KP == 2) {                                                                 \
      const uint4 v = wld4((const uint4*)q);                                        \
      bf[BUF][gi][g][0] = v.x; bf[BUF][gi][g][1] = v.y;                             \
      bf[BUF][gi][g][2] = v.z; bf[BUF][gi][g][3] = v.w;                             \
    } else {                                                                        \
      const uint2 v = wld2((const uint2*)q);                                        \
      bf[BUF][gi][g][0] = v.x; bf[BUF][gi][g][1] = v.y;                             \
    }                                                                               \
  }

// The KP k blocks held in BUF, against the A fragments read out of hs.
#define KMMA(BUF, KB0)                                                             \
  _Pragma("unroll") for (int kk = 0; kk < KP; ++kk) {                               \
    const int kb = (KB0) + kk;                                                      \
    _Pragma("unroll") for (int mb = 0; mb < MB; ++mb) {                             \
      const uint2 a0 = *(const uint2*)(hs + (16 * mb + le) * HSTRIDE + kb * 16 + 4 * lj); \
      const uint2 a1 =                                                              \
          *(const uint2*)(hs + (16 * mb + le + 8) * HSTRIDE + kb * 16 + 4 * lj);    \
      _Pragma("unroll") for (int gi = 0; gi < GRP; ++gi)                            \
          _Pragma("unroll") for (int g = 0; g < 3; ++g)                             \
              MMA16816(acc[mb][gi][g][0], acc[mb][gi][g][1], acc[mb][gi][g][2],     \
                       acc[mb][gi][g][3], a0.x, a1.x, a0.y, a1.y,                   \
                       bf[BUF][gi][g][2 * kk], bf[BUF][gi][g][2 * kk + 1]);         \
    }                                                                               \
  }

// Issue the next k step's weight load, then spend this one's mmas covering its
// latency.  The stream is ~80% of the L2's achievable bandwidth, so the load is
// hundreds of cycles out; without this the mmas that consume a fragment sit
// right behind the load that produced it and the tensor pipe drains.
#define KSTEP(CUR, OFF)                                                            \
  {                                                                                \
    const int kb0 = (kbi + (OFF) + krot) & 15;                                      \
    if (kbi + (OFF) + KP < 16) {                                                    \
      const int kn = (kbi + (OFF) + KP + krot) & 15;                                 \
      WLOAD(CUR ^ 1, kn)                                                            \
    }                                                                               \
    KMMA(CUR, kb0)                                                                  \
  }

  // ---- gate matmul of GRP gate blocks: acc[.][gi] = W_{jb0+gi}^T * h(hs) ----
  auto kblk = [&](const __half* __restrict__ Wl, int jb0) {
#pragma unroll
    for (int mb = 0; mb < MB; ++mb)
#pragma unroll
      for (int gi = 0; gi < GRP; ++gi)
#pragma unroll
        for (int g = 0; g < 3; ++g)
#pragma unroll
          for (int i = 0; i < 4; ++i) acc[mb][gi][g][i] = 0.f;

    if constexpr (PF) {
      WLOAD(0, (krot & 15))
#pragma unroll 1
      for (int kbi = 0; kbi < 16; kbi += 2 * KP) {
        KSTEP(0, 0)
        KSTEP(1, KP)
      }
    } else {
#pragma unroll 1
      for (int kbi = 0; kbi < 16; kbi += KP) {
        const int kb0 = (kbi + krot) & 15;
        WLOAD(0, kb0)
        KMMA(0, kb0)
      }
    }
  };

  // ---- publish the new h ---------------------------------------------------
  auto flush = [&]() {
#pragma unroll
    for (int ng = 0; ng < NG; ++ng) {
      const int hp = hpos(8 * (NG * w + ng) + 2 * lj);
#pragma unroll
      for (int mb = 0; mb < MB; ++mb) {
        *(__half2*)(hs + (16 * mb + le) * HSTRIDE + hp) = nh[ng][mb][0];
        *(__half2*)(hs + (16 * mb + le + 8) * HSTRIDE + hp) = nh[ng][mb][1];
      }
    }
  };

  // ---- highway: state (=out) read-modify-write, new h into nh[ng] ----------
  // sidx(l, j0, env0 + 16*mb + le, NPst) with j0 even, le < 8 and env0 % 16 == 0
  // collapses to row (l*128 + j0/2) at offset mb*32 + le*4; state rows below SHR
  // live in shared memory, at the padded row stride, in the same shape.  Both
  // arms are spelled out so each keeps its own address space.
// ABLATE is a timing-only knob (the results are wrong): 1 drops the
// transcendentals, 2 also drops the recurrent-state load/store, so the
// difference against a normal build prices each pipe against the matmul.
#ifndef ABLATE
#define ABLATE 0
#endif
#if XMUFU
#define XM_ADD(VAR, SRC)                                                           \
  {                                                                                \
    float xa_ = 0.f;                                                               \
    _Pragma("unroll") for (int q_ = 0; q_ < XMUFU; ++q_) xa_ += TANH_C((SRC) + (float)q_); \
    (VAR) = fmaf(zz_, xa_, (VAR));                                                 \
  }
#else
#define XM_ADD(VAR, SRC)
#endif
#if ABLATE == 0
#define HW_BODY(GI)                                                                    \
    _Pragma("unroll") for (int i = 0; i < 4; ++i) {                                \
      const float o =                                                              \
          sa[i] + SGH(acc[mb][GI][1][i]) * (TANH_C(acc[mb][GI][0][i]) - sa[i]);     \
      const float p = SGH(acc[mb][GI][2][i]);                                      \
      on[i] = o;                                                                   \
      hn[i] = p * o + (1.f - p) * xs[mb][i];                                       \
      XM_ADD(on[i], o)                                                             \
    }
#else
#define HW_BODY(GI)                                                                    \
    _Pragma("unroll") for (int i = 0; i < 4; ++i) {                                \
      const float o = fmaf(0.25f, acc[mb][GI][1][i], sa[i]);                        \
      const float p = 0.25f * acc[mb][GI][2][i];                                    \
      on[i] = o;                                                                   \
      hn[i] = p * o + (1.f - p) * xs[mb][i];                                       \
    }
#endif
#define HIGHWAY(SP, GI, NI)                                                        \
  _Pragma("unroll") for (int mb = 0; mb < MB; ++mb) {                              \
    float* __restrict__ sp = (SP) + mb * 32;                                       \
    float sa[4] = {0.f, 0.f, 0.f, 0.f};                                            \
    if (!fst && ABLATE < 2) {                                                      \
      const float4 sv = *(const float4*)sp;                                        \
      sa[0] = sv.x; sa[1] = sv.y; sa[2] = sv.z; sa[3] = sv.w;                      \
    }                                                                              \
    float on[4], hn[4];                                                            \
    HW_BODY(GI)                                                                    \
    if (ABLATE < 2) *(float4*)sp = make_float4(on[0], on[1], on[2], on[3]);        \
    nh[NI][mb][0] = __floats2half2_rn(hn[0], hn[1]);                               \
    nh[NI][mb][1] = __floats2half2_rn(hn[2], hn[3]);                               \
  }

#define EPILOGUE(L, JB, GI, NI)                                                    \
  {                                                                                \
    const int r_ = (L) * 128 + 4 * (JB) + lj;                                       \
    if (SHR > 0 && r_ < SHR) {                                                      \
      float* __restrict__ sp0 = sst + (size_t)r_ * SSTRIDE(E) + le * 4;              \
      HIGHWAY(sp0, GI, NI)                                                          \
    } else {                                                                        \
      float* __restrict__ sp0 =                                                     \
          stbase + (size_t)r_ * ((size_t)NPst * 2) + eoff + le * 4;                  \
      HIGHWAY(sp0, GI, NI)                                                          \
    }                                                                               \
  }

  // ------------------------- layer 0: folded encoder+gates, one k block
  // hs still holds the previous step's h here, which nothing reads any more, so
  // the new h can be published without a barrier in front of it.
#pragma unroll
  for (int ng = 0; ng < NG; ++ng) {
    const int jb = NG * w + ng;
    float xs[MB][4];  // the encoder output = this layer's highway input
    unsigned b0[4][2];
#pragma unroll
    for (int g = 0; g < 4; ++g) {
      const uint2 v = *(const uint2*)(P.w0 + (((size_t)jb * 4 + g) * 32 + lane) * 4);
      b0[g][0] = v.x;
      b0[g][1] = v.y;
    }
#pragma unroll
    for (int mb = 0; mb < MB; ++mb) {
#pragma unroll
      for (int g = 0; g < 3; ++g)
#pragma unroll
        for (int i = 0; i < 4; ++i) acc[mb][0][g][i] = 0.f;
#pragma unroll
      for (int i = 0; i < 4; ++i) xs[mb][i] = 0.f;
      const uint2 a0 = *(const uint2*)(oh + (16 * mb + le) * 16 + 4 * lj);
      const uint2 a1 = *(const uint2*)(oh + (16 * mb + le + 8) * 16 + 4 * lj);
#pragma unroll
      for (int g = 0; g < 3; ++g)
        MMA16816(acc[mb][0][g][0], acc[mb][0][g][1], acc[mb][0][g][2], acc[mb][0][g][3], a0.x,
                 a1.x, a0.y, a1.y, b0[g][0], b0[g][1]);
      MMA16816(xs[mb][0], xs[mb][1], xs[mb][2], xs[mb][3], a0.x, a1.x, a0.y, a1.y, b0[3][0],
               b0[3][1]);
    }
    EPILOGUE(0, jb, 0, ng)
  }
  flush();
  __syncthreads();

  // ------------------------- recurrent layers 1 and 2
#pragma unroll 1
  for (int l = 1; l < 3; ++l) {
    const __half* __restrict__ Wl = P.wg + (size_t)(l - 1) * HID * GATES;
#pragma unroll
    for (int ng0 = 0; ng0 < NG; ng0 += GRP) {
      kblk(Wl, NG * w + ng0);
#pragma unroll
      for (int gi = 0; gi < GRP; ++gi) {
        const int jb = NG * w + ng0 + gi;
        float xs[MB][4];
        const int hp = hpos(8 * jb + 2 * lj);
#pragma unroll
        for (int mb = 0; mb < MB; ++mb) {
          const __half2 x0 = *(const __half2*)(hs + (16 * mb + le) * HSTRIDE + hp);
          const __half2 x1 = *(const __half2*)(hs + (16 * mb + le + 8) * HSTRIDE + hp);
          xs[mb][0] = __low2float(x0);
          xs[mb][1] = __high2float(x0);
          xs[mb][2] = __low2float(x1);
          xs[mb][3] = __high2float(x1);
        }
        EPILOGUE(l, jb, gi, ng0 + gi)
      }
    }
    __syncthreads();  // everyone is done reading the old h
    flush();
    __syncthreads();
  }
#undef HIGHWAY
#undef EPILOGUE

  // ------------------------------------------------------------------ logits
  // NACT*E dot products spread over the block's 1024/NG threads (E=96 at NG=4
  // needs two rounds).
#pragma unroll
  for (int r = 0; r < (NACT * E * NG + 1023) / 1024; ++r) {
    const int q = tid + r * (1024 / NG);
    if (q >= NACT * E) break;
    const int a = q & 3, ee = q >> 2;
    const __half* __restrict__ hp = hs + (size_t)ee * HSTRIDE;
    const __half* __restrict__ wp = was + a * WASTRIDE;
    float s = 0.f;
#pragma unroll
    for (int k = 0; k < HID; k += 8) {
      const uint4 hv = *(const uint4*)(hp + k);
      const uint4 wv = *(const uint4*)(wp + k);
      const __half2* hh = (const __half2*)&hv;
      const __half2* ww = (const __half2*)&wv;
#pragma unroll
      for (int u = 0; u < 4; ++u) {
        const float2 hf = __half22float2(hh[u]), wf = __half22float2(ww[u]);
        s = fmaf(hf.x, wf.x, s);
        s = fmaf(hf.y, wf.y, s);
      }
    }
    s += P.ba[a];
    lg_s[a * E + ee] = s;
    if (last) P.lgo[(size_t)(env0 + ee) * NACT + a] = s;
  }
  __syncthreads();
}

// argmax of the 4 logits of env `tid`, ties to the lowest index (torch.argmax)
__device__ __forceinline__ int pick(const float* __restrict__ lg_s, int E, int tid) {
  const float l0 = lg_s[0 * E + tid], l1 = lg_s[1 * E + tid];
  const float l2 = lg_s[2 * E + tid], l3 = lg_s[3 * E + tid];
  int a = 0;
  float best = l0;
  if (l1 > best) { best = l1; a = 1; }
  if (l2 > best) { best = l2; a = 2; }
  if (l3 > best) { a = 3; }
  return a;
}

// mt_init leaves raw MT19937 words in ag/fd; tempering all 4n of them there
// costs it 90 cycles a generation, so the first reader does it instead -- once
// per env instead of once per env-step, and off the serial engine's critical
// path.  temper is a bijection and temper(0)%11 == 0, so the zeroed padding
// envs still start at (0,0).
__device__ __forceinline__ int mt_cook1(int w) {
  unsigned x = (unsigned)w;
  x ^= x >> 11;
  x ^= (x << 7) & 0x9d2c5680u;
  x ^= (x << 15) & 0xefc60000u;
  x ^= x >> 18;
  return (int)(x % 11u);
}
__device__ __forceinline__ int2 mt_cook(int2 v) {
  return make_int2(mt_cook1(v.x), mt_cook1(v.y));
}

// ---------------------------------------------------------------------------
// speculative kernel: one block owns E envs for the whole horizon
// ---------------------------------------------------------------------------
template <int E, int NG, int KP, int GRP>
__global__ __launch_bounds__(32 * (32 / NG), 1) void rollout_sp(const RParams P) {
  constexpr int WARPS = 32 / NG, THREADS = 32 * WARPS;
  static_assert(E % 16 == 0 && 32 % NG == 0 && (KP == 1 || KP == 2), "cfg");
  SH_DECL(E)
  (void)shi;

  constexpr int SHR = sh_rows(E);
  const int tid = threadIdx.x;
  const int H = P.horizon, nbsp = P.nbsp;
  float* __restrict__ stbase = P.state + (size_t)blockIdx.x * ((size_t)E * GATES);
  float* __restrict__ sst = (float*)(sraw + shmem_for(E));  // [SHR][SSTRIDE]

  for (int i = tid; i < E * 16; i += THREADS)
    oh[i] = __float2half((i & 15) == hpos(4) ? 1.0f : 0.0f);
  for (int i = tid; i < NACT * HID; i += THREADS) was[(i >> 8) * WASTRIDE + (i & 255)] = P.wa[i];
  __syncthreads();

  for (int c = blockIdx.x; c < P.nchunk; c += nbsp) {
    const int env0 = c * E;
    if (tid < E) {
      const int e = env0 + tid;
      const int2 a = mt_cook(*(const int2*)(P.ag + 2 * e));
      const int2 f = mt_cook(*(const int2*)(P.fd + 2 * e));
      es[tid] = make_int4(a.x, a.y, f.x, f.y);
      rs[tid] = (u64)((long long)e + P.seedbase);
      rws[tid] = 0.f;
    }

    for (int t = 0; t < H; ++t) {
      if (tid < E) {
        const int4 s = es[tid];
        const __half2 v01 = __floats2half2_rn((float)(s.z - s.x) * RBOARD,
                                              (float)(s.w - s.y) * RBOARD);
        const __half2 v23 = __floats2half2_rn((float)s.x * RBOARD1, (float)s.y * RBOARD1);
        uint4 pk;
        pk.x = *(const unsigned*)&v01;  // hpos(0), hpos(1) = 0, 1
        pk.y = 0u;
        pk.z = *(const unsigned*)&v23;  // hpos(2), hpos(3) = 4, 5
        pk.w = 0u;
        *(uint4*)(oh + tid * 16) = pk;
      }
      __syncthreads();

      step_layers<E, NG, KP, SHR, GRP>(P, hs, oh, was, lg_s, sst, stbase, E, 0, t == 0,
                                       t + 1 == H, env0);

      int hit = 0;
      if (tid < E) {
        const int a = pick(lg_s, E, tid);
        int4 s = es[tid];
        const int agx = min(max(s.x + ((a == 3) - (a == 2)), 0), BOARD - 1);
        const int agy = min(max(s.y + ((a == 1) - (a == 0)), 0), BOARD - 1);
        // speculation: some env eats on every step, so the shared LCG is
        // advanced twice per step unconditionally.  The two reductions mod 11
        // are not: a u64 remainder is a 128-bit multiply-high plus fixup, and
        // the greedy agents stop eating for good after ~10 steps, so past that
        // no lane of the warp takes this branch at all.
        const u64 r = lcg_step(rs[tid]), r2 = lcg_step(r);
        rs[tid] = r2;
        if (env0 + tid < P.n && agx == s.z && agy == s.w) {
          hit = 1;
          rws[tid] += 1.0f;
          s.z = (int)(r % BOARD);
          s.w = (int)(r2 % BOARD);
        }
        es[tid] = make_int4(agx, agy, s.z, s.w);
      }
      const int nh = __syncthreads_count(hit);
      if (tid == 0 && nh && t + 1 < H) atomicAdd(&P.ctr[t], nh);
    }

    if (tid < E) {
      const int e = env0 + tid;
      const int4 s = es[tid];
      *(longlong2*)(P.agout + 2 * e) = make_longlong2(s.x, s.y);
      P.rew[e] = rws[tid];
    }
  }

  // Verify.  A wrong K is only *observable* through a respawn, so the run is
  // exact iff no env ate after a step on which nobody ate: the speculation
  // K(t) = t+1 is right for every env that eats at t as long as every earlier
  // step also had an eater.  So ctr[] must be a (possibly empty) run of nonzero
  // counts followed by zeros -- which is what the greedy policy does, since
  // agents settle into fixed points and stop eating for good after ~10 steps.
  //
  // Sound because ctr[] is itself exact up to the first bad respawn: if the
  // first env to use a wrong K eats at step t, then ctr[t] > 0 and some earlier
  // ctr[s] == 0, and both of those counts were computed from exact state.
  // (Eating on the last step is not recorded: its respawn is never read.)
  __threadfence();
  if (tid == 0) {
    if (atomicAdd(P.done, 1) == nbsp - 1) {
      int ok = 1, seen = 0;
      const volatile int* cv = (const volatile int*)P.ctr;
      for (int i = H - 2; i >= 0; --i) {
        if (cv[i]) seen = 1;
        else if (seen) { ok = 0; break; }
      }
      *P.flag = ok;
    }
  }
}

// ---------------------------------------------------------------------------
// exact lockstep fallback.  E envs per block, NG gate blocks per warp
// (WARPS = 32/NG covers all 32 gate blocks), KP k-blocks per weight load.
// ---------------------------------------------------------------------------
template <int E, int NG, int KP, int GRP>
__global__ __launch_bounds__(32 * (32 / NG), 1) void rollout_ls(const RParams P) {
  constexpr int WARPS = 32 / NG, THREADS = 32 * WARPS;
  static_assert(E % 16 == 0 && 32 % NG == 0 && (KP == 1 || KP == 2), "cfg");
  if (*(const volatile int*)P.flag) return;  // rollout_sp already got it right
  SH_DECL(E)
  (void)es; (void)rs; (void)rws;

  const int tid = threadIdx.x;
  const int NP = P.np, H = P.horizon, nb = P.nblocks;
  int* __restrict__ hist = P.hist + (size_t)blockIdx.x * H;

  for (int i = tid; i < E * 16; i += THREADS)
    oh[i] = __float2half((i & 15) == hpos(4) ? 1.0f : 0.0f);
  for (int i = tid; i < NACT * HID; i += THREADS) was[(i >> 8) * WASTRIDE + (i & 255)] = P.wa[i];

  int Kacc = 0;  // thread 0 only: # steps <= tk that had a global hit
  int tk = -1;
  __syncthreads();

  for (int t = 0; t < H; ++t) {
    int blockhit = 0;
    const bool first = (t == 0);
    for (int c = blockIdx.x; c < P.nchunk; c += nb) {
      const int env0 = c * E;

      if (tid < E) {
        const int e = env0 + tid;
        int2 a = *(const int2*)(P.ag + 2 * e), f = *(const int2*)(P.fd + 2 * e);
        if (first) {  // raw MT words on the first touch; cook and store back
          a = mt_cook(a);
          f = mt_cook(f);
          *(int2*)(P.ag + 2 * e) = a;
          *(int2*)(P.fd + 2 * e) = f;
        }
        const __half2 v01 = __floats2half2_rn((float)(f.x - a.x) * RBOARD,
                                              (float)(f.y - a.y) * RBOARD);
        const __half2 v23 = __floats2half2_rn((float)a.x * RBOARD1, (float)a.y * RBOARD1);
        uint4 pk;
        pk.x = *(const unsigned*)&v01;
        pk.y = 0u;
        pk.z = *(const unsigned*)&v23;
        pk.w = 0u;
        *(uint4*)(oh + tid * 16) = pk;
        if (first) P.rew[e] = 0.f;
      }
      __syncthreads();

      step_layers<E, NG, KP, 0, GRP>(P, hs, oh, was, lg_s, nullptr, P.state, NP,
                                     (env0 >> 4) * 32, first, t + 1 == H, env0);

      int hit = 0, e = 0;
      if (tid < E) {
        e = env0 + tid;
        const int a = pick(lg_s, E, tid);
        const int2 ap = *(const int2*)(P.ag + 2 * e);
        const int agx = min(max(ap.x + ((a == 3) - (a == 2)), 0), BOARD - 1);
        const int agy = min(max(ap.y + ((a == 1) - (a == 0)), 0), BOARD - 1);
        *(int2*)(P.ag + 2 * e) = make_int2(agx, agy);
        if (t + 1 == H) *(longlong2*)(P.agout + 2 * e) = make_longlong2(agx, agy);
        const int2 fp = *(const int2*)(P.fd + 2 * e);
        if (e < P.n && agx == fp.x && agy == fp.y) {
          hit = 1;
          P.rew[e] += 1.0f;
        }
      }
      const int nh = __syncthreads_count(hit);
      if (nh) {
        blockhit = 1;
        if (t + 1 < H) {
          if (tid == 0) {
            while (tk < t - 1) {
              ++tk;
              int a = hist[tk];
              if (!a) {
                volatile const int* cv = (volatile const int*)P.lsc;
                int v;
                do { v = cv[tk]; } while ((v & 1023) != nb);
                a = (v >> 10) != 0;
              }
              Kacc += a;
            }
            shi[0] = Kacc + 1;
          }
          __syncthreads();
          if (hit) {
            const int K = shi[0];
            u64 r = (u64)((long long)e + P.seedbase);
            for (int i = 0; i < 2 * K - 1; ++i) r = lcg_step(r);
            const int nfx = (int)(r % BOARD);
            r = lcg_step(r);
            const int nfy = (int)(r % BOARD);
            *(int2*)(P.fd + 2 * e) = make_int2(nfx, nfy);
          }
          __syncthreads();
        }
      }
    }  // chunk loop

    if (tid == 0) {
      hist[t] = blockhit;
      if (t + 1 < H) atomicAdd(&P.lsc[t], 1 + (blockhit ? 1024 : 0));
    }
  }
}

// ---------------------------------------------------------------------------
// initial positions, drawn on the device
//
// The reference seeds a CPU MT19937 and calls torch.randint(0, 11) twice, so the
// initial agent then food coordinates are mt19937() % 11 in memory order.  Doing
// that on the host costs 46-477 us of the graded wall time (it dominated it at
// n=65536), so run the engine here instead: one block of 624 threads holds the
// state in shared memory and emits one 624-word block per twist.
//
// The in-place twist writes state[k] from state[k+397], which is the *new* value
// once k >= 227, so the naive parallel form is three dependent phases.  But the
// twist is GF(2)-linear, so substituting the earlier phases back in expresses
// every new word as an xor of at most three T(k) = tw(old[k], old[k+1]) terms
// and one old word:
//   k <   227: old[k+397] ^ T(k)
//   k <   454: old[k+170] ^ T(k-227) ^ T(k)
//   k <   623: old[k- 57] ^ T(k-454) ^ T(k-227) ^ T(k)
// with 623 the one word whose y takes the *new* low bits, from new[0].
//
// Each T is wanted by up to three threads, so publishing them through shared
// costs a second barrier per generation and saves two thirds of the twist
// arithmetic and three of the five shared loads: measured 566 -> 400 cycles a
// generation.  (Folding the T publication into the state publication to get
// back to one barrier was tried and is slower -- the lane at each warp's edge
// cannot see its neighbour's new word, and redoing that word costs more than
// the barrier.)  Tempering and the %11 are left to the rollout's first read of
// ag/fd, which is another 90 cycles a generation off this kernel.
#define MT_N 624
#define MT_M 397
#define MT_Q (MT_N - MT_M)  // 227

__device__ __forceinline__ unsigned mt_tw(unsigned u, unsigned v) {
  return (((u & 0x80000000u) | (v & 0x7fffffffu)) >> 1) ^
         ((unsigned)(-(int)(v & 1u)) & 0x9908b0dfu);
}

__global__ __launch_bounds__(MT_N) void mt_init(int* __restrict__ ag, int* __restrict__ fd,
                                                int* __restrict__ aux, int auxn, unsigned seed,
                                                int n, int np) {
  __shared__ unsigned buf[2][MT_N];
  __shared__ unsigned ts[MT_N];
  const int i = threadIdx.x;
  // padding envs are never scored but do get stepped; keep them deterministic
  for (int k = 2 * n + i; k < 2 * np; k += MT_N) {
    ag[k] = 0;
    fd[k] = 0;
  }
  for (int k = i; k < auxn; k += MT_N) aux[k] = 0;  // the rollout's counters
  if (i == 0) {  // init_with_uint32: a sequential recurrence, ~4 us
    unsigned s = seed;
    buf[0][0] = s;
    for (int j = 1; j < MT_N; ++j) {
      s = 1812433253u * (s ^ (s >> 30)) + (unsigned)j;
      buf[0][j] = s;
    }
  }
  __syncthreads();

  const int total = 4 * n;  // 2n agent coords then 2n food coords
  const int off = MT_M - (i >= MT_Q ? MT_Q : 0) - (i >= 2 * MT_Q ? MT_Q : 0);
  int cur = 0;
  for (int base = 0; base < total; base += MT_N) {
    const unsigned* __restrict__ src = buf[cur];
    unsigned* __restrict__ dst = buf[cur ^ 1];
    const unsigned t_i = mt_tw(src[i], src[i + 1 == MT_N ? 0 : i + 1]);
    ts[i] = t_i;
    __syncthreads();

    unsigned nv = t_i ^ src[i + off];
    if (i >= MT_Q) nv ^= ts[i - MT_Q];
    if (i >= 2 * MT_Q) nv ^= ts[i - 2 * MT_Q];
    // ts[623] was formed against src[0]; the real y there takes the new low bits
    if (i == MT_N - 1) nv ^= t_i ^ mt_tw(src[MT_N - 1], ts[0] ^ src[MT_M]);
    dst[i] = nv;

    const int k = base + i;
    if (k < total) {
      if (k < 2 * n)
        ag[k] = (int)nv;
      else
        fd[k - 2 * n] = (int)nv;
    }
    cur ^= 1;
    __syncthreads();
  }
}

// ---------------------------------------------------------------------------
// host dispatch
// ---------------------------------------------------------------------------
// Per-E tuple (NG, KP, GRP): NG gate blocks per warp, so the block is 1024/NG
// threads and a thread may hold 16*NG registers; KP k blocks per weight load;
// GRP gate blocks per matmul pass (A-fragment reuse, 12*MB*GRP accumulators).
// Overridable one field at a time from the command line for sweeps, e.g.
// -DNG64=4 -DGRP64=4 (nvcc rejects a comma-valued -D, so no tuple form).
// NG=4 (256 threads) was measured 27% slower than NG=2 at E=64: one block per
// SM means 8 warps, i.e. 2 per scheduler, which cannot hide the L2 latency of
// the weight stream.  So every config runs 512 threads / 128 registers, and GRP
// is capped by that: 12*MB*GRP accumulators plus ~50 other live registers.
#ifndef NG16
#define NG16 2
#endif
#ifndef KP16
#define KP16 2
#endif
#ifndef GRP16
#define GRP16 2
#endif
#ifndef NG32
#define NG32 2
#endif
#ifndef KP32
#define KP32 2
#endif
#ifndef GRP32
#define GRP32 2
#endif
#ifndef NG48
#define NG48 2
#endif
#ifndef KP48
#define KP48 2
#endif
#ifndef GRP48
#define GRP48 2
#endif
#ifndef NG64
#define NG64 2
#endif
#ifndef KP64
#define KP64 2
#endif
#ifndef GRP64
#define GRP64 1
#endif
#ifndef NG80
#define NG80 2
#endif
#ifndef KP80
#define KP80 2
#endif
#ifndef GRP80
#define GRP80 1
#endif
#ifndef NG96
#define NG96 2
#endif
#ifndef KP96
#define KP96 2
#endif
#ifndef GRP96
#define GRP96 1
#endif
#define CFG_LIST                \
  CFG(16, NG16, KP16, GRP16)    \
  CFG(32, NG32, KP32, GRP32)    \
  CFG(48, NG48, KP48, GRP48)    \
  CFG(64, NG64, KP64, GRP64)    \
  CFG(80, NG80, KP80, GRP80)    \
  CFG(96, NG96, KP96, GRP96)

int max_resident(int E) {
  int nb = 0, nb2 = 0, dev = 0;
  cudaGetDevice(&dev);
  cudaDeviceProp prop;
  cudaGetDeviceProperties(&prop, dev);
  size_t shm = shmem_for(E), sh1 = shmem_sp(E);
  switch (E) {
#define CFG(EE, NG, KP, GRP)                                                             \
  case EE: {                                                                             \
    auto k1 = rollout_sp<EE, NG, KP, GRP>;                                               \
    auto k2 = rollout_ls<EE, NG, KP, GRP>;                                               \
    cudaFuncSetAttribute(k1, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)sh1);      \
    cudaFuncSetAttribute(k2, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shm);      \
    cudaOccupancyMaxActiveBlocksPerMultiprocessor(&nb, k1, 32 * (32 / NG), sh1);          \
    cudaOccupancyMaxActiveBlocksPerMultiprocessor(&nb2, k2, 32 * (32 / NG), shm);         \
    break;                                                                               \
  }
    CFG_LIST
#undef CFG
    default: TORCH_CHECK(false, "bad E ", E);
  }
  return (nb < nb2 ? nb : nb2) * prop.multiProcessorCount;
}

// Split out of rollout() so run() can fire it as soon as the position buffer
// exists: the engine is one block on one SM for 10-80 us, which is long enough
// to hide the rest of run()'s allocations and the rollout launch behind it.
// It also clears aux, saving a separate fill kernel.
void mt_seed(at::Tensor ag, at::Tensor fd, at::Tensor aux, long long seed, long long n,
             long long np) {
  mt_init<<<1, MT_N, 0, at::cuda::getCurrentCUDAStream()>>>(
      ag.data_ptr<int>(), fd.data_ptr<int>(), aux.data_ptr<int>(), (int)aux.numel(),
      (unsigned)seed, (int)n, (int)np);
}

void rollout(at::Tensor w0, at::Tensor wg, at::Tensor wa, at::Tensor ba, at::Tensor state,
             at::Tensor ag, at::Tensor agout, at::Tensor fd, at::Tensor rew, at::Tensor lgo,
             at::Tensor aux, at::Tensor hist, long long seed, long long n, long long np,
             long long nchunk, long long nblocks, long long nbsp, long long horizon,
             long long E) {
  const long long seedbase = seed * 10007;
  RParams P;
  P.w0 = (const __half*)w0.data_ptr<at::Half>();
  P.wg = (const __half*)wg.data_ptr<at::Half>();
  P.wa = (const __half*)wa.data_ptr<at::Half>();
  P.ba = ba.data_ptr<float>();
  P.state = state.data_ptr<float>();
  P.ag = ag.data_ptr<int>();
  P.agout = (long long*)agout.data_ptr<int64_t>();
  P.fd = fd.data_ptr<int>();
  P.rew = rew.data_ptr<float>();
  P.lgo = lgo.data_ptr<float>();
  int* a = aux.data_ptr<int>();
  P.ctr = a;
  P.lsc = a + horizon;
  P.done = a + 2 * horizon;
  P.flag = a + 2 * horizon + 1;
  P.hist = hist.data_ptr<int>();
  P.seedbase = seedbase;
  P.n = (int)n;
  P.np = (int)np;
  P.nchunk = (int)nchunk;
  P.nblocks = (int)nblocks;
  P.nbsp = (int)nbsp;
  P.horizon = (int)horizon;
  size_t shm = shmem_for((int)E), sh1 = shmem_sp((int)E);
  cudaStream_t s = at::cuda::getCurrentCUDAStream();
  switch ((int)E) {
#define CFG(EE, NG, KP, GRP)                                                             \
  case EE: {                                                                             \
    auto k1 = rollout_sp<EE, NG, KP, GRP>;                                               \
    auto k2 = rollout_ls<EE, NG, KP, GRP>;                                               \
    cudaFuncSetAttribute(k1, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)sh1);      \
    cudaFuncSetAttribute(k2, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shm);      \
    if (nbsp > 0) k1<<<(int)nbsp, 32 * (32 / NG), sh1, s>>>(P);                           \
    k2<<<(int)nblocks, 32 * (32 / NG), shm, s>>>(P);                                      \
    break;                                                                               \
  }
    CFG_LIST
#undef CFG
    default: TORCH_CHECK(false, "bad E ", E);
  }
}

// ---------------------------------------------------------------------------
// policy_forward (reference-order math, correctness path; not perf critical)
// ---------------------------------------------------------------------------
__global__ __launch_bounds__(256) void pf_kernel(const float* __restrict__ obs,
                                                 const float* __restrict__ st,
                                                 const float* __restrict__ wenc,
                                                 const float* __restrict__ benc,
                                                 const float* __restrict__ wgt,
                                                 const float* __restrict__ wa,
                                                 const float* __restrict__ ba,
                                                 const float* __restrict__ wv,
                                                 const float* __restrict__ bv,
                                                 float* __restrict__ logits,
                                                 float* __restrict__ nst,
                                                 float* __restrict__ val, int n, int tf32) {
  const int e = blockIdx.x;
  const int j = threadIdx.x;
  if (e >= n) return;
  const bool T = tf32 != 0;
  __shared__ float hs[HID];
  float ob[4];
#pragma unroll
  for (int i = 0; i < 4; ++i) ob[i] = rnd(obs[(size_t)e * 4 + i], T);
  float h = 0.f;
#pragma unroll
  for (int i = 0; i < 4; ++i) h = fmaf(rnd(wenc[i * HID + j], T), ob[i], h);
  h += benc[j];
  hs[j] = h;
  __syncthreads();
  for (int l = 0; l < 3; ++l) {
    const float* W = wgt + (size_t)l * HID * GATES;
    float zh = 0.f, zg = 0.f, zp = 0.f;
    for (int k = 0; k < HID; ++k) {
      float hv = rnd(hs[k], T);
      const float* wp = W + (size_t)k * GATES + j;
      zh = fmaf(rnd(wp[0], T), hv, zh);
      zg = fmaf(rnd(wp[HID], T), hv, zg);
      zp = fmaf(rnd(wp[2 * HID], T), hv, zp);
    }
    float s = st[((size_t)e * 3 + l) * HID + j];
    float o = s + sgm(zg) * (tanhf(zh) - s);
    float p = sgm(zp);
    float hn = p * o + (1.f - p) * hs[j];
    nst[((size_t)e * 3 + l) * HID + j] = o;
    __syncthreads();
    hs[j] = hn;
    __syncthreads();
  }
  if (j < NACT) {
    float s = 0.f;
    for (int k = 0; k < HID; ++k) s = fmaf(rnd(wa[j * HID + k], T), rnd(hs[k], T), s);
    logits[(size_t)e * NACT + j] = s + ba[j];
  } else if (j == 4) {
    float s = 0.f;
    for (int k = 0; k < HID; ++k) s = fmaf(rnd(wv[k], T), rnd(hs[k], T), s);
    val[e] = s + bv[0];
  }
}

std::vector<at::Tensor> policy_forward_cuda(at::Tensor obs, at::Tensor state, at::Tensor wenc,
                                            at::Tensor benc, at::Tensor wgt, at::Tensor wa,
                                            at::Tensor ba, at::Tensor wv, at::Tensor bv,
                                            int64_t tf32) {
  int n = (int)obs.size(0);
  auto opt = obs.options();
  at::Tensor logits = at::empty({n, NACT}, opt);
  at::Tensor nst = at::empty({n, 3, HID}, opt);
  at::Tensor val = at::empty({n}, opt);
  cudaStream_t s = at::cuda::getCurrentCUDAStream();
  pf_kernel<<<n, HID, 0, s>>>(obs.data_ptr<float>(), state.data_ptr<float>(),
                              wenc.data_ptr<float>(), benc.data_ptr<float>(),
                              wgt.data_ptr<float>(), wa.data_ptr<float>(), ba.data_ptr<float>(),
                              wv.data_ptr<float>(), bv.data_ptr<float>(), logits.data_ptr<float>(),
                              nst.data_ptr<float>(), val.data_ptr<float>(), n, (int)tf32);
  return {logits, nst, val};
}

// ---------------------------------------------------------------------------
// env_step (reference-exact, used by check.py)
// ---------------------------------------------------------------------------
__global__ void es_move(const float* __restrict__ agent, const float* __restrict__ food,
                        const int64_t* __restrict__ act, float* __restrict__ nag,
                        float* __restrict__ rew, int* __restrict__ anyhit, int n) {
  int e = blockIdx.x * blockDim.x + threadIdx.x;
  if (e >= n) return;
  float x = agent[(size_t)e * 2 + 0], y = agent[(size_t)e * 2 + 1];
  int64_t a = act[e];
  x += (float)((a == 3) - (a == 2));
  y += (float)((a == 1) - (a == 0));
  x = fminf(fmaxf(x, 0.f), (float)(BOARD - 1));
  y = fminf(fmaxf(y, 0.f), (float)(BOARD - 1));
  nag[(size_t)e * 2 + 0] = x;
  nag[(size_t)e * 2 + 1] = y;
  int hit = (x == food[(size_t)e * 2 + 0] && y == food[(size_t)e * 2 + 1]);
  rew[e] = (float)hit;
  if (hit) atomicOr(anyhit, 1);
}

__global__ void es_food(const float* __restrict__ food, const float* __restrict__ rew,
                        const int64_t* __restrict__ rng, float* __restrict__ nfood,
                        int64_t* __restrict__ nrng, const int* __restrict__ anyhit, int n) {
  int e = blockIdx.x * blockDim.x + threadIdx.x;
  if (e >= n) return;
  float fx = food[(size_t)e * 2 + 0], fy = food[(size_t)e * 2 + 1];
  int64_t r = rng[e];
  if (*anyhit) {
    u64 u = lcg_step((u64)r);
    float nx = (float)(u % BOARD);
    u = lcg_step(u);
    float ny = (float)(u % BOARD);
    r = (int64_t)u;
    if (rew[e] != 0.f) { fx = nx; fy = ny; }
  }
  nfood[(size_t)e * 2 + 0] = fx;
  nfood[(size_t)e * 2 + 1] = fy;
  nrng[e] = r;
}

std::vector<at::Tensor> env_step_cuda(at::Tensor agent, at::Tensor food, at::Tensor act,
                                      at::Tensor rng) {
  int n = (int)agent.size(0);
  auto fopt = agent.options();
  at::Tensor nag = at::empty_like(agent);
  at::Tensor nfd = at::empty_like(food);
  at::Tensor rew = at::empty({n}, fopt);
  at::Tensor nrng = at::empty_like(rng);
  at::Tensor anyhit = at::zeros({1}, agent.options().dtype(at::kInt));
  cudaStream_t s = at::cuda::getCurrentCUDAStream();
  int thr = 256, blk = (n + thr - 1) / thr;
  es_move<<<blk, thr, 0, s>>>(agent.data_ptr<float>(), food.data_ptr<float>(),
                              act.data_ptr<int64_t>(), nag.data_ptr<float>(),
                              rew.data_ptr<float>(), anyhit.data_ptr<int>(), n);
  es_food<<<blk, thr, 0, s>>>(food.data_ptr<float>(), rew.data_ptr<float>(),
                              rng.data_ptr<int64_t>(), nfd.data_ptr<float>(),
                              nrng.data_ptr<int64_t>(), anyhit.data_ptr<int>(), n);
  return {nag, nfd, rew, nrng};
}

__global__ void mt_cook_all(int* __restrict__ p, int m) {
  const int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i < m) p[i] = mt_cook1(p[i]);
}

// exposed for validation against torch.randint on a seeded CPU generator
at::Tensor init_positions(int64_t n, int64_t np, int64_t seed) {
  at::Tensor t = at::empty({2, np, 2}, at::TensorOptions().dtype(at::kInt).device(at::kCUDA));
  at::Tensor a = at::empty({4}, at::TensorOptions().dtype(at::kInt).device(at::kCUDA));
  int* p = t.data_ptr<int>();
  cudaStream_t s = at::cuda::getCurrentCUDAStream();
  mt_init<<<1, MT_N, 0, s>>>(p, p + 2 * np, a.data_ptr<int>(), 4, (unsigned)seed, (int)n, (int)np);
  const int m = (int)(4 * np);
  mt_cook_all<<<(m + 255) / 256, 256, 0, s>>>(p, m);  // mt_init now stores raw words
  return t;
}

std::vector<int64_t> prof_read(bool reset) {
  u64 host[8];
  cudaMemcpyFromSymbol(host, g_prof, sizeof(host), 0, cudaMemcpyDeviceToHost);
  if (reset) {
    u64 z[8] = {0, 0, 0, 0, 0, 0, 0, 0};
    cudaMemcpyToSymbol(g_prof, z, sizeof(z), 0, cudaMemcpyHostToDevice);
  }
  std::vector<int64_t> out(8);
  for (int i = 0; i < 8; ++i) out[i] = (int64_t)host[i];
  return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("prof_read", &prof_read);
  m.def("init_positions", &init_positions, "device-side MT19937 initial positions");
  m.def("mt_seed", &mt_seed, "device-side MT19937 initial positions + aux clear");
  m.def("rollout", &rollout, "fused rollout");
  m.def("max_resident", &max_resident, "max resident blocks");
  m.def("sh_rows", [](int64_t E) { return (int64_t)sh_rows((int)E); },
        "recurrent state rows (of 384) kept in shared memory by rollout_sp");
  m.def("policy_forward", &policy_forward_cuda, "policy forward");
  m.def("env_step", &env_step_cuda, "env step");
}

20260725_023503_or-opus_anthropic_claude-opus-5_04_grid_mingru_sps