KernelBench cuda · RTX PRO 6000

Grid + MinGRU SPS DeepSeek V4.1 Flash

28.6%geomean peak fraction across shapes

manually audited: clean

DeepSeek V4.1 Flash wrote a hand-rolled fp32 SM120 rollout: a cp.async double-buffered register-tile GEMM with the 3-layer MinGRU highway fused into the accumulators, the whole horizon replayed from a CUDA graph. It reproduces the reference environment down to the batch-global any-hit RNG gating that two other cells on this problem shortcut.

harnessdeepseek-claudeagent session3h 51mtotal wall3h 52mcheck25sbenchmark2soutput tokens532,911cost$31.54gpu-lock wait2h 5mgpu-lock held29mregimethroughput

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)
"""CUDA implementation of the vectorized grid-foraging env + 3-layer MinGRU policy.

Reference semantics live in reference.py.  All heavy math runs in hand written
CUDA kernels (see the CUDA source below); Python only launches them.

Layout
------
* ``enc_kernel``   : obs (N,4)            -> h0 (N,256)      (Linear 4->256)
* ``layer_kernel`` : fused  h @ W_l^T  (256->768)  +  MinGRU highway update.
  Computes a (BM x BN) tile of the 768 gate pre-activations entirely in
  registers, applies the elementwise MinGRU recurrence from the register tile
  and writes h_next / state directly -- the 768 wide ``gates`` matrix never
  touches global memory.
  The 768 gate columns are re-ordered on the host so that column ``c`` is
  ``(hidden i, gate g)`` with ``c = 3*i + g``; a thread owning ``TN`` columns
  (TN a multiple of 3) therefore owns complete (zh,zg,zp) triples.
* ``head_kernel``  : h3 -> logits (N,4), value (N,), greedy actions
* env kernels      : move / hit / LCG respawn, matching reference.env_step
"""
from __future__ import annotations

import os

import torch
import torch.nn as nn

from torch.utils.cpp_extension import load_inline

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

os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")

_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <c10/cuda/CUDAStream.h>
#include <cstdint>

#define BOARD 11
#define OBS_DIM 4
#define HIDDEN 256
#define GRU_LAYERS 3
#define NUM_ACTIONS 4
#define GRU_OUT 768

// ---------------------------------------------------------------------------
// LCG identical to reference._lcg_step (int64 wrapping mult, mask to 63 bits)
// ---------------------------------------------------------------------------
__device__ __forceinline__ long long lcg_step(long long r) {
    unsigned long long u = (unsigned long long)r;
    u = u * 6364136223846793005ULL + 1ULL;
    return (long long)(u & 0x7FFFFFFFFFFFFFFFULL);
}

__device__ __forceinline__ void cpa16(void* s, const void* g) {
    unsigned su = (unsigned)__cvta_generic_to_shared(s);
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(su), "l"(g));
}
__device__ __forceinline__ void cp_commit() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void cp_wait0() { asm volatile("cp.async.wait_group 0;"); }

// ---------------------------------------------------------------------------
// Layer 0 encoder: h0[e][j] = b_enc[j] + sum_k obs[e][k] * w_enc[j][k]
// ---------------------------------------------------------------------------
__global__ void enc_kernel(const float* __restrict__ obs,
                           const float* __restrict__ w_enc,
                           const float* __restrict__ b_enc,
                           float* __restrict__ h0, int N) {
    long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    long long total = (long long)N * HIDDEN;
    if (i >= total) return;
    int e = (int)(i >> 8);
    int j = (int)(i & 255);
    const float* o = obs + (long long)e * OBS_DIM;
    const float* w = w_enc + (long long)j * OBS_DIM;
    float s = b_enc[j];
    #pragma unroll
    for (int k = 0; k < OBS_DIM; ++k) s = fmaf(o[k], w[k], s);
    h0[i] = s;
}

