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