KernelBench cuda · RTX PRO 6000
Grid + MinGRU SPS Claude Fable 5.1
manually audited: clean
Claude Fable 5.1 through OpenRouter (or-fable harness), max effort, unlimited budget, RTX PRO 6000, 2h30m to a voluntary stop, 17 builds from a first PASS at 0.2205 to 0.7057. Cluster-cooperative persistent kernel with the whole rollout resident: layer 1 collapsed algebraically into a K=4 product (W1 @ W_enc composed in fp64 on the host), layers 2 and 3 on mma.sync m16n8k16 with split-fp16 operands (hi + lo, weights pre-scaled by 2^11, three MMAs per tile) for fp32-level logits (max abs error 3e-8 against the reference), layer outputs exchanged between CTAs through L2 with cp.async and hardware cluster barriers, the environment step fused with a deferred, exactly replayed batch-wide LCG advance (per-step any-hit flags plus completion counters, no grid barrier), and initial positions from a bit-exact on-GPU MT19937 matching torch's CPU randint. Fresh reward, position and logit tensors every call; the only caches are a weight packing keyed on (id, data_ptr, _version) and an (N,H) workspace whose state buffer is returned as a view. Contended in-run grade 0.7056; isolated sequential regrade on the same box 2026-09-04 (clocks reset first): 0.7093 (+0.5%), 100.9M to 111.5M steps/s across the four shapes. Probe (probe.log): on all four deck shapes with seeds 42 and 123 every environment's position trajectory equals the reference (envs_diverged 0), rewards match, logits within 3.4e-8, state cos 1.00000, and an in-place weight overwrite tracks the new reference (ALL_POSITIONS_EQUAL). One disclosed integrity episode: the agent ran nvidia-smi -lgc 2400,2400 at 12:57Z for A/B timing hygiene (locked and unlocked numbers agreed within 1%), reset with -rgc at 13:19Z, and the graded check and benchmark ran at 13:41Z on default boost clocks; the regrade reset clocks again before measuring. Templates byte-identical, no network, no foreign reads, no key strings. OpenRouter cost $53.40 (599k output tokens, 36.0M cache reads). Second on the problem behind or-opus 1.961.
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 megakernel solution: vectorized 11x11 grid foraging + 3x MinGRU(h=256) greedy rollout.
Design (kernel source embedded below as _KERNEL_SRC / _HOST_SRC):
* one persistent launch; thread-block clusters of 4 CTAs own 96-env row blocks for the whole horizon,
each CTA computing 64 hidden units (192 gate columns) of every layer for its rows
* layer 1 is collapsed algebraically: gates_1 = W_1 (W_enc obs + b_enc) = (W_1 W_enc) obs + W_1 b_enc,
a K=4 product with weights composed in fp64; only layers 2 and 3 run the 256-wide GEMM
* GEMMs on tensor cores (mma.sync m16n8k16) with split precision x = hi + lo (fp16 each, weights
pre-scaled by 2^11, 3 MMAs per tile) -> fp32-level gate accuracy (logits within ~3e-8 of the reference)
* activations kept in shared memory in MMA-fragment-interleaved order; layer outputs exchanged between
the 4 CTAs through L2 with cp.async and the hardware cluster barrier; the encoder is recomputed locally
* env step (argmax action, clamped move, food hit, LCG respawn) fused into the kernel; the reference's
batch-wide "any env hit food" RNG rule is honoured exactly with deferred, on-demand RNG advancement
* initial agent/food positions are generated on the GPU with a bit-exact MT19937 (torch CPU randint)
"""
from __future__ import annotations
import hashlib
import os
from pathlib import Path
import torch
import torch.nn as nn
BOARD = 11
OBS_DIM = 4
HIDDEN = 256
GRU_LAYERS = 3
NUM_ACTIONS = 4
GRU_OUT = 3 * HIDDEN
_HERE = Path(__file__).resolve().parent
# ----------------------------------------------------------------------------- build extension
try: # ninja may only be available as a python package
import ninja # type: ignore
os.environ["PATH"] = ninja.BIN_DIR + os.pathsep + os.environ.get("PATH", "")
except Exception: # pragma: no cover
pass
_CPP_DECLS = r"""
#include <torch/extension.h>
#include <vector>
int gm_max_clusters();
std::vector<torch::Tensor> gm_rollout(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc, torch::Tensor benc,
torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv, int64_t N, int64_t H, int64_t seed);
std::vector<torch::Tensor> gm_policy_forward(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc,
torch::Tensor benc, torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv,
torch::Tensor obs, torch::Tensor state);
std::vector<torch::Tensor> gm_env_step(torch::Tensor agent, torch::Tensor food, torch::Tensor actions,
torch::Tensor rng);
"""
_KERNEL_SRC = r"""
// Grid-foraging env + 3x MinGRU(h=256) policy rollout megakernel for sm_120 (v6).
//
// Structure
// * persistent thread-block clusters of S=4 CTAs; each cluster owns row blocks of R=96 envs
// * CTA rank c computes hidden units [64c, 64c+64) of every layer (192 gate columns) for its 96 rows
// * GEMMs run on tensor cores (mma.sync m16n8k16 fp16 -> fp32) with split precision:
// x = hi + lo (hi, lo fp16), weights pre-scaled by 2^11; 3 MMAs per tile (hh, hl, lh)
// which gives ~fp32-level accuracy for the gate pre-activations.
// * layer 1 is collapsed algebraically: gates_1 = W_1 (W_enc obs + b_enc) = (W_1 W_enc) obs + W_1 b_enc,
// a K=4 product per gate column, so only layers 2 and 3 run the 256-wide GEMM.
// * activations live in shared memory in MMA-fragment-interleaved order (rows r and r+8 paired), so a
// single 16-byte load yields a complete A fragment.
// * layer outputs are exchanged between the 4 CTAs through L2 (24 KB slices), synced with the
// hardware cluster barrier; remote chunks stream into a 3-slot smem ring via cp.async.
// * env step is computed redundantly by every CTA of the cluster (deterministic integer math).
// * the reference advances the food RNG only if ANY env in the whole batch hit food that step.
// Nothing observable depends on a row's RNG until that row hits food, so RNG advancement is
// deferred: each cluster publishes per-step "any hit" flags + a completion counter, and a block
// that hits at step t first waits until every cluster has finished step t-1, then replays the
// global flag history for its rows. Exact semantics, single launch, no grid barrier.
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cooperative_groups.h>
#include <cstdint>
namespace cg = cooperative_groups;
namespace gm {
constexpr int BOARD = 11;
constexpr int HID = 256;
constexpr int R = 96; // rows (envs) per block
constexpr int S = 4; // cluster size
constexpr int UNITS = HID / S; // hidden units per CTA (64)
constexpr int NT = 256; // threads per CTA
constexpr int MT = R / 16; // m-tiles (6)
constexpr int REC_BYTES = 256; // record = rows (r, r+8) of one m-tile x one k32 chunk, hi+lo fp16
constexpr int CHUNK_BYTES = (R / 2) * REC_BYTES; // 48 records = 12288
constexpr int CHUNK_U4 = CHUNK_BYTES / 16; // 768
constexpr int NSLOT = 5; // 2 own chunks + 3-slot ring for remote chunks
constexpr int SMEM_SLOTS = NSLOT * CHUNK_BYTES; // 61440
constexpr int SMEM_OBS = SMEM_SLOTS; // R*4 floats
constexpr int SMEM_WENC = SMEM_OBS + R * 4 * 4; // my 64 units x 4 floats
constexpr int SMEM_BENC = SMEM_WENC + UNITS * 4 * 4; // 64 floats
constexpr int SMEM_W1P = SMEM_BENC + UNITS * 4; // composed layer-1 weights [3 gates][64 units][4] (x2048)
constexpr int SMEM_B1P = SMEM_W1P + 3 * UNITS * 4 * 4; // composed layer-1 bias [3][64] (x2048)
constexpr int SMEM_BA = SMEM_B1P + 3 * UNITS * 4; // logits-head B fragments: 2 chunks x 2 hilo x 32 lanes x 16 B
constexpr int SMEM_TOTAL = SMEM_BA + 2 * 2 * 32 * 16; // 69376
constexpr float INV_SCALE = 1.0f / 2048.0f;
enum Mode { MODE_RUN = 0, MODE_FWD = 2 };
struct Params {
const uint4* __restrict__ Bgru; // [3][S][8 warps][8 chunks][2 hilo][3 gates][32 lanes] uint4
const uint4* __restrict__ Ba; // [S][2 chunks][2 hilo][32 lanes] uint4
const float* __restrict__ Wenc; // [256][4]
const float* __restrict__ benc; // [256]
const float* __restrict__ W1p; // [768][4] 2048 * (w_gru[0] @ w_enc)
const float* __restrict__ b1p; // [768] 2048 * (w_gru[0] @ b_enc)
const float* __restrict__ ba; // [4]
const float* __restrict__ bv; // [1]
const float* state_in; // [npad][3][256] (fwd mode input; run mode == state_out)
float* state_out; // [npad][3][256]
uint4* slices; // [nclusters][2 bufs][S][2 chunks] (CHUNK_BYTES each)
float* partials; // [nclusters][S][R][8]
int* env; // [S][npad][8] int32 words: ax ay fx fy rng_lo rng_hi reward rng_step
const int* env_init; // [npad][4] ax, ay, fx, fy
int* flags; // [H] any-hit flag per step (global OR over all envs)
unsigned int* done; // [H] number of clusters that finished step t
float* out_rewards; // [N]
long long* out_pos; // [N][2]
float* out_logits; // [N][4]
float* out_value; // [N]
const float* __restrict__ obs_in; // [N][4] (fwd mode)
int N, npad, nblocks, nclusters, H, mode;
long long seed;
};
// ------------------------------------------------------------------ device helpers
__device__ __forceinline__ uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }
__device__ __forceinline__ void mma16816(float* c, const uint4& a, uint32_t b0, uint32_t 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"(c[0]), "+f"(c[1]), "+f"(c[2]), "+f"(c[3])
: "r"(a.x), "r"(a.y), "r"(a.z), "r"(a.w), "r"(b0), "r"(b1));
}
__device__ __forceinline__ void cp_async16(uint32_t saddr, const void* gptr) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" [REDACTED: IP]"r"(saddr), "l"(gptr));
}
__device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;\n" [REDACTED: IP]); }
template <int N>
__device__ __forceinline__ void cp_async_wait() { asm volatile("cp.async.wait_group %0;\n" [REDACTED: IP]"n"(N)); }
__device__ __forceinline__ void cluster_arrive() { asm volatile("barrier.cluster.arrive.release.aligned;\n" ::: "memory"); }
__device__ __forceinline__ void cluster_wait() { asm volatile("barrier.cluster.wait.acquire.aligned;\n" ::: "memory"); }
__device__ __forceinline__ void cluster_sync_all() { cluster_arrive(); cluster_wait(); }
__device__ __forceinline__ float ex2_approx(float x) {
float y;
asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
__device__ __forceinline__ float rcp_approx(float x) {
float y;
asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
// sigmoid(zg), tanh(zh), sigmoid(zp) with one reciprocal: 3 MUFU.EX2 + 1 MUFU.RCP per element
__device__ __forceinline__ void gates3(float zh, float zg, float zp, float& sg, float& th, float& sp) {
const float eg = ex2_approx(fminf(-1.4426950408889634f * zg, 40.0f));
const float eh = ex2_approx(fminf(-2.8853900817779268f * zh, 40.0f));
const float ep = ex2_approx(fminf(-1.4426950408889634f * zp, 40.0f));
const float a = 1.0f + eg, b = 1.0f + eh, c = 1.0f + ep;
const float ab = a * b;
const float r = rcp_approx(ab * c);
sg = (b * c) * r;
th = (1.0f - eh) * ((a * c) * r);
sp = ab * r;
}
__device__ __forceinline__ uint32_t pack_half2(float lo_val, float hi_val) {
__half2 h = __floats2half2_rn(lo_val, hi_val); // .x = first arg (low 16 bits)
return *reinterpret_cast<uint32_t*>(&h);
}
__device__ __forceinline__ void split_hilo(float x, float& hi, float& lo) {
hi = __half2float(__float2half_rn(x));
lo = x - hi;
}
__device__ __forceinline__ unsigned long long lcg(unsigned long long r) {
return (r * 6364136223846793005ULL + 1ULL) & 0x7FFFFFFFFFFFFFFFULL;
}
// Byte offset inside a chunk slot of the 4-byte word holding units (k, k+1) of row r (k even, 0..30),
// hilo 0 = fp16 hi, 1 = fp16 lo. Record = (m-tile, g8); 4 blocks of 64 B: [hi t0][hi t1][lo t0][lo t1],
// block index XOR-swizzled by the record parity; inside a block: quad q at 16q, word = 2*(j&1) + half.
__device__ __forceinline__ int a_word_off(int r, int k, int hilo) {
const int m = r >> 4, g8 = r & 7, half = (r >> 3) & 1;
const int q = (k & 7) >> 1, j = k >> 3;
const int blk = ((j >> 1) + 2 * hilo) ^ (g8 & 1);
return (m * 8 + g8) * REC_BYTES + blk * 64 + 16 * q + 4 * (2 * (j & 1) + half);
}
// ------------------------------------------------------------------ the megakernel
__global__ void __launch_bounds__(NT, 1) rollout_kernel(const Params p) {
extern __shared__ __align__(128) unsigned char smem[];
float* obs_s = reinterpret_cast<float*>(smem + SMEM_OBS);
float* wenc_s = reinterpret_cast<float*>(smem + SMEM_WENC);
float* benc_s = reinterpret_cast<float*>(smem + SMEM_BENC);
float* w1p_s = reinterpret_cast<float*>(smem + SMEM_W1P);
float* b1p_s = reinterpret_cast<float*>(smem + SMEM_B1P);
uint4* ba_s = reinterpret_cast<uint4*>(smem + SMEM_BA);
cg::cluster_group cluster = cg::this_cluster();
const int rank = (int)cluster.block_rank();
const int cid = blockIdx.x / S;
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const int q = lane & 3;
const int g8 = lane >> 2; // row within m-tile (0..7)
const int mode = p.mode;
// encoder weights of my 64 units (highway input h0) and the composed layer-1 gate weights of my columns
for (int i = tid; i < UNITS * 4; i += NT) wenc_s[i] = p.Wenc[(rank * UNITS) * 4 + i];
if (tid < UNITS) benc_s[tid] = p.benc[rank * UNITS + tid];
for (int i = tid; i < 3 * UNITS * 4; i += NT) {
const int g = i / (UNITS * 4), rem = i % (UNITS * 4);
w1p_s[i] = p.W1p[(g * HID + rank * UNITS) * 4 + rem];
}
for (int i = tid; i < 3 * UNITS; i += NT) b1p_s[i] = p.b1p[(i / UNITS) * HID + rank * UNITS + (i % UNITS)];
if (tid < 128) ba_s[tid] = p.Ba[(size_t)rank * 128 + tid];
const bool single_block = (p.nblocks <= p.nclusters);
// per-thread constants
const int par = g8 & 1;
const int unit_chunk = warp >> 2; // which own chunk holds my 8 units
const int jj = warp & 3; // unit-group (word index) inside the chunk
// A-fragment load offsets for my lane inside a record: block b' = b ^ par
const int off_h0 = ((0 ^ par) * 64) + 16 * q; // hi, k16 tile 0
const int off_h1 = ((1 ^ par) * 64) + 16 * q; // hi, k16 tile 1
const int off_l0 = ((2 ^ par) * 64) + 16 * q; // lo, tile 0
const int off_l1 = ((3 ^ par) * 64) + 16 * q; // lo, tile 1
// epilogue word offsets (inside the record of my rows) for my unit pair (jj, q)
const int ep_hi = (((jj >> 1) + 0) ^ par) * 64 + 16 * q + 4 * (2 * (jj & 1));
const int ep_lo = (((jj >> 1) + 2) ^ par) * 64 + 16 * q + 4 * (2 * (jj & 1));
unsigned char* own_slot = smem + unit_chunk * CHUNK_BYTES;
auto kchunk = [&](int c) -> int {
if (c < 2) return 2 * rank + c;
const int j = c - 2;
return 2 * ((rank + 1 + (j >> 1)) & (S - 1)) + (j & 1);
};
uint4* my_slice_g0 = p.slices + ((size_t)((cid * 2 + 0) * S + rank)) * (2 * CHUNK_U4); // written after layer 0
uint4* my_slice_g1 = p.slices + ((size_t)((cid * 2 + 1) * S + rank)) * (2 * CHUNK_U4); // written after layer 1
const int nsteps = (mode == MODE_FWD) ? 1 : p.H;
const uint32_t slots_saddr = smem_u32(smem);
uint4 B0[6], B1[6]; // B fragments, alternating buffers: [hilo*3 + g]
const uint4* Bbase = p.Bgru + ((size_t)((1 * S + rank) * 8 + warp)) * (8 * 192) + lane; // layer index 1
#pragma unroll
for (int i = 0; i < 6; ++i) B0[i] = __ldg(Bbase + kchunk(0) * 192 + i * 32);
int ax = 0, ay = 0, fx = -1, fy = -1, rng_step = -1, rew = 0;
unsigned long long rng = 0;
for (int t = 0; t < nsteps; ++t) {
for (int blk = cid; blk < p.nblocks; blk += p.nclusters) {
const int row_base = blk * R;
// ---------------- phase 0: observations
if (tid < R) {
const int grow = row_base + tid;
if (mode == MODE_FWD) {
float4 o = make_float4(0.f, 0.f, 0.f, 0.f);
if (grow < p.N) o = *reinterpret_cast<const float4*>(p.obs_in + (size_t)grow * 4);
*reinterpret_cast<float4*>(obs_s + tid * 4) = o;
} else {
if (t == 0) {
const int4 e = *reinterpret_cast<const int4*>(p.env_init + (size_t)grow * 4);
ax = e.x; ay = e.y; fx = e.z; fy = e.w;
rng = (unsigned long long)((long long)grow + p.seed * 10007LL);
rew = 0; rng_step = -1;
} else if (!single_block) {
const int* e = p.env + ((size_t)rank * p.npad + grow) * 8;
const int4 e0 = *reinterpret_cast<const int4*>(e);
const int4 e1 = *reinterpret_cast<const int4*>(e + 4);
ax = e0.x; ay = e0.y; fx = e0.z; fy = e0.w;
rng = ((unsigned long long)(unsigned)e1.x) | ((unsigned long long)(unsigned)e1.y << 32);
rew = e1.z; rng_step = e1.w;
}
float4 o;
o.x = __fdiv_rn((float)(fx - ax), (float)BOARD);
o.y = __fdiv_rn((float)(fy - ay), (float)BOARD);
o.z = __fdiv_rn((float)ax, (float)(BOARD - 1));
o.w = __fdiv_rn((float)ay, (float)(BOARD - 1));
*reinterpret_cast<float4*>(obs_s + tid * 4) = o;
}
}
__syncthreads();
// ---------------- 3 MinGRU layers
for (int l = 0; l < 3; ++l) {
float acc[MT][3][4];
#pragma unroll
for (int m = 0; m < MT; ++m)
#pragma unroll
for (int g = 0; g < 3; ++g)
#pragma unroll
for (int e = 0; e < 4; ++e) acc[m][g][e] = 0.f;
const uint4* Bl = p.Bgru + ((size_t)((l * S + rank) * 8 + warp)) * (8 * 192) + lane;
const uint4* Bl_next = (l < 2) ? (Bl + (size_t)S * 8 * 8 * 192)
: (p.Bgru + ((size_t)((1 * S + rank) * 8 + warp)) * (8 * 192) + lane);
if (l == 0) {
// ---------------- layer 1 collapsed: acc = 2048 * (W1p obs + b1p) for my fragment elements
const int u0 = warp * 8 + 2 * q; // local unit index of my pair
float wv[3][2][4], bvv[3][2];
#pragma unroll
for (int g = 0; g < 3; ++g) {
#pragma unroll
for (int b = 0; b < 2; ++b) {
const float4 w = *reinterpret_cast<const float4*>(w1p_s + (g * UNITS + u0 + b) * 4);
wv[g][b][0] = w.x; wv[g][b][1] = w.y; wv[g][b][2] = w.z; wv[g][b][3] = w.w;
bvv[g][b] = b1p_s[g * UNITS + u0 + b];
}
}
#pragma unroll
for (int m = 0; m < MT; ++m) {
#pragma unroll
for (int half = 0; half < 2; ++half) {
const int r = 16 * m + g8 + 8 * half;
const float4 o = *reinterpret_cast<const float4*>(obs_s + r * 4);
#pragma unroll
for (int g = 0; g < 3; ++g) {
#pragma unroll
for (int b = 0; b < 2; ++b) {
float v = bvv[g][b];
v = fmaf(wv[g][b][0], o.x, v);
v = fmaf(wv[g][b][1], o.y, v);
v = fmaf(wv[g][b][2], o.z, v);
v = fmaf(wv[g][b][3], o.w, v);
acc[m][g][2 * half + b] = v;
}
}
}
}
} else {
// one chunk of the k-loop; Bcur holds this chunk's B fragments, Bnext receives the next chunk's
#define GM_CHUNK(c, Bcur, Bnext) \
{ \
if ((c) >= 2) { \
if ((c) <= 6) cp_async_wait<1>(); else cp_async_wait<0>(); \
} \
__syncthreads(); \
{ \
if ((c) < 6) { \
const int j = (c); \
const int src_rank = (rank + 1 + (j >> 1)) & (S - 1); \
const uint4* src = p.slices + ((size_t)((cid * 2 + (l - 1)) * S + src_rank)) * (2 * CHUNK_U4) + \
(j & 1) * CHUNK_U4; \
const uint32_t dst = slots_saddr + (2 + (j % 3)) * CHUNK_BYTES; \
_Pragma("unroll") \
for (int i = 0; i < 3; ++i) cp_async16(dst + (tid + i * NT) * 16, src + tid + i * NT); \
cp_async_commit(); \
} \
} \
{ \
const uint4* nb = ((c) < 7) ? (Bl + kchunk((c) + 1) * 192) : (Bl_next + kchunk(0) * 192); \
_Pragma("unroll") \
for (int i = 0; i < 6; ++i) Bnext[i] = __ldg(nb + i * 32); \
} \
const unsigned char* slot = smem + (((c) < 2) ? (c) : (2 + (((c) - 2) % 3))) * CHUNK_BYTES; \
const unsigned char* rec0 = slot + g8 * REC_BYTES; \
uint4 ah0 = *reinterpret_cast<const uint4*>(rec0 + off_h0); \
uint4 ah1 = *reinterpret_cast<const uint4*>(rec0 + off_h1); \
uint4 al0 = *reinterpret_cast<const uint4*>(rec0 + off_l0); \
uint4 al1 = *reinterpret_cast<const uint4*>(rec0 + off_l1); \
_Pragma("unroll") \
for (int m = 0; m < MT; ++m) { \
uint4 nh0, nh1, nl0, nl1; \
if (m + 1 < MT) { \
const unsigned char* recn = rec0 + (m + 1) * 8 * REC_BYTES; \
nh0 = *reinterpret_cast<const uint4*>(recn + off_h0); \
nh1 = *reinterpret_cast<const uint4*>(recn + off_h1); \
nl0 = *reinterpret_cast<const uint4*>(recn + off_l0); \
nl1 = *reinterpret_cast<const uint4*>(recn + off_l1); \
} \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], ah0, Bcur[g].x, Bcur[g].y); \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], ah0, Bcur[3 + g].x, Bcur[3 + g].y); \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], al0, Bcur[g].x, Bcur[g].y); \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], ah1, Bcur[g].z, Bcur[g].w); \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], ah1, Bcur[3 + g].z, Bcur[3 + g].w); \
_Pragma("unroll") \
for (int g = 0; g < 3; ++g) mma16816(acc[m][g], al1, Bcur[g].z, Bcur[g].w); \
if (m + 1 < MT) { ah0 = nh0; ah1 = nh1; al0 = nl0; al1 = nl1; } \
} \
}
#pragma unroll 1
for (int c = 0; c < 8; c += 2) {
GM_CHUNK(c, B0, B1)
GM_CHUNK(c + 1, B1, B0)
}
#undef GM_CHUNK
}
__syncthreads(); // everyone done reading own slots before the epilogue overwrites them
// ---------------- epilogue: gates -> state update, highway; write h_new into own slot
{
const int ubase = rank * UNITS + warp * 8 + 2 * q; // global unit index of my pair
float2 sv[MT][2];
#pragma unroll
for (int m = 0; m < MT; ++m) {
#pragma unroll
for (int half = 0; half < 2; ++half) {
const int grow = row_base + 16 * m + g8 + 8 * half;
const size_t sidx = ((size_t)grow * 3 + l) * HID + ubase;
sv[m][half] = make_float2(0.f, 0.f);
if ((mode != MODE_FWD && t > 0) || (mode == MODE_FWD && grow < p.N)) {
sv[m][half] = *reinterpret_cast<const float2*>(p.state_in + sidx);
}
}
}
#pragma unroll
for (int m = 0; m < MT; ++m) {
#pragma unroll
for (int half = 0; half < 2; ++half) {
const int r = 16 * m + g8 + 8 * half;
const int grow = row_base + r;
const bool valid = grow < p.N;
unsigned char* rec = own_slot + (m * 8 + g8) * REC_BYTES;
uint32_t* hi_p = reinterpret_cast<uint32_t*>(rec + ep_hi + 4 * half);
uint32_t* lo_p = reinterpret_cast<uint32_t*>(rec + ep_lo + 4 * half);
float hold0, hold1;
if (l == 0) {
const float4 o = *reinterpret_cast<const float4*>(obs_s + r * 4);
const int u0 = warp * 8 + 2 * q;
const float4 w0 = *reinterpret_cast<const float4*>(wenc_s + u0 * 4);
const float4 w1 = *reinterpret_cast<const float4*>(wenc_s + (u0 + 1) * 4);
hold0 = benc_s[u0] + w0.x * o.x + w0.y * o.y + w0.z * o.z + w0.w * o.w;
hold1 = benc_s[u0 + 1] + w1.x * o.x + w1.y * o.y + w1.z * o.z + w1.w * o.w;
} else {
const float2 hhf = __half22float2(*reinterpret_cast<const __half2*>(hi_p));
const float2 hlf = __half22float2(*reinterpret_cast<const __half2*>(lo_p));
hold0 = hhf.x + hlf.x;
hold1 = hhf.y + hlf.y;
}
const float s0 = sv[m][half].x, s1 = sv[m][half].y;
const size_t sidx = ((size_t)grow * 3 + l) * HID + ubase;
const float zh0 = acc[m][0][2 * half + 0] * INV_SCALE, zh1 = acc[m][0][2 * half + 1] * INV_SCALE;
const float zg0 = acc[m][1][2 * half + 0] * INV_SCALE, zg1 = acc[m][1][2 * half + 1] * INV_SCALE;
const float zp0 = acc[m][2][2 * half + 0] * INV_SCALE, zp1 = acc[m][2][2 * half + 1] * INV_SCALE;
float sg0, th0, p0, sg1, th1, p1;
gates3(zh0, zg0, zp0, sg0, th0, p0);
gates3(zh1, zg1, zp1, sg1, th1, p1);
const float out0 = s0 + sg0 * (th0 - s0);
const float out1 = s1 + sg1 * (th1 - s1);
const float hn0 = p0 * out0 + (1.0f - p0) * hold0;
const float hn1 = p1 * out1 + (1.0f - p1) * hold1;
if (mode != MODE_FWD || valid) {
*reinterpret_cast<float2*>(p.state_out + sidx) = make_float2(out0, out1);
}
float a_h, a_l, b_h, b_l;
split_hilo(hn0, a_h, a_l);
split_hilo(hn1, b_h, b_l);
*hi_p = pack_half2(a_h, b_h);
*lo_p = pack_half2(a_l, b_l);
}
}
}
__syncthreads();
if (l < 2) {
const uint4* src = reinterpret_cast<const uint4*>(smem);
uint4* dstg = (l == 0) ? my_slice_g0 : my_slice_g1;
#pragma unroll
for (int i = 0; i < 6; ++i) dstg[tid + i * NT] = src[tid + i * NT];
cluster_sync_all();
}
}
// ---------------- partial logits from my 64 units of h3: warps 0..5, one m-tile each
if (warp < MT) {
float lacc[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int ch = 0; ch < 2; ++ch) {
const uint4 bh = ba_s[(ch * 2 + 0) * 32 + lane];
const uint4 bl = ba_s[(ch * 2 + 1) * 32 + lane];
const unsigned char* rec = smem + ch * CHUNK_BYTES + (warp * 8 + g8) * REC_BYTES;
const uint4 ah0 = *reinterpret_cast<const uint4*>(rec + off_h0);
const uint4 ah1 = *reinterpret_cast<const uint4*>(rec + off_h1);
const uint4 al0 = *reinterpret_cast<const uint4*>(rec + off_l0);
const uint4 al1 = *reinterpret_cast<const uint4*>(rec + off_l1);
mma16816(lacc, ah0, bh.x, bh.y);
mma16816(lacc, ah0, bl.x, bl.y);
mma16816(lacc, al0, bh.x, bh.y);
mma16816(lacc, ah1, bh.z, bh.w);
mma16816(lacc, ah1, bl.z, bl.w);
mma16816(lacc, al1, bh.z, bh.w);
}
float* pp = p.partials + ((size_t)(cid * S + rank)) * (R * 8);
const int r0 = warp * 16 + g8;
*reinterpret_cast<float2*>(pp + r0 * 8 + 2 * q) = make_float2(lacc[0], lacc[1]);
*reinterpret_cast<float2*>(pp + (r0 + 8) * 8 + 2 * q) = make_float2(lacc[2], lacc[3]);
}
cluster_sync_all();
// ---------------- logits, action, env step (redundant on every CTA of the cluster)
int hit = 0;
if (tid < R) {
const int grow = row_base + tid;
const bool valid = grow < p.N;
const float* pp = p.partials + ((size_t)(cid * S)) * (R * 8) + tid * 8;
float lg[5];
#pragma unroll
for (int a = 0; a < 5; ++a) {
float v = 0.f;
#pragma unroll
for (int s = 0; s < S; ++s) v += pp[s * (R * 8) + a];
lg[a] = v * INV_SCALE + ((a < 4) ? p.ba[a] : p.bv[0]);
}
if (mode == MODE_FWD) {
if (rank == 0 && valid) {
*reinterpret_cast<float4*>(p.out_logits + (size_t)grow * 4) = make_float4(lg[0], lg[1], lg[2], lg[3]);
p.out_value[grow] = lg[4];
}
} else {
int act = 0;
float best = lg[0];
if (lg[1] > best) { best = lg[1]; act = 1; }
if (lg[2] > best) { best = lg[2]; act = 2; }
if (lg[3] > best) { best = lg[3]; act = 3; }
const int dx = (act == 2) ? -1 : ((act == 3) ? 1 : 0);
const int dy = (act == 0) ? -1 : ((act == 1) ? 1 : 0);
ax = min(max(ax + dx, 0), BOARD - 1);
ay = min(max(ay + dy, 0), BOARD - 1);
hit = (ax == fx && ay == fy && valid) ? 1 : 0;
rew += hit;
if (t == p.H - 1 && rank == 0 && valid) {
p.out_pos[(size_t)grow * 2 + 0] = ax;
p.out_pos[(size_t)grow * 2 + 1] = ay;
p.out_rewards[grow] = (float)rew;
*reinterpret_cast<float4*>(p.out_logits + (size_t)grow * 4) = make_float4(lg[0], lg[1], lg[2], lg[3]);
}
}
}
if (mode != MODE_FWD) {
const int anyhit = __syncthreads_or(hit);
if (anyhit) {
if (tid == 0 && t > 0) {
while (*((volatile unsigned int*)(p.done + (t - 1))) < (unsigned)p.nclusters) { __nanosleep(64); }
__threadfence();
}
__syncthreads();
if (tid < R) {
for (int tt = rng_step + 1; tt < t; ++tt) {
if (__ldcg(p.flags + tt) != 0) { rng = lcg(lcg(rng)); }
}
rng = lcg(rng);
const int nfx = (int)(rng % (unsigned long long)BOARD);
rng = lcg(rng);
const int nfy = (int)(rng % (unsigned long long)BOARD);
if (hit) { fx = nfx; fy = nfy; }
rng_step = t;
}
if (tid == NT - 1 && rank == 0) p.flags[t] = 1;
}
if (tid < R && !single_block) {
const int grow = row_base + tid;
int* e = p.env + ((size_t)rank * p.npad + grow) * 8;
*reinterpret_cast<int4*>(e) = make_int4(ax, ay, fx, fy);
*reinterpret_cast<int4*>(e + 4) = make_int4((int)(unsigned)(rng & 0xFFFFFFFFULL), (int)(unsigned)(rng >> 32), rew, rng_step);
}
}
} // blocks
if (mode != MODE_FWD && rank == 0 && tid == NT - 1) {
__threadfence();
atomicAdd(p.done + t, 1u);
}
} // steps
}
// ------------------------------------------------------------------ MT19937 init (matches torch CPU randint)
// torch.randint(0, 11, ...) on a CPU generator seeded with `seed` = mt19937 32-bit outputs mod 11, in order:
// agent (N,2) row-major first, then food (N,2). One block of 640 threads; the twist of one generation is
// done in 3 dependency phases (i<227 | 227<=i<454 | i>=454) on a double-buffered state.
__global__ void __launch_bounds__(640) mt_init_kernel(int* env_init, int N, int npad, unsigned int seed) {
__shared__ unsigned int st[2][624];
const int tid = threadIdx.x;
if (tid == 0) {
unsigned x = seed;
st[0][0] = x;
for (int i = 1; i < 624; ++i) { x = 1812433253u * (x ^ (x >> 30)) + (unsigned)i; st[0][i] = x; }
}
__syncthreads();
const int total = 4 * N;
const int ngen = (total + 623) / 624;
int cur = 0;
for (int g = 0; g < ngen; ++g) {
const unsigned* a = st[cur];
unsigned* b = st[cur ^ 1];
if (tid < 227) {
const int i = tid;
const unsigned y = (a[i] & 0x80000000u) | (a[i + 1] & 0x7fffffffu);
b[i] = a[i + 397] ^ (y >> 1) ^ ((y & 1u) ? 0x9908b0dfu : 0u);
}
__syncthreads();
if (tid < 227) {
const int i = 227 + tid;
const unsigned y = (a[i] & 0x80000000u) | (a[i + 1] & 0x7fffffffu);
b[i] = b[i - 227] ^ (y >> 1) ^ ((y & 1u) ? 0x9908b0dfu : 0u);
}
__syncthreads();
if (tid < 170) {
const int i = 454 + tid;
const unsigned nxt = (i == 623) ? b[0] : a[i + 1];
const unsigned y = (a[i] & 0x80000000u) | (nxt & 0x7fffffffu);
b[i] = b[i - 227] ^ (y >> 1) ^ ((y & 1u) ? 0x9908b0dfu : 0u);
}
__syncthreads();
if (tid < 624) {
const int j = g * 624 + tid;
if (j < total) {
unsigned y = b[tid];
y ^= (y >> 11);
y ^= (y << 7) & 0x9d2c5680u;
y ^= (y << 15) & 0xefc60000u;
y ^= (y >> 18);
const int v = (int)(y % 11u);
if (j < 2 * N) env_init[(j >> 1) * 4 + (j & 1)] = v;
else { const int jj = j - 2 * N; env_init[(jj >> 1) * 4 + 2 + (jj & 1)] = v; }
}
}
cur ^= 1;
}
for (int r = N + tid; r < npad; r += 640) {
env_init[r * 4 + 0] = 0; env_init[r * 4 + 1] = 0; env_init[r * 4 + 2] = -1; env_init[r * 4 + 3] = -1;
}
}
// ------------------------------------------------------------------ standalone env_step kernels
__global__ void env_move_kernel(const float* agent, const float* food, const long long* actions,
float* agent_out, float* reward, int* hitbuf, int* anyhit, int N) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= N) return;
int ax = (int)agent[i * 2 + 0], ay = (int)agent[i * 2 + 1];
const int fx = (int)food[i * 2 + 0], fy = (int)food[i * 2 + 1];
const long long act = actions[i];
const int dx = (act == 2) ? -1 : ((act == 3) ? 1 : 0);
const int dy = (act == 0) ? -1 : ((act == 1) ? 1 : 0);
ax = min(max(ax + dx, 0), BOARD - 1);
ay = min(max(ay + dy, 0), BOARD - 1);
agent_out[i * 2 + 0] = (float)ax; agent_out[i * 2 + 1] = (float)ay;
const int hit = (ax == fx && ay == fy) ? 1 : 0;
reward[i] = (float)hit;
hitbuf[i] = hit;
if (hit) atomicOr(anyhit, 1);
}
__global__ void env_respawn_kernel(const float* food, const int* hitbuf, const long long* rng_in,
const int* anyhit, float* food_out, long long* rng_out, int N) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= N) return;
unsigned long long r = (unsigned long long)rng_in[i];
float fx = food[i * 2 + 0], fy = food[i * 2 + 1];
if (*anyhit) {
r = lcg(r);
const int nfx = (int)(r % (unsigned long long)BOARD);
r = lcg(r);
const int nfy = (int)(r % (unsigned long long)BOARD);
if (hitbuf[i]) { fx = (float)nfx; fy = (float)nfy; }
}
food_out[i * 2 + 0] = fx; food_out[i * 2 + 1] = fy;
rng_out[i] = (long long)r;
}
} // namespace gm
"""
_HOST_SRC = r"""
// Host-side launch wrappers (appended after gm_kernel.cu in the load_inline CUDA source).
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <vector>
#include <algorithm>
namespace gm {
static int g_max_clusters = -1;
static void ensure_attrs() {
if (g_max_clusters > 0) return;
cudaFuncSetAttribute(rollout_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_TOTAL);
cudaLaunchConfig_t cfg = {};
cfg.gridDim = dim3(S * 47, 1, 1);
cfg.blockDim = dim3(NT, 1, 1);
cfg.dynamicSmemBytes = SMEM_TOTAL;
cudaLaunchAttribute at[1];
at[0].id = cudaLaunchAttributeClusterDimension;
at[0].val.clusterDim.x = S; at[0].val.clusterDim.y = 1; at[0].val.clusterDim.z = 1;
cfg.attrs = at; cfg.numAttrs = 1;
int n = 0;
cudaError_t e = cudaOccupancyMaxActiveClusters(&n, rollout_kernel, &cfg);
if (e != cudaSuccess || n <= 0) {
cudaGetLastError();
n = 46;
}
g_max_clusters = n;
}
static void launch_rollout(const Params& p, int nclusters, cudaStream_t stream) {
cudaLaunchConfig_t cfg = {};
cfg.gridDim = dim3(S * nclusters, 1, 1);
cfg.blockDim = dim3(NT, 1, 1);
cfg.dynamicSmemBytes = SMEM_TOTAL;
cfg.stream = stream;
cudaLaunchAttribute at[1];
at[0].id = cudaLaunchAttributeClusterDimension;
at[0].val.clusterDim.x = S; at[0].val.clusterDim.y = 1; at[0].val.clusterDim.z = 1;
cfg.attrs = at; cfg.numAttrs = 1;
cudaError_t e = cudaLaunchKernelEx(&cfg, rollout_kernel, p);
TORCH_CHECK(e == cudaSuccess, "rollout launch failed: ", cudaGetErrorString(e));
}
int max_clusters() { ensure_attrs(); return g_max_clusters; }
// Workspace cached across calls with the same (N, H); outputs are fresh tensors each call.
struct Workspace {
int N = -1, H = -1, nclusters = 0, nblocks = 0, npad = 0;
torch::Tensor state, slices, partials, env, env_init, flags;
};
static Workspace g_ws;
std::vector<torch::Tensor> rollout(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc, torch::Tensor benc,
torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv, int64_t N, int64_t H, int64_t seed) {
ensure_attrs();
const c10::cuda::CUDAGuard guard(Bgru.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto dev = Bgru.device();
auto f32 = torch::TensorOptions().dtype(torch::kFloat32).device(dev);
auto i32 = torch::TensorOptions().dtype(torch::kInt32).device(dev);
auto i64 = torch::TensorOptions().dtype(torch::kInt64).device(dev);
Workspace& ws = g_ws;
if (ws.N != (int)N || ws.H != (int)H || !ws.state.defined() || ws.state.device() != dev) {
ws.N = (int)N; ws.H = (int)H;
ws.nblocks = (int)((N + R - 1) / R);
ws.npad = ws.nblocks * R;
ws.nclusters = std::min(ws.nblocks, g_max_clusters);
ws.state = torch::empty({(int64_t)ws.npad, 3, HID}, f32);
ws.slices = torch::empty({(int64_t)ws.nclusters * 2 * S * 2 * CHUNK_U4 * 4}, i32);
ws.partials = torch::empty({(int64_t)ws.nclusters * S * R * 8}, f32);
ws.env = torch::empty({(int64_t)S * ws.npad * 8}, i32);
ws.env_init = torch::empty({(int64_t)ws.npad * 4}, i32);
ws.flags = torch::empty({2 * (int64_t)H}, i32);
}
cudaMemsetAsync(ws.flags.data_ptr<int>(), 0, sizeof(int) * 2 * H, stream);
auto rewards = torch::empty({N}, f32);
auto positions = torch::empty({N, 2}, i64);
auto logits = torch::empty({N, 4}, f32);
mt_init_kernel<<<1, 640, 0, stream>>>(ws.env_init.data_ptr<int>(), (int)N, ws.npad, (unsigned)(uint64_t)seed);
Params p;
p.Bgru = reinterpret_cast<const uint4*>(Bgru.data_ptr());
p.Ba = reinterpret_cast<const uint4*>(Ba.data_ptr());
p.Wenc = Wenc.data_ptr<float>();
p.benc = benc.data_ptr<float>();
p.W1p = W1p.data_ptr<float>();
p.b1p = b1p.data_ptr<float>();
p.ba = ba.data_ptr<float>();
p.bv = bv.data_ptr<float>();
p.state_in = ws.state.data_ptr<float>();
p.state_out = ws.state.data_ptr<float>();
p.slices = reinterpret_cast<uint4*>(ws.slices.data_ptr());
p.partials = ws.partials.data_ptr<float>();
p.env = ws.env.data_ptr<int>();
p.env_init = ws.env_init.data_ptr<int>();
p.flags = ws.flags.data_ptr<int>();
p.done = reinterpret_cast<unsigned int*>(ws.flags.data_ptr<int>() + H);
p.out_rewards = rewards.data_ptr<float>();
p.out_pos = reinterpret_cast<long long*>(positions.data_ptr<int64_t>());
p.out_logits = logits.data_ptr<float>();
p.out_value = nullptr;
p.obs_in = nullptr;
p.N = (int)N; p.npad = ws.npad; p.nblocks = ws.nblocks; p.nclusters = ws.nclusters; p.H = (int)H;
p.seed = (long long)seed;
p.mode = MODE_RUN;
launch_rollout(p, ws.nclusters, stream);
return {rewards, positions, logits, ws.state.narrow(0, 0, N)};
}
std::vector<torch::Tensor> policy_forward_cuda(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc,
torch::Tensor benc, torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv,
torch::Tensor obs, torch::Tensor state) {
ensure_attrs();
const c10::cuda::CUDAGuard guard(Bgru.device());
auto stream = at::cuda::getCurrentCUDAStream();
const int64_t N = obs.size(0);
const int nblocks = (int)((N + R - 1) / R);
const int nclusters = std::min(nblocks, g_max_clusters);
auto dev = Bgru.device();
auto f32 = torch::TensorOptions().dtype(torch::kFloat32).device(dev);
auto i32 = torch::TensorOptions().dtype(torch::kInt32).device(dev);
auto slices = torch::empty({(int64_t)nclusters * 2 * S * 2 * CHUNK_U4 * 4}, i32);
auto partials = torch::empty({(int64_t)nclusters * S * R * 8}, f32);
auto logits = torch::empty({N, 4}, f32);
auto value = torch::empty({N}, f32);
auto new_state = torch::empty({N, 3, HID}, f32);
Params p;
p.Bgru = reinterpret_cast<const uint4*>(Bgru.data_ptr());
p.Ba = reinterpret_cast<const uint4*>(Ba.data_ptr());
p.Wenc = Wenc.data_ptr<float>();
p.benc = benc.data_ptr<float>();
p.W1p = W1p.data_ptr<float>();
p.b1p = b1p.data_ptr<float>();
p.ba = ba.data_ptr<float>();
p.bv = bv.data_ptr<float>();
p.state_in = state.data_ptr<float>();
p.state_out = new_state.data_ptr<float>();
p.slices = reinterpret_cast<uint4*>(slices.data_ptr());
p.partials = partials.data_ptr<float>();
p.env = nullptr; p.env_init = nullptr; p.flags = nullptr; p.done = nullptr;
p.out_rewards = nullptr; p.out_pos = nullptr;
p.out_logits = logits.data_ptr<float>();
p.out_value = value.data_ptr<float>();
p.obs_in = obs.data_ptr<float>();
p.N = (int)N; p.npad = nblocks * R; p.nblocks = nblocks; p.nclusters = nclusters; p.H = 1;
p.seed = 0;
p.mode = MODE_FWD;
launch_rollout(p, nclusters, stream);
return {logits, new_state, value};
}
std::vector<torch::Tensor> env_step_cuda(torch::Tensor agent, torch::Tensor food, torch::Tensor actions,
torch::Tensor rng) {
const c10::cuda::CUDAGuard guard(agent.device());
auto stream = at::cuda::getCurrentCUDAStream();
const int64_t N = agent.size(0);
auto dev = agent.device();
auto f32 = torch::TensorOptions().dtype(torch::kFloat32).device(dev);
auto i32 = torch::TensorOptions().dtype(torch::kInt32).device(dev);
auto i64 = torch::TensorOptions().dtype(torch::kInt64).device(dev);
auto agent_out = torch::empty({N, 2}, f32);
auto food_out = torch::empty({N, 2}, f32);
auto reward = torch::empty({N}, f32);
auto rng_out = torch::empty({N}, i64);
auto hitbuf = torch::empty({N}, i32);
auto anyhit = torch::zeros({1}, i32);
const int threads = 256;
const int blocks = (int)((N + threads - 1) / threads);
if (N > 0) {
env_move_kernel<<<blocks, threads, 0, stream>>>(agent.data_ptr<float>(), food.data_ptr<float>(),
reinterpret_cast<const long long*>(actions.data_ptr<int64_t>()), agent_out.data_ptr<float>(),
reward.data_ptr<float>(), hitbuf.data_ptr<int>(),
anyhit.data_ptr<int>(), (int)N);
env_respawn_kernel<<<blocks, threads, 0, stream>>>(food.data_ptr<float>(), hitbuf.data_ptr<int>(),
reinterpret_cast<const long long*>(rng.data_ptr<int64_t>()), anyhit.data_ptr<int>(),
food_out.data_ptr<float>(), reinterpret_cast<long long*>(rng_out.data_ptr<int64_t>()),
(int)N);
}
return {agent_out, food_out, reward, rng_out};
}
} // namespace gm
int gm_max_clusters() { return gm::max_clusters(); }
std::vector<torch::Tensor> gm_rollout(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc, torch::Tensor benc,
torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv, int64_t N, int64_t H, int64_t seed) {
return gm::rollout(Bgru, Ba, Wenc, benc, W1p, b1p, ba, bv, N, H, seed);
}
std::vector<torch::Tensor> gm_policy_forward(torch::Tensor Bgru, torch::Tensor Ba, torch::Tensor Wenc,
torch::Tensor benc, torch::Tensor W1p, torch::Tensor b1p,
torch::Tensor ba, torch::Tensor bv,
torch::Tensor obs, torch::Tensor state) {
return gm::policy_forward_cuda(Bgru, Ba, Wenc, benc, W1p, b1p, ba, bv, obs, state);
}
std::vector<torch::Tensor> gm_env_step(torch::Tensor agent, torch::Tensor food, torch::Tensor actions,
torch::Tensor rng) {
return gm::env_step_cuda(agent, food, actions, rng);
}
"""
_CUDA_SRC = _KERNEL_SRC + "\n" + _HOST_SRC
_ext = None
def _get_ext():
global _ext
if _ext is None:
from torch.utils.cpp_extension import load_inline
_ext = load_inline(
name="gm_mingru_sps_" + hashlib.sha1(_CUDA_SRC.encode()).hexdigest()[:10],
cpp_sources=_CPP_DECLS,
cuda_sources=_CUDA_SRC,
functions=["gm_max_clusters", "gm_rollout", "gm_policy_forward", "gm_env_step"],
extra_cuda_cflags=["-O3", "-std=c++17", "-lineinfo", "--expt-relaxed-constexpr"],
verbose=False,
)
return _ext
# ----------------------------------------------------------------------------- model (same params as reference)
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 forward(self, obs: torch.Tensor, state: torch.Tensor):
return policy_forward(self, obs, state)
# ----------------------------------------------------------------------------- weight packing
_OFF = [0, 1, 8, 9, 16, 17, 24, 25] # permuted-k layout inside a k32 chunk: 8 values per lane-quad
_WCACHE: dict = {}
def _split_hilo(x: torch.Tensor):
hi = x.half()
lo = (x - hi.float()).half()
return hi, lo
def _pack_weights(model: Model):
"""Pre-permute weights into mma.m16n8k16 B-fragment order, scaled by 2^11, as fp16 hi/lo pairs."""
params = [model.w_enc, model.b_enc, model.w_gru, model.w_a, model.b_a, model.w_v, model.b_v]
key = (id(model), tuple(p.data_ptr() for p in params), tuple(p._version for p in params))
hit = _WCACHE.get(key)
if hit is not None:
return hit
dev = model.w_gru.device
with torch.no_grad():
w_gru = model.w_gru.detach().float() # (3, 768, 256)
off = torch.tensor(_OFF, device=dev)
c, w, ch, g, t, i = torch.meshgrid(
torch.arange(4, device=dev), torch.arange(8, device=dev), torch.arange(8, device=dev),
torch.arange(3, device=dev), torch.arange(32, device=dev), torch.arange(8, device=dev),
indexing="ij",
)
u = 64 * c + 8 * w + t // 4
k = 32 * ch + 2 * (t % 4) + off[i]
rows = g * HIDDEN + u
vals = w_gru[:, rows, k] * 2048.0 # (3, 4, 8, 8, 3, 32, 8)
hi, lo = _split_hilo(vals)
bgru = torch.stack([hi, lo], dim=4).contiguous() # (3,4,8,8,2,3,32,8)
wa_ext = torch.zeros(8, HIDDEN, device=dev, dtype=torch.float32)
wa_ext[:4] = model.w_a.detach().float()
wa_ext[4] = model.w_v.detach().float()[0]
c2, ch2, t2, i2 = torch.meshgrid(
torch.arange(4, device=dev), torch.arange(2, device=dev), torch.arange(32, device=dev),
torch.arange(8, device=dev), indexing="ij",
)
a = t2 // 4
k2 = 32 * ch2 + 2 * (t2 % 4) + off[i2]
u2 = 64 * c2 + k2
vals2 = wa_ext[a, u2] * 2048.0 # (4, 2, 32, 8)
hi2, lo2 = _split_hilo(vals2)
ba_pack = torch.stack([hi2, lo2], dim=2).contiguous() # (4,2,2,32,8)
# layer 1 collapsed: gates_1 = W_1 (W_enc obs + b_enc) = (W_1 W_enc) obs + W_1 b_enc (composed in fp64)
w1 = model.w_gru.detach()[0].double()
w1p = (w1 @ model.w_enc.detach().double()) * 2048.0 # (768, 4)
b1p = (w1 @ model.b_enc.detach().double()) * 2048.0 # (768,)
packed = (
bgru.view(torch.int16).contiguous(),
ba_pack.view(torch.int16).contiguous(),
model.w_enc.detach().float().contiguous(),
model.b_enc.detach().float().contiguous(),
w1p.float().contiguous(),
b1p.float().contiguous(),
model.b_a.detach().float().contiguous(),
model.b_v.detach().float().contiguous(),
)
_WCACHE.clear()
_WCACHE[key] = packed
return packed
# ----------------------------------------------------------------------------- public API
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,)."""
ext = _get_ext()
pw = _pack_weights(model)
obs_c = obs.detach().float().contiguous()
state_c = state.detach().float().contiguous()
logits, new_state, value = ext.gm_policy_forward(*pw, obs_c, state_c)
return logits, new_state, value
def env_step(agent: torch.Tensor, food: torch.Tensor, actions: torch.Tensor, rng_state: torch.Tensor):
ext = _get_ext()
a, f, r, g = ext.gm_env_step(
agent.detach().float().contiguous(),
food.detach().float().contiguous(),
actions.detach().long().contiguous(),
rng_state.detach().long().contiguous(),
)
return a, f, r, g
def run(num_envs: int, horizon: int, seed: int, model: Model | None = None) -> dict:
device = torch.device("cuda:0")
if model is None:
model = Model()
model = model.to(device).eval()
ext = _get_ext()
pw = _pack_weights(model)
if int(num_envs) <= 0 or int(horizon) <= 0:
return {
"rewards": torch.zeros(max(int(num_envs), 0), device=device),
"positions": torch.zeros(max(int(num_envs), 0), 2, dtype=torch.int64, device=device),
"last_logits": torch.zeros(max(int(num_envs), 0), NUM_ACTIONS, device=device),
"state": torch.zeros(max(int(num_envs), 0), GRU_LAYERS, HIDDEN, device=device),
}
with torch.no_grad():
rewards, positions, last_logits, state = ext.gm_rollout(*pw, int(num_envs), int(horizon), int(seed))
return {
"rewards": rewards,
"positions": positions,
"last_logits": last_logits,
"state": state,
}
def get_init_inputs():
return []
def get_inputs():
return []
20260904_111027_or-fable_anthropic_claude-fable-5-1_04_grid_mingru_sps