// ---------------------------------------------------------------------------
// Fused  (BM x BN) x K  GEMM  +  MinGRU highway recurrence.
//
//   h_in   (M, 256)         input activations (layer L-1 output / enc output)
//   st     (M, 3, 256)      GRU state, row stride 768 (updated in place)
//   Wt     (256, 768)       weight, column c = 3*i + g  ->  w_gru[l][g*256+i]
//   h_out  (M, 256)         gated output activations
// ---------------------------------------------------------------------------
template<int BM,int BN,int BK,int TM,int TN,int TH,int PAD,int MINB=0>
__global__ void __launch_bounds__(TH, MINB > 0 ? MINB : 1)
layer_kernel(const float* __restrict__ h_in,
             float* __restrict__ st,
             const float* __restrict__ Wt,
             float* __restrict__ h_out,
             int M, int layer)
{
    constexpr int TX = BN / TN, TY = BM / TM;
    constexpr int BKP = BK + PAD;
    constexpr int NA = BM * (BK / 4) / TH;
    constexpr int NB = BK * (BN / 4) / TH;
    constexpr int HU = TN / 3;                 // hidden units owned per thread
    extern __shared__ float smem[];
    float* As = smem;                          // [2][BM][BKP]
    float* Bs = smem + 2 * BM * BKP;           // [2][BK][BN]

    const int tx = threadIdx.x % TX, ty = threadIdx.x / TX;
    const int m0 = blockIdx.y * BM, n0 = blockIdx.x * BN;

    #pragma unroll
    for (int s = 0; s < NA; ++s) {
        int idx = threadIdx.x + s * TH;
        int m = idx / (BK / 4), u = idx % (BK / 4);
        cpa16(&As[m * BKP + u * 4], &h_in[(long long)(m0 + m) * HIDDEN + u * 4]);
    }
    #pragma unroll
    for (int s = 0; s < NB; ++s) {
        int idx = threadIdx.x + s * TH;
        int k = idx / (BN / 4), v = idx % (BN / 4);
        cpa16(&Bs[k * BN + v * 4], &Wt[(long long)k * GRU_OUT + n0 + v * 4]);
    }
    cp_commit();
    cp_wait0();
    __syncthreads();

    // The recurrence operands (h_in row, prior state) are only needed by the
    // epilogue.  Reading them up front costs TM*8 live registers across the
    // whole k-loop, which measurably throttles the FFMA stream; they are warm
    // in L1/L2 by then anyway.
    const int hb = n0 / 3 + tx * HU;
    const int sbase = layer * HIDDEN;
    float hprev[TM][HU], stval[TM][HU];

    float acc[TM][TN];
    #pragma unroll
    for (int i = 0; i < TM; ++i)
        #pragma unroll
        for (int j = 0; j < TN; ++j) acc[i][j] = 0.f;

    constexpr int NV = TN / 4;
    int cur = 0;
    for (int k0 = 0; k0 < HIDDEN; k0 += BK) {
        int nxt = k0 + BK;
        if (nxt < HIDDEN) {
            float* a2 = As + (cur ^ 1) * BM * BKP;
            float* b2 = Bs + (cur ^ 1) * BK * BN;
            #pragma unroll
            for (int s = 0; s < NA; ++s) {
                int idx = threadIdx.x + s * TH;
                int m = idx / (BK / 4), u = idx % (BK / 4);
                cpa16(&a2[m * BKP + u * 4], &h_in[(long long)(m0 + m) * HIDDEN + nxt + u * 4]);
            }
            #pragma unroll
            for (int s = 0; s < NB; ++s) {
                int idx = threadIdx.x + s * TH;
                int k = idx / (BN / 4), v = idx % (BN / 4);
                cpa16(&b2[k * BN + v * 4], &Wt[(long long)(nxt + k) * GRU_OUT + n0 + v * 4]);
            }
            cp_commit();
        }
        const float* a1 = As + cur * BM * BKP;
        const float* b1 = Bs + cur * BK * BN;
        #pragma unroll
        for (int kk = 0; kk < BK; ++kk) {
            float a[TM];
            #pragma unroll
            for (int i = 0; i < TM; ++i) a[i] = a1[(ty * TM + i) * BKP + kk];
            float4 bv[NV];
            #pragma unroll
            for (int v = 0; v < NV; ++v) bv[v] = *(const float4*)&b1[kk * BN + tx * TN + v * 4];
            #pragma unroll
            for (int i = 0; i < TM; ++i)
                #pragma unroll
                for (int v = 0; v < NV; ++v) {
                    acc[i][v*4+0] = fmaf(a[i], bv[v].x, acc[i][v*4+0]);
                    acc[i][v*4+1] = fmaf(a[i], bv[v].y, acc[i][v*4+1]);
                    acc[i][v*4+2] = fmaf(a[i], bv[v].z, acc[i][v*4+2]);
                    acc[i][v*4+3] = fmaf(a[i], bv[v].w, acc[i][v*4+3]);
                }
        }
        if (nxt < HIDDEN) { cp_wait0(); __syncthreads(); }
        cur ^= 1;
    }

    // ---- MinGRU highway recurrence, straight out of the register tile ----
    #pragma unroll
    for (int i = 0; i < TM; ++i) {
        const int env = m0 + ty * TM + i;
        const float* hip = h_in + (long long)env * HIDDEN + hb;
        const float* sip = st + (long long)env * (GRU_LAYERS * HIDDEN) + sbase + hb;
        #pragma unroll
        for (int q = 0; q < HU / 4; ++q) {
            const float4 hv = *(const float4*)(hip + 4 * q);
            const float4 sv = *(const float4*)(sip + 4 * q);
            hprev[i][4*q+0]=hv.x; hprev[i][4*q+1]=hv.y;
            hprev[i][4*q+2]=hv.z; hprev[i][4*q+3]=hv.w;
            stval[i][4*q+0]=sv.x; stval[i][4*q+1]=sv.y;
            stval[i][4*q+2]=sv.z; stval[i][4*q+3]=sv.w;
        }
    }
    #pragma unroll
    for (int i = 0; i < TM; ++i) {
        int env = m0 + ty * TM + i;
        float* stp = st + (long long)env * (GRU_LAYERS * HIDDEN) + sbase + hb;
        float* hp  = h_out + (long long)env * HIDDEN + hb;
        float ov[HU], wv[HU];
        #pragma unroll
        for (int u = 0; u < HU; ++u) {
            float zh = acc[i][3*u+0], zg = acc[i][3*u+1], zp = acc[i][3*u+2];
            float s  = stval[i][u];
            float out = s + (1.f / (1.f + expf(-zg))) * (tanhf(zh) - s);
            float p   = 1.f / (1.f + expf(-zp));
            ov[u] = out;
            wv[u] = p * out + (1.f - p) * hprev[i][u];
        }
        #pragma unroll
        for (int q = 0; q < HU / 4; ++q) {
            *(float4*)(stp + 4 * q) = make_float4(ov[4*q], ov[4*q+1], ov[4*q+2], ov[4*q+3]);
            *(float4*)(hp  + 4 * q) = make_float4(wv[4*q], wv[4*q+1], wv[4*q+2], wv[4*q+3]);
        }
    }
}

// ---------------------------------------------------------------------------
// heads: logits = w_a @ h + b_a ; value = w_v @ h + b_v ; actions = argmax
// one warp per env
// ---------------------------------------------------------------------------
// One warp per env.  The env's 256 activations are pulled into registers once
// and reused by all five dot products (four logits + value); reading h every
// pass instead costs 5x the load traffic and forces a serial FMA chain per
// action.  The head weights are staged in shared memory once per block.
__global__ void head_kernel(const float* __restrict__ h,
                            const float* __restrict__ w_a,
                            const float* __restrict__ b_a,
                            const float* __restrict__ w_v,
                            const float* __restrict__ b_v,
                            float* __restrict__ logits,
                            float* __restrict__ value,
                            long long* __restrict__ actions,
                            int N) {
    constexpr int PER = HIDDEN / 32;          // activations per lane
    __shared__ float sw_a[NUM_ACTIONS * HIDDEN];
    __shared__ float sw_v[HIDDEN];
    __shared__ float sb_a[NUM_ACTIONS], sb_v;
    for (int i = threadIdx.x; i < NUM_ACTIONS * HIDDEN; i += blockDim.x) sw_a[i] = w_a[i];
    for (int i = threadIdx.x; i < HIDDEN; i += blockDim.x) sw_v[i] = w_v[i];
    if (threadIdx.x < NUM_ACTIONS) sb_a[threadIdx.x] = b_a[threadIdx.x];
    if (threadIdx.x == 0) sb_v = b_v[0];
    __syncthreads();

    const int lane = threadIdx.x & 31;
    const int e = blockIdx.x * (blockDim.x >> 5) + (threadIdx.x >> 5);
    if (e >= N) return;
    const float* hk = h + (long long)e * HIDDEN + lane;
    float hv[PER];
    #pragma unroll
    for (int i = 0; i < PER; ++i) hv[i] = hk[32 * i];

    float l[NUM_ACTIONS];
    #pragma unroll
    for (int a = 0; a < NUM_ACTIONS; ++a) {
        const float* wa = sw_a + a * HIDDEN + lane;
        float s = 0.f;
        #pragma unroll
        for (int i = 0; i < PER; ++i) s = fmaf(hv[i], wa[32 * i], s);
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
        s += sb_a[a];
        l[a] = s;
        if (lane == 0) logits[(long long)e * NUM_ACTIONS + a] = s;
    }
    float v = 0.f;
    #pragma unroll
    for (int i = 0; i < PER; ++i) v = fmaf(hv[i], sw_v[lane + 32 * i], v);
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) v += __shfl_down_sync(0xffffffffu, v, o);
    if (lane == 0) {
        value[e] = v + sb_v;
        int best = 0;
        #pragma unroll
        for (int a = 1; a < NUM_ACTIONS; ++a) if (l[a] > l[best]) best = a;
        actions[e] = (long long)best;
    }
}

// ---------------------------------------------------------------------------
// env_step: move + hit detection + global any-hit flag (reference semantics)
// ---------------------------------------------------------------------------
__global__ void env_move_kernel(const float* __restrict__ agent,
                                const float* __restrict__ food,
                                const long long* __restrict__ actions,
                                float* __restrict__ agent_out,
                                unsigned char* __restrict__ hit,
                                int* __restrict__ flag,
                                int N) {
    int e = blockIdx.x * blockDim.x + threadIdx.x;
    if (e >= N) return;
    long long a = actions[e];
    float x = agent[(long long)e * 2 + 0];
    float y = agent[(long long)e * 2 + 1];
    if (a == 0) y -= 1.f;
    else if (a == 1) y += 1.f;
    else if (a == 2) x -= 1.f;
    else x += 1.f;
    x = fminf(fmaxf(x, 0.f), (float)(BOARD - 1));
    y = fminf(fmaxf(y, 0.f), (float)(BOARD - 1));
    agent_out[(long long)e * 2 + 0] = x;
    agent_out[(long long)e * 2 + 1] = y;
    bool h = (x == food[(long long)e * 2 + 0]) && (y == food[(long long)e * 2 + 1]);
    hit[e] = h ? 1 : 0;
    if (h) atomicOr(flag, 1);
}

__global__ void env_rng_kernel(const unsigned char* __restrict__ hit,
                               float* __restrict__ food,
                               long long* __restrict__ rng,
                               const int* __restrict__ flag,
                               int* __restrict__ other_flag,
                               float* __restrict__ reward,
                               const float* __restrict__ agent,
                               float* __restrict__ obs,
                               int N) {
    __shared__ int sflag;
    if (threadIdx.x == 0) {
        sflag = *flag;
        if (blockIdx.x == 0) *other_flag = 0;
    }
    __syncthreads();
    int any = sflag;
    int e = blockIdx.x * blockDim.x + threadIdx.x;
    if (e >= N) return;
    if (any) {
        long long r1 = lcg_step(rng[e]);
        long long r2 = lcg_step(r1);
        if (hit[e]) {
            food[(long long)e * 2 + 0] = (float)(r1 % BOARD);
            food[(long long)e * 2 + 1] = (float)(r2 % BOARD);
        }
        rng[e] = r2;
    }
    reward[e] += (float)hit[e];

    // obs for the next step, from the post-move agent (this same thread owns the
    // row, so the food it may have just respawned is already visible).
    float ax = agent[(long long)e * 2 + 0];
    float ay = agent[(long long)e * 2 + 1];
    float* o = obs + (long long)e * OBS_DIM;
    o[0] = (food[(long long)e * 2 + 0] - ax) / (float)BOARD;
    o[1] = (food[(long long)e * 2 + 1] - ay) / (float)BOARD;
    o[2] = ax / (float)(BOARD - 1);
    o[3] = ay / (float)(BOARD - 1);
}

// ---------------------------------------------------------------------------
// obs from agent/food (also used to feed the next policy step)
// ---------------------------------------------------------------------------
__global__ void obs_kernel(const float* __restrict__ agent,
                           const float* __restrict__ food,
                           float* __restrict__ obs, int N) {
    int e = blockIdx.x * blockDim.x + threadIdx.x;
    if (e >= N) return;
    float ax = agent[(long long)e * 2 + 0];
    float ay = agent[(long long)e * 2 + 1];
    float fx = food[(long long)e * 2 + 0];
    float fy = food[(long long)e * 2 + 1];
    float* o = obs + (long long)e * OBS_DIM;
    o[0] = (fx - ax) / (float)BOARD;
    o[1] = (fy - ay) / (float)BOARD;
    o[2] = ax / (float)(BOARD - 1);
    o[3] = ay / (float)(BOARD - 1);
}

// ---------------------------------------------------------------------------
// host launchers
// ---------------------------------------------------------------------------
static inline int cdiv(long long a, int b) { return (int)((a + b - 1) / b); }
#define GSTR at::cuda::getCurrentCUDAStream()

void enc_launch(torch::Tensor obs, torch::Tensor w_enc, torch::Tensor b_enc,
                torch::Tensor h0) {
    int N = (int)obs.size(0);
    if (N == 0) return;
    enc_kernel<<<cdiv((long long)N * HIDDEN, 256), 256, 0, GSTR>>>(
        obs.data_ptr<float>(), w_enc.data_ptr<float>(), b_enc.data_ptr<float>(),
        h0.data_ptr<float>(), N);
}

// sm_120 exposes only 100 KB of shared memory per SM (not the 227 KB of the
// datacentre parts) and 64 K registers, so the resident-block count -- hence
// how much latency the SM can hide -- is set by BM*(BK+PAD)+BK*BN.  The tile
// parameters below are fixed to the winner of an end-to-end sweep.
template<int BM,int BN,int BK,int TM,int TN,int TH,int PAD,int MINB=0>
static void layer_run(const float* h, float* st, const float* W, float* out,
                      int M, int layer) {
    constexpr int SMEM = 2 * (BM * (BK + PAD) + BK * BN) * 4;
    static int ok = 0;
    if (ok == 0) {
        cudaError_t e = cudaFuncSetAttribute(layer_kernel<BM,BN,BK,TM,TN,TH,PAD,MINB>,
                             cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
        ok = (e == cudaSuccess) ? 1 : -1;
    }
    if (ok < 0) return;
    dim3 grid(GRU_OUT / BN, M / BM);
    layer_kernel<BM,BN,BK,TM,TN,TH,PAD,MINB><<<grid, TH, SMEM, GSTR>>>(h, st, W, out, M, layer);
}

// Winner of an end-to-end sweep of 18 tile shapes (see scratch/sweep.py):
// BM=64 BN=96 BK=16 TM=4 TN=12 TH=128 PAD=0 with __launch_bounds__ minBlocks=3.
//
// BK=16 with no padding is the minimum shared memory that still lets ptxas
// vectorise the B operand: at BN=96/TN=12 the B-tile stride of 12 floats maps
// threads 0..7 onto banks {0-3,12-15,24-27,4-7,16-19,28-31,8-11,20-23}, i.e.
// all 32 banks conflict-free, while A is a broadcast -- so padding only wasted
// shared memory and cost a resident block.
//
// minBlocks=3 lets ptxas spend up to 170 registers instead of capping it at the
// 128 that 4 blocks would allow.  Every explicit MINB (2/3/5) beat the implicit
// 128-register allocation, most visibly on the smallest shape (4096/32), where
// the 128-register build ran a reproducible ~6% slower.
void layer_launch(torch::Tensor h_in, torch::Tensor st, torch::Tensor Wt,
                  torch::Tensor h_out, int layer) {
    int M = (int)h_in.size(0);
    if (M == 0) return;
    layer_run<64, 96, 16, 4, 12, 128, 0, 3>(
        h_in.data_ptr<float>(), st.data_ptr<float>(), Wt.data_ptr<float>(),
        h_out.data_ptr<float>(), M, layer);
}

void head_launch(torch::Tensor h, torch::Tensor w_a, torch::Tensor b_a,
                 torch::Tensor w_v, torch::Tensor b_v, torch::Tensor logits,
                 torch::Tensor value, torch::Tensor actions) {
    int N = (int)h.size(0);
    if (N == 0) return;
    head_kernel<<<cdiv(N * 32, 256), 256, 0, GSTR>>>(
        h.data_ptr<float>(), w_a.data_ptr<float>(), b_a.data_ptr<float>(),
        w_v.data_ptr<float>(), b_v.data_ptr<float>(), logits.data_ptr<float>(),
        value.data_ptr<float>(), reinterpret_cast<long long*>(actions.data_ptr<int64_t>()), N);
}

void env_move_launch(torch::Tensor agent, torch::Tensor food, torch::Tensor actions,
                     torch::Tensor agent_out, torch::Tensor hit, torch::Tensor flag) {
    int N = (int)agent.size(0);
    if (N == 0) return;
    env_move_kernel<<<cdiv(N, 256), 256, 0, GSTR>>>(
        agent.data_ptr<float>(), food.data_ptr<float>(),
        reinterpret_cast<long long*>(actions.data_ptr<int64_t>()), agent_out.data_ptr<float>(),
        hit.data_ptr<unsigned char>(), flag.data_ptr<int>(), N);
}

void env_rng_launch(torch::Tensor hit, torch::Tensor food, torch::Tensor rng,
                    torch::Tensor flag, torch::Tensor other_flag, torch::Tensor reward,
                    torch::Tensor agent, torch::Tensor obs) {
    int N = (int)rng.size(0);
    if (N == 0) return;
    env_rng_kernel<<<cdiv(N, 256), 256, 0, GSTR>>>(
        hit.data_ptr<unsigned char>(), food.data_ptr<float>(),
        reinterpret_cast<long long*>(rng.data_ptr<int64_t>()), flag.data_ptr<int>(), other_flag.data_ptr<int>(),
        reward.data_ptr<float>(), agent.data_ptr<float>(), obs.data_ptr<float>(), N);
}

void obs_launch(torch::Tensor agent, torch::Tensor food, torch::Tensor obs) {
    int N = (int)agent.size(0);
    if (N == 0) return;
    obs_kernel<<<cdiv(N, 256), 256, 0, GSTR>>>(agent.data_ptr<float>(), food.data_ptr<float>(),
                                      obs.data_ptr<float>(), N);
}

void zero_launch(torch::Tensor t) {
    if (t.numel() == 0) return;
    cudaMemsetAsync(t.data_ptr(), 0, t.numel() * t.element_size(), GSTR.stream());
}

void copy_launch(torch::Tensor dst, torch::Tensor src) {
    if (dst.numel() == 0) return;
    cudaMemcpyAsync(dst.data_ptr(), src.data_ptr(), dst.numel() * dst.element_size(),
                    cudaMemcpyDeviceToDevice, GSTR.stream());
}

"""

_CPP_SRC = r"""
#include <torch/extension.h>
void enc_launch(torch::Tensor obs, torch::Tensor w_enc, torch::Tensor b_enc, torch::Tensor h0);
void layer_launch(torch::Tensor h_in, torch::Tensor st, torch::Tensor Wt,
                  torch::Tensor h_out, int layer);
void head_launch(torch::Tensor h, torch::Tensor w_a, torch::Tensor b_a, torch::Tensor w_v,
                 torch::Tensor b_v, torch::Tensor logits, torch::Tensor value,
                 torch::Tensor actions);
void env_move_launch(torch::Tensor agent, torch::Tensor food, torch::Tensor actions,
                     torch::Tensor agent_out, torch::Tensor hit, torch::Tensor flag);
void env_rng_launch(torch::Tensor hit, torch::Tensor food, torch::Tensor rng,
                    torch::Tensor flag, torch::Tensor other_flag, torch::Tensor reward,
                    torch::Tensor agent, torch::Tensor obs);
void obs_launch(torch::Tensor agent, torch::Tensor food, torch::Tensor obs);
void zero_launch(torch::Tensor t);
void copy_launch(torch::Tensor dst, torch::Tensor src);
"""

def _build_dir():
    d = os.path.join(os.path.expanduser("~"), ".cache", "torch_extensions",
                     "grid_mingru_fused")
    os.makedirs(d, exist_ok=True)
    return d


def _load():
    return load_inline(
        name="grid_mingru_fused",
        cpp_sources=[_CPP_SRC],
        cuda_sources=[_CUDA_SRC],
        build_directory=_build_dir(),
        functions=[
            "enc_launch", "layer_launch", "head_launch",
            "env_move_launch", "env_rng_launch", "obs_launch",
            "zero_launch", "copy_launch",
        ],
        extra_cuda_cflags=["-O3", "-arch=sm_120", "--use_fast_math"],
        verbose=False,
    )


_mod = _load()


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)

    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)


def _wkey(model):
    return tuple((p._version, p.data_ptr()) for p in model.parameters())


def _params(model):
    return (model.w_enc, model.b_enc, model.w_gru, model.w_a, model.b_a,
            model.w_v, model.b_v)


def _permute_gru(w_gru):
    """(L,768,256) -> (L,256,768) with column c = 3*i+g mapping to row g*256+i."""
    return w_gru.detach().reshape(GRU_LAYERS, 3, HIDDEN, HIDDEN) \
                .permute(0, 2, 1, 3).reshape(GRU_LAYERS, GRU_OUT, HIDDEN) \
                .transpose(1, 2).contiguous()


def _get_perm(model):
    w_gru = model.w_gru
    key = (w_gru._version, w_gru.data_ptr())
    cache = getattr(model, "_perm_cache", None)
    if cache is not None and cache[0] == key:
        return cache[1]
    perm = _permute_gru(w_gru)
    model._perm_cache = (key, perm)
    return perm


# Row tile of the fused layer kernel; callers must pad M to a multiple of it.
_BM = 64


def _pad_rows(t, bm=None):
    bm = _BM if bm is None else bm
    n = t.shape[0]
    m = ((n + bm - 1) // bm) * bm
    if m == n:
        return t
    return torch.cat([t, torch.zeros(m - n, *t.shape[1:], device=t.device, dtype=t.dtype)], 0)


def policy_forward(model, obs: torch.Tensor, state: torch.Tensor):
    """obs (N,4), state (N,3,256) -> logits (N,4), new_state (N,3,256), value (N,)"""
    obs = obs.contiguous().float()
    state = state.contiguous().float()
    N = obs.shape[0]
    dev = obs.device
    w_enc, b_enc, w_gru, w_a, b_a, w_v, b_v = _params(model)
    perm = _get_perm(model)

    obs_p = _pad_rows(obs)
    M = obs_p.shape[0]
    h_a = torch.zeros(M, HIDDEN, device=dev, dtype=torch.float32)
    h_b = torch.empty_like(h_a)
    st = torch.zeros(M, GRU_LAYERS, HIDDEN, device=dev, dtype=torch.float32)
    st[:N].copy_(state.reshape(N, GRU_LAYERS, HIDDEN))
    logits = torch.empty(N, NUM_ACTIONS, device=dev, dtype=torch.float32)
    value = torch.empty(N, device=dev, dtype=torch.float32)
    actions = torch.empty(N, device=dev, dtype=torch.int64)

    _mod.enc_launch(obs_p, w_enc, b_enc, h_a)
    for layer in range(GRU_LAYERS):
        _mod.layer_launch(h_a, st, perm[layer], h_b, layer)
        h_a, h_b = h_b, h_a
    _mod.head_launch(h_a, w_a, b_a, w_v, b_v, logits, value, actions)
    return logits, st[:N].clone(), value


def env_step(agent, food, actions, rng_state):
    """Same semantics as reference.env_step (global any-hit advances all RNGs)."""
    agent = agent.contiguous().float()
    food = food.contiguous().float()
    actions = actions.contiguous().to(torch.int64)
    rng_state = rng_state.contiguous().to(torch.int64)
    N = agent.shape[0]
    dev = agent.device
    agent_out = torch.empty_like(agent)
    hit = torch.empty(N, device=dev, dtype=torch.uint8)
    food = food.clone()
    flags = torch.zeros(2, device=dev, dtype=torch.int32)
    _mod.env_move_launch(agent, food, actions, agent_out, hit, flags[0])
    # the rng kernel also refreshes obs from the post-move agent; unused here
    _mod.env_rng_launch(hit, food, rng_state, flags[0], flags[1],
                        torch.zeros(N, device=dev), agent_out,
                        torch.empty(N, OBS_DIM, device=dev))
    return agent_out, food, hit.float(), rng_state


class _Runner:
    """Owns the per-(num_envs,horizon) buffers and the captured CUDA graph."""

    def __init__(self, model, num_envs, horizon):
        self.model = model
        self.num_envs = num_envs
        self.horizon = horizon
        dev = torch.device("cuda:0")
        self.dev = dev
        self.M = ((num_envs + _BM - 1) // _BM) * _BM

        self.h_a = torch.zeros(self.M, HIDDEN, device=dev)
        self.h_b = torch.empty_like(self.h_a)
        self.state = torch.zeros(self.M, GRU_LAYERS, HIDDEN, device=dev)
        self.obs = torch.empty(self.M, OBS_DIM, device=dev)
        self.hit = torch.empty(self.M, device=dev, dtype=torch.uint8)
        self.agent = torch.empty(self.M, 2, device=dev)
        self.agent_out = torch.empty_like(self.agent)
        self.food = torch.empty(self.M, 2, device=dev)
        self.rng = torch.empty(self.M, device=dev, dtype=torch.int64)
        self.rewards = torch.empty(self.M, device=dev)
        self.logits = torch.empty(self.M, NUM_ACTIONS, device=dev)
        self.value = torch.empty(self.M, device=dev)
        self.actions = torch.empty(self.M, device=dev, dtype=torch.int64)
        self.flags = torch.zeros(2, device=dev, dtype=torch.int32)
        # Sources for the captured ``_reset`` copies.  Replaying the graph
        # always reads these buffers, so a new seed only needs to rewrite their
        # contents -- the graph itself stays valid.
        self.agent0 = torch.zeros(self.M, 2, device=dev)
        self.food0 = torch.zeros(self.M, 2, device=dev)
        self.rng0 = torch.zeros(self.M, dtype=torch.int64, device=dev)
        self.seed = None
        self.graph = None
        self.perm_key = None
        self.perm = None

    def _step(self, t):
        m = _mod
        p = self.perm
        m.enc_launch(self.obs, self.model.w_enc, self.model.b_enc, self.h_a)
        m.layer_launch(self.h_a, self.state, p[0], self.h_b, 0)
        m.layer_launch(self.h_b, self.state, p[1], self.h_a, 1)
        m.layer_launch(self.h_a, self.state, p[2], self.h_b, 2)
        m.head_launch(self.h_b, self.model.w_a, self.model.b_a,
                      self.model.w_v, self.model.b_v,
                      self.logits, self.value, self.actions)
        m.env_move_launch(self.agent, self.food, self.actions, self.agent_out,
                          self.hit, self.flags[t & 1])
        m.env_rng_launch(self.hit, self.food, self.rng, self.flags[t & 1],
                         self.flags[(t + 1) & 1], self.rewards,
                         self.agent_out, self.obs)
        self.agent, self.agent_out = self.agent_out, self.agent

    def _reset(self, agent0, food0, rng0):
        m = _mod
        m.copy_launch(self.agent, agent0)
        m.copy_launch(self.food, food0)
        m.copy_launch(self.rng, rng0)
        m.zero_launch(self.state)
        m.zero_launch(self.rewards)
        m.zero_launch(self.flags)
        m.obs_launch(self.agent, self.food, self.obs)

    def _fill_init(self, seed):
        """Rewrite the captured reset sources for ``seed`` (mirrors reference.run)."""
        n = self.num_envs
        g = torch.Generator(device="cpu")
        g.manual_seed(seed)
        self.agent0[:n].copy_(torch.randint(0, BOARD, (n, 2), generator=g).float())
        self.food0[:n].copy_(torch.randint(0, BOARD, (n, 2), generator=g).float())
        self.rng0[:n].copy_(torch.arange(n, dtype=torch.int64) + seed * 10007)
        self.seed = seed

    def capture(self, seed):
        self._fill_init(seed)
        self.perm = _get_perm(self.model)
        self.perm_key = _wkey(self.model)

        s = torch.cuda.Stream()
        s.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(s):
            for _ in range(2):
                self._reset(self.agent0, self.food0, self.rng0)
                for t in range(self.horizon):
                    self._step(t)
        torch.cuda.current_stream().wait_stream(s)
        torch.cuda.synchronize()
        gph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(gph):
            self._reset(self.agent0, self.food0, self.rng0)
            for t in range(self.horizon):
                self._step(t)
        self.graph = gph
        torch.cuda.synchronize()

    def run(self, seed):
        if self.graph is None or self.perm_key != _wkey(self.model):
            self.capture(seed)
        elif seed != self.seed:
            self._fill_init(seed)
        self.graph.replay()
        return self._collect()

    def _collect(self):
        n = self.num_envs
        return {
            "rewards": self.rewards[:n].clone(),
            "positions": self.agent[:n].round().long(),
            "last_logits": self.logits[:n].clone(),
            "state": self.state[:n].clone(),
        }


_RUNNERS = {}


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

    key = (num_envs, horizon, _wkey(model))
    runner = _RUNNERS.get(key)
    if runner is None:
        runner = _Runner(model, num_envs, horizon)
        _RUNNERS[key] = runner
    return runner.run(seed)

20260910_202109_deepseek-claude_deepseek-flash_04_grid_mingru_sps