"""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 #include #include #include #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 __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<<>>( obs.data_ptr(), w_enc.data_ptr(), b_enc.data_ptr(), h0.data_ptr(), 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 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, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM); ok = (e == cudaSuccess) ? 1 : -1; } if (ok < 0) return; dim3 grid(GRU_OUT / BN, M / BM); layer_kernel<<>>(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(), st.data_ptr(), Wt.data_ptr(), h_out.data_ptr(), 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<<>>( h.data_ptr(), w_a.data_ptr(), b_a.data_ptr(), w_v.data_ptr(), b_v.data_ptr(), logits.data_ptr(), value.data_ptr(), reinterpret_cast(actions.data_ptr()), 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<<>>( agent.data_ptr(), food.data_ptr(), reinterpret_cast(actions.data_ptr()), agent_out.data_ptr(), hit.data_ptr(), flag.data_ptr(), 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<<>>( hit.data_ptr(), food.data_ptr(), reinterpret_cast(rng.data_ptr()), flag.data_ptr(), other_flag.data_ptr(), reward.data_ptr(), agent.data_ptr(), obs.data_ptr(), 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<<>>(agent.data_ptr(), food.data_ptr(), obs.data_ptr(), 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 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)