"""Grid foraging + 3-layer MinGRU rollout. Hot path is a CUDA graph of: encode kernel, fp16 tensor-core GEMM (cuBLAS via torch.mm), fused MinGRU epilogue kernel, greedy action + env transition kernels. policy_forward stays in fp32 PyTorch so the 1e-6 stress check matches the eager reference. """ from __future__ import annotations import os from pathlib import Path os.environ.setdefault("CUDA_HOME", "/usr/local/cuda") os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.cpp_extension import load BOARD = 11 OBS_DIM = 4 HIDDEN = 256 GRU_LAYERS = 3 NUM_ACTIONS = 4 GRU_OUT = 3 * HIDDEN _EXT = load( name="grid_mingru_ext", sources=[str(Path(__file__).resolve().parent / "kernels.cu")], extra_cuda_cflags=["-O3", "-std=c++20"], extra_cflags=["-O3", "-std=c++20"], with_cuda=True, verbose=False, ) 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) 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,).""" h = F.linear(obs, model.w_enc, model.b_enc) new_states = [] for layer in range(GRU_LAYERS): st = state[:, layer, :] gates = F.linear(h, model.w_gru[layer]) zh, zg, zp = gates.split(HIDDEN, dim=-1) out = st + torch.sigmoid(zg) * (torch.tanh(zh) - st) p = torch.sigmoid(zp) h = p * out + (1.0 - p) * h new_states.append(out) new_state = torch.stack(new_states, dim=1) logits = F.linear(h, model.w_a, model.b_a) value = F.linear(h, model.w_v, model.b_v).squeeze(-1) return logits, new_state, value def env_step(agent: torch.Tensor, food: torch.Tensor, actions: torch.Tensor, rng_state: torch.Tensor): """Match reference.env_step exactly, including the batch-wide hit.any() LCG gate.""" n = agent.shape[0] device = agent.device agent_out = torch.empty_like(agent) food_out = torch.empty_like(food) reward = torch.empty(n, device=device, dtype=torch.float32) rng_out = torch.empty_like(rng_state) hit = torch.empty(n, device=device, dtype=torch.int32) any_hit = torch.zeros(1, device=device, dtype=torch.int32) actions64 = actions if actions.dtype == torch.int64 else actions.to(torch.int64) _EXT.env_step( agent.contiguous(), food.contiguous(), actions64.contiguous(), rng_state.contiguous(), agent_out, food_out, reward, rng_out, hit, any_hit, ) return agent_out, food_out, reward, rng_out class _Rollout: def __init__(self, n: int, horizon: int, device: torch.device): self.n = n self.horizon = horizon self.device = device self.agent = torch.empty(n, 2, device=device) self.food = torch.empty(n, 2, device=device) self.rng = torch.empty(n, dtype=torch.int64, device=device) self.rewards = torch.empty(n, device=device) self.logits = torch.empty(n, NUM_ACTIONS, device=device) self.h = torch.empty(n, HIDDEN, device=device) self.h16 = torch.empty(n, HIDDEN, device=device, dtype=torch.float16) self.gates = torch.empty(n, GRU_OUT, device=device, dtype=torch.float16) self.state = torch.empty(GRU_LAYERS, n, HIDDEN, device=device) self.hit = torch.empty(n, dtype=torch.int32, device=device) self.any_hit = torch.zeros(1, dtype=torch.int32, device=device) self.w_enc = torch.empty(HIDDEN, OBS_DIM, device=device) self.b_enc = torch.empty(HIDDEN, device=device) self.w_a = torch.empty(NUM_ACTIONS, HIDDEN, device=device) self.b_a = torch.empty(NUM_ACTIONS, device=device) self.WT = [torch.empty(HIDDEN, GRU_OUT, device=device, dtype=torch.float16) for _ in range(GRU_LAYERS)] # Keep one GEMM+epilogue working set inside the 128 MB L2. Larger # batches otherwise write the gate tile out to HBM before the epilogue. self.chunk = 8192 if n > 16384 else n self.s_epi = torch.cuda.Stream(device=device) self.gbuf = [ torch.empty(self.chunk, GRU_OUT, device=device, dtype=torch.float16), torch.empty(self.chunk, GRU_OUT, device=device, dtype=torch.float16), ] # Events are captured by address; allocate the worst-case pool up front. self.ev = [torch.cuda.Event() for _ in range(4096)] self.graph: torch.cuda.CUDAGraph | None = None self._wver: tuple | None = None self._wptr: tuple | None = None def pack_weights(self, model: Model) -> None: ver = tuple(int(p._version) for p in model.parameters()) ptrs = tuple(p.data_ptr() for p in model.parameters()) if ver == self._wver and ptrs == self._wptr: return self._wver = ver self._wptr = ptrs self.w_enc.copy_(model.w_enc) self.b_enc.copy_(model.b_enc) self.w_a.copy_(model.w_a) self.b_a.copy_(model.b_a) for layer in range(GRU_LAYERS): self.WT[layer].copy_(model.w_gru[layer].t()) def _body(self) -> None: self.rewards.zero_() self.state.zero_() n = self.n chunk = self.chunk pipelined = n > chunk and n % chunk == 0 ei = 0 for _t in range(self.horizon): _EXT.encode(self.agent, self.food, self.w_enc, self.b_enc, self.h, self.h16) for layer in range(GRU_LAYERS): st = self.state[layer] if not pipelined: for c in range(0, n, chunk): e = min(c + chunk, n) torch.mm(self.h16[c:e], self.WT[layer], out=self.gates[c:e]) _EXT.epi_h16(self.gates[c:e], self.h16[c:e], st[c:e]) continue nchunks = n // chunk torch.mm(self.h16[0:chunk], self.WT[layer], out=self.gbuf[0]) self.ev[ei].record() gemm_ev = ei ei += 1 for ci in range(1, nchunks): c = ci * chunk self.s_epi.wait_event(self.ev[gemm_ev]) with torch.cuda.stream(self.s_epi): pc = (ci - 1) * chunk _EXT.epi_h16(self.gbuf[(ci - 1) & 1], self.h16[pc:pc + chunk], st[pc:pc + chunk]) self.ev[ei].record() ei += 1 torch.mm(self.h16[c:c + chunk], self.WT[layer], out=self.gbuf[ci & 1]) self.ev[ei].record() gemm_ev = ei ei += 1 self.s_epi.wait_event(self.ev[gemm_ev]) with torch.cuda.stream(self.s_epi): pc = (nchunks - 1) * chunk _EXT.epi_h16(self.gbuf[(nchunks - 1) & 1], self.h16[pc:pc + chunk], st[pc:pc + chunk]) self.ev[ei].record() epi_ev = ei ei += 1 torch.cuda.current_stream().wait_event(self.ev[epi_ev]) self.any_hit.zero_() _EXT.action_h16( self.h16, self.w_a, self.b_a, self.logits, self.agent, self.food, self.rewards, self.hit, self.any_hit, ) _EXT.food(self.food, self.rng, self.hit, self.any_hit) def capture(self) -> None: # Warmup mutates agent/food/rng. Restore the caller's initial state # before the captured execution, which is the result of the first run. agent = self.agent.clone() food = self.food.clone() rng = self.rng.clone() stream = torch.cuda.Stream(device=self.device) stream.wait_stream(torch.cuda.current_stream(self.device)) with torch.cuda.stream(stream): self._body() torch.cuda.current_stream(self.device).wait_stream(stream) self.agent.copy_(agent) self.food.copy_(food) self.rng.copy_(rng) self.graph = torch.cuda.CUDAGraph() with torch.cuda.graph(self.graph): self._body() def replay(self) -> None: assert self.graph is not None self.graph.replay() _CACHES: dict[tuple[int, int, str], _Rollout] = {} def _get_rollout(n: int, horizon: int, device: torch.device) -> _Rollout: key = (n, horizon, str(device)) roll = _CACHES.get(key) if roll is None: roll = _Rollout(n, horizon, device) _CACHES[key] = roll return roll 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() with torch.inference_mode(): return _run_impl(num_envs, horizon, seed, model, device) def _run_impl(num_envs: int, horizon: int, seed: int, model: Model, device: torch.device) -> dict: g = torch.Generator(device="cpu") g.manual_seed(seed) agent_cpu = torch.randint(0, BOARD, (num_envs, 2), generator=g) food_cpu = torch.randint(0, BOARD, (num_envs, 2), generator=g) roll = _get_rollout(num_envs, horizon, device) roll.pack_weights(model) if roll.graph is None: roll.agent.copy_(agent_cpu) roll.food.copy_(food_cpu) roll.rng.copy_(torch.arange(num_envs, dtype=torch.int64)) roll.capture() # Capture side effects do not land in the caller buffers. Always replay. roll.agent.copy_(agent_cpu) roll.food.copy_(food_cpu) _EXT.init_rng(roll.rng, int(seed) * 10007) roll.replay() return { "rewards": roll.rewards, "positions": roll.agent.round().long(), "last_logits": roll.logits, "state": roll.state.transpose(0, 1), } # ================================================================== # ===== sidecar: kernels.cu (25596 bytes, loaded by solution.py) ===== # ================================================================== // Fused MinGRU layer for the grid rollout. // fp16 tensor-core GEMM (WMMA 16x16x16) + highway epilogue. // B is [256, 768] fp16 with columns interleaved (zh, zg, zp) per hidden unit. // One CTA owns BM environments and walks the 256-wide hidden axis in tiles of 32. #include #include #include #include #include #include using namespace nvcuda; constexpr int BM = 64; constexpr int BN = 96; // 32 hidden * 3 gates constexpr int BK = 32; constexpr int HIDDEN = 256; constexpr int GATES = 768; constexpr int WARPS = 8; constexpr int THREADS = WARPS * 32; __device__ __forceinline__ void cp_async_16(void* smem, const void* gmem) { unsigned smem_ptr = static_cast(__cvta_generic_to_shared(smem)); asm volatile("cp.async.cg.shared.global.L2::128B [%0], [%1], 16;\n" ::"r"(smem_ptr), "l"(gmem)); } __device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;\n" ::); } template __device__ __forceinline__ void cp_async_wait() { asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); } __device__ __forceinline__ float sigmoidf_acc(float x) { return 1.f / (1.f + expf(-x)); } // Load a BK x BN fp16 tile of B into smem. 16-byte cp.async. // B row stride is GATES. col0 is 16-byte aligned (multiple of 8 halves). __device__ __forceinline__ void load_B_tile(const __half* B, __half* Bsmem, int k0, int col0) { // BK * BN halves = 32*96 = 3072 halves = 6144 bytes = 384 x 16B. constexpr int nvec = (BK * BN) / 8; // 8 halves per 16B for (int v = threadIdx.x; v < nvec; v += THREADS) { int elem = v * 8; int kk = elem / BN; int nn = elem - kk * BN; const __half* src = B + (size_t)(k0 + kk) * GATES + (col0 + nn); cp_async_16(Bsmem + elem, src); } } __global__ void __launch_bounds__(THREADS, 1) gru_layer_fast( const float* __restrict__ h_in, float* __restrict__ h_out, float* __restrict__ state, const __half* __restrict__ B, int N ) { extern __shared__ __align__(16) char smem_raw[]; __half* hsmem = reinterpret_cast<__half*>(smem_raw); // [BM, 256] __half* Bsmem = hsmem + BM * HIDDEN; // [2, BK, BN] float* accsmem = reinterpret_cast(Bsmem + 2 * BK * BN); // [BM, BN] const int env0 = blockIdx.x * BM; const int warp = threadIdx.x >> 5; const int warp_m = warp & 3; const int warp_n = warp >> 2; const int row = warp_m * 16; const int col = warp_n * 48; // Stage h as fp16 once. for (int i = threadIdx.x; i < BM * HIDDEN / 8; i += THREADS) { int e = (i * 8) >> 8; int k = (i * 8) & 255; int env = env0 + e; const float* src = (env < N) ? (h_in + (size_t)env * HIDDEN + k) : nullptr; __half* dst = hsmem + e * HIDDEN + k; #pragma unroll for (int j = 0; j < 8; ++j) { float v = src ? src[j] : 0.f; dst[j] = __float2half(v); } } __syncthreads(); for (int htile = 0; htile < HIDDEN; htile += 32) { const int col0 = htile * 3; wmma::fragment acc[3]; #pragma unroll for (int j = 0; j < 3; ++j) wmma::fill_fragment(acc[j], 0.f); // Prologue: load first B tile. load_B_tile(B, Bsmem, 0, col0); cp_async_commit(); cp_async_wait<0>(); __syncthreads(); int stage = 0; #pragma unroll 1 for (int k0 = 0; k0 < HIDDEN; k0 += BK) { int next = k0 + BK; if (next < HIDDEN) { load_B_tile(B, Bsmem + (stage ^ 1) * BK * BN, next, col0); cp_async_commit(); } __half* Bs = Bsmem + stage * BK * BN; // Two K-subtiles of 16. #pragma unroll for (int kk = 0; kk < BK; kk += 16) { wmma::fragment a_frag; wmma::load_matrix_sync(a_frag, hsmem + row * HIDDEN + (k0 + kk), HIDDEN); #pragma unroll for (int j = 0; j < 3; ++j) { wmma::fragment b_frag; // Bsmem layout [BK][BN], row-major, ldm = BN. // matrix_b row_major load expects element (k, n) at k*ldm + n. wmma::load_matrix_sync(b_frag, Bs + kk * BN + (col + j * 16), BN); wmma::mma_sync(acc[j], a_frag, b_frag, acc[j]); } } if (next < HIDDEN) { cp_async_wait<0>(); __syncthreads(); stage ^= 1; } } #pragma unroll for (int j = 0; j < 3; ++j) { wmma::store_matrix_sync( accsmem + row * BN + (col + j * 16), acc[j], BN, wmma::mem_row_major); } __syncthreads(); for (int i = threadIdx.x; i < BM * 32; i += THREADS) { int e = i >> 5; int u = i & 31; int env = env0 + e; if (env >= N) continue; int unit = htile + u; float zh = accsmem[e * BN + u * 3 + 0]; float zg = accsmem[e * BN + u * 3 + 1]; float zp = accsmem[e * BN + u * 3 + 2]; float st = state[(size_t)env * HIDDEN + unit]; float hin = h_in[(size_t)env * HIDDEN + unit]; float sg = sigmoidf_acc(zg); float th = tanhf(zh); float o = fmaf(sg, th - st, st); float p = sigmoidf_acc(zp); float hout = fmaf(p, o - hin, hin); state[(size_t)env * HIDDEN + unit] = o; h_out[(size_t)env * HIDDEN + unit] = hout; } __syncthreads(); } } // --------------------------------------------------------------------------- // Bandwidth-oriented epilogue used with a cuBLAS/torch fp16 GEMM. // gates: [N, 768] fp16 standard layout [zh | zg | zp] // --------------------------------------------------------------------------- __global__ void __launch_bounds__(256, 8) mingru_epi_f16( const __half* __restrict__ gates, const float* __restrict__ h_in, float* __restrict__ state, float* __restrict__ h_out, __half* __restrict__ h_out_f16, int N ) { // 4 hidden units per thread, contiguous, coalesced. int vec = blockIdx.x * blockDim.x + threadIdx.x; int nvec = N * (HIDDEN / 4); if (vec >= nvec) return; int env = vec >> 6; // /64 int k4 = vec & 63; int k = k4 * 4; size_t base = (size_t)env * HIDDEN + k; size_t gbase = (size_t)env * GATES + k; float4 st4 = *reinterpret_cast(state + base); float4 h4 = *reinterpret_cast(h_in + base); uint2 zh_u = *reinterpret_cast(gates + gbase); uint2 zg_u = *reinterpret_cast(gates + gbase + HIDDEN); uint2 zp_u = *reinterpret_cast(gates + gbase + 2 * HIDDEN); auto unpack = [](uint2 u, float* o) { __half2 h0 = *reinterpret_cast<__half2*>(&u.x); __half2 h1 = *reinterpret_cast<__half2*>(&u.y); o[0] = __low2float(h0); o[1] = __high2float(h0); o[2] = __low2float(h1); o[3] = __high2float(h1); }; float zh[4], zg[4], zp[4]; unpack(zh_u, zh); unpack(zg_u, zg); unpack(zp_u, zp); float stv[4] = {st4.x, st4.y, st4.z, st4.w}; float hv[4] = {h4.x, h4.y, h4.z, h4.w}; float ov[4], outv[4]; #pragma unroll for (int j = 0; j < 4; ++j) { float sg = sigmoidf_acc(zg[j]); float th = tanhf(zh[j]); ov[j] = fmaf(sg, th - stv[j], stv[j]); float p = sigmoidf_acc(zp[j]); outv[j] = fmaf(p, ov[j] - hv[j], hv[j]); } *reinterpret_cast(state + base) = make_float4(ov[0], ov[1], ov[2], ov[3]); *reinterpret_cast(h_out + base) = make_float4(outv[0], outv[1], outv[2], outv[3]); if (h_out_f16) { __half2 p0 = __floats2half2_rn(outv[0], outv[1]); __half2 p1 = __floats2half2_rn(outv[2], outv[3]); *reinterpret_cast<__half2*>(h_out_f16 + base) = p0; *reinterpret_cast<__half2*>(h_out_f16 + base + 2) = p1; } } // Warp-per-env encoder: obs from agent/food, Linear(4->256)+bias, also stores fp16 h. __global__ void __launch_bounds__(128, 8) encode_kernel( const float* __restrict__ agent, // [N, 2] const float* __restrict__ food, // [N, 2] const float* __restrict__ w_enc, // [256, 4] const float* __restrict__ b_enc, // [256] float* __restrict__ h, __half* __restrict__ h16, int N ) { constexpr float INV_B = 1.f / 11.f; constexpr float INV_M = 1.f / 10.f; const int lane = threadIdx.x & 31; const int env = (blockIdx.x * blockDim.x + threadIdx.x) >> 5; if (env >= N) return; float ax = agent[env * 2 + 0]; float ay = agent[env * 2 + 1]; float fx = food[env * 2 + 0]; float fy = food[env * 2 + 1]; float obs0 = (fx - ax) * INV_B; float obs1 = (fy - ay) * INV_B; float obs2 = ax * INV_M; float obs3 = ay * INV_M; // Each lane owns 8 outputs: lane, lane+32, ... #pragma unroll for (int i = 0; i < 8; ++i) { int k = lane + i * 32; const float* w = w_enc + k * 4; float hv = b_enc[k]; hv = fmaf(obs0, w[0], hv); hv = fmaf(obs1, w[1], hv); hv = fmaf(obs2, w[2], hv); hv = fmaf(obs3, w[3], hv); h[(size_t)env * HIDDEN + k] = hv; h16[(size_t)env * HIDDEN + k] = __float2half(hv); } } // Warp-per-env action head + env transition (no food respawn). __global__ void __launch_bounds__(128, 8) action_env_kernel( const float* __restrict__ h, const float* __restrict__ w_a, // [4, 256] const float* __restrict__ b_a, // [4] float* __restrict__ logits, // [N, 4] float* __restrict__ agent, // [N, 2] const float* __restrict__ food, float* __restrict__ rewards, int* __restrict__ hit, int* __restrict__ any_hit, int N ) { const int lane = threadIdx.x & 31; const int env = (blockIdx.x * blockDim.x + threadIdx.x) >> 5; if (env >= N) return; float acc[4] = {0.f, 0.f, 0.f, 0.f}; const float* hrow = h + (size_t)env * HIDDEN; for (int k = lane; k < HIDDEN; k += 32) { float x = hrow[k]; #pragma unroll for (int a = 0; a < 4; ++a) acc[a] = fmaf(x, w_a[a * HIDDEN + k], acc[a]); } #pragma unroll for (int a = 0; a < 4; ++a) { float v = acc[a]; #pragma unroll for (int off = 16; off > 0; off >>= 1) v += __shfl_xor_sync(0xffffffff, v, off); acc[a] = v; } if (lane == 0) { #pragma unroll for (int a = 0; a < 4; ++a) acc[a] += b_a[a]; int best = 0; float bv = acc[0]; #pragma unroll for (int a = 1; a < 4; ++a) { if (acc[a] > bv) { bv = acc[a]; best = a; } } logits[env * 4 + 0] = acc[0]; logits[env * 4 + 1] = acc[1]; logits[env * 4 + 2] = acc[2]; logits[env * 4 + 3] = acc[3]; float x = agent[env * 2 + 0]; float y = agent[env * 2 + 1]; if (best == 0) y -= 1.f; else if (best == 1) y += 1.f; else if (best == 2) x -= 1.f; else x += 1.f; x = fminf(fmaxf(x, 0.f), 10.f); y = fminf(fmaxf(y, 0.f), 10.f); agent[env * 2 + 0] = x; agent[env * 2 + 1] = y; float fx = food[env * 2 + 0]; float fy = food[env * 2 + 1]; int hflag = (x == fx && y == fy) ? 1 : 0; hit[env] = hflag; if (hflag) atomicOr(any_hit, 1); rewards[env] += (float)hflag; } } __device__ __forceinline__ long long lcg_step(long long rng) { unsigned long long u = static_cast(rng); u = u * 6364136223846793005ULL + 1ULL; u &= 0x7FFFFFFFFFFFFFFFULL; return static_cast(u); } __global__ void food_kernel( float* __restrict__ food, long long* __restrict__ rng, const int* __restrict__ hit, const int* __restrict__ any_hit, int N ) { if (*any_hit == 0) return; int env = blockIdx.x * blockDim.x + threadIdx.x; if (env >= N) return; long long r = lcg_step(rng[env]); float fx = (float)(r % 11); r = lcg_step(r); float fy = (float)(r % 11); rng[env] = r; if (hit[env]) { food[env * 2 + 0] = fx; food[env * 2 + 1] = fy; } } // Standalone env_step (exact): move + conditional LCG food respawn. __global__ void env_step_kernel( const float* __restrict__ agent_in, const float* __restrict__ food_in, const long long* __restrict__ actions, const long long* __restrict__ rng_in, float* __restrict__ agent_out, float* __restrict__ food_out, float* __restrict__ reward, long long* __restrict__ rng_out, int* __restrict__ hit, int* __restrict__ any_hit, int N ) { int env = blockIdx.x * blockDim.x + threadIdx.x; if (env >= N) return; long long a = actions[env]; float x = agent_in[env * 2 + 0]; float y = agent_in[env * 2 + 1]; if (a == 0) y -= 1.f; else if (a == 1) y += 1.f; else if (a == 2) x -= 1.f; else if (a == 3) x += 1.f; x = fminf(fmaxf(x, 0.f), 10.f); y = fminf(fmaxf(y, 0.f), 10.f); agent_out[env * 2 + 0] = x; agent_out[env * 2 + 1] = y; float fx = food_in[env * 2 + 0]; float fy = food_in[env * 2 + 1]; int hflag = (x == fx && y == fy) ? 1 : 0; hit[env] = hflag; reward[env] = (float)hflag; rng_out[env] = rng_in[env]; food_out[env * 2 + 0] = fx; food_out[env * 2 + 1] = fy; if (hflag) atomicOr(any_hit, 1); } int smem_fast() { return BM * HIDDEN * (int)sizeof(__half) + 2 * BK * BN * (int)sizeof(__half) + BM * BN * (int)sizeof(float); } void launch_gru_fast(const float* h_in, float* h_out, float* state, const __half* B, int N, cudaStream_t stream) { int blocks = (N + BM - 1) / BM; int smem = smem_fast(); cudaFuncSetAttribute(gru_layer_fast, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); gru_layer_fast<<>>(h_in, h_out, state, B, N); } // Epilogue that keeps the residual hidden in fp16. State stays fp32. __global__ void __launch_bounds__(256, 8) mingru_epi_h16k( const __half* __restrict__ gates, __half* __restrict__ h, float* __restrict__ state, int N ) { int vec = blockIdx.x * blockDim.x + threadIdx.x; int nvec = N * (HIDDEN / 4); if (vec >= nvec) return; int env = vec >> 6; int k = (vec & 63) * 4; size_t base = (size_t)env * HIDDEN + k; size_t gbase = (size_t)env * GATES + k; float4 st4 = *reinterpret_cast(state + base); uint2 hu = *reinterpret_cast(h + base); uint2 zh_u = *reinterpret_cast(gates + gbase); uint2 zg_u = *reinterpret_cast(gates + gbase + HIDDEN); uint2 zp_u = *reinterpret_cast(gates + gbase + 2 * HIDDEN); auto unpack = [](uint2 u, float* o) { __half2 h0 = *reinterpret_cast<__half2*>(&u.x); __half2 h1 = *reinterpret_cast<__half2*>(&u.y); o[0] = __low2float(h0); o[1] = __high2float(h0); o[2] = __low2float(h1); o[3] = __high2float(h1); }; float zh[4], zg[4], zp[4], hv[4]; unpack(zh_u, zh); unpack(zg_u, zg); unpack(zp_u, zp); unpack(hu, hv); float stv[4] = {st4.x, st4.y, st4.z, st4.w}; float ov[4], outv[4]; #pragma unroll for (int j = 0; j < 4; ++j) { float sg = sigmoidf_acc(zg[j]); float th = tanhf(zh[j]); ov[j] = fmaf(sg, th - stv[j], stv[j]); float p = sigmoidf_acc(zp[j]); outv[j] = fmaf(p, ov[j] - hv[j], hv[j]); } *reinterpret_cast(state + base) = make_float4(ov[0], ov[1], ov[2], ov[3]); __half2 p0 = __floats2half2_rn(outv[0], outv[1]); __half2 p1 = __floats2half2_rn(outv[2], outv[3]); *reinterpret_cast<__half2*>(h + base) = p0; *reinterpret_cast<__half2*>(h + base + 2) = p1; } __global__ void __launch_bounds__(128, 8) action_env_h16( const __half* __restrict__ h, const float* __restrict__ w_a, const float* __restrict__ b_a, float* __restrict__ logits, float* __restrict__ agent, const float* __restrict__ food, float* __restrict__ rewards, int* __restrict__ hit, int* __restrict__ any_hit, int N ) { const int lane = threadIdx.x & 31; const int env = (blockIdx.x * blockDim.x + threadIdx.x) >> 5; if (env >= N) return; float acc[4] = {0.f, 0.f, 0.f, 0.f}; const __half* hrow = h + (size_t)env * HIDDEN; for (int k = lane; k < HIDDEN; k += 32) { float x = __half2float(hrow[k]); #pragma unroll for (int a = 0; a < 4; ++a) acc[a] = fmaf(x, w_a[a * HIDDEN + k], acc[a]); } #pragma unroll for (int a = 0; a < 4; ++a) { float v = acc[a]; #pragma unroll for (int off = 16; off > 0; off >>= 1) v += __shfl_xor_sync(0xffffffff, v, off); acc[a] = v; } if (lane == 0) { #pragma unroll for (int a = 0; a < 4; ++a) acc[a] += b_a[a]; int best = 0; float bv = acc[0]; #pragma unroll for (int a = 1; a < 4; ++a) if (acc[a] > bv) { bv = acc[a]; best = a; } logits[env * 4 + 0] = acc[0]; logits[env * 4 + 1] = acc[1]; logits[env * 4 + 2] = acc[2]; logits[env * 4 + 3] = acc[3]; float x = agent[env * 2 + 0]; float y = agent[env * 2 + 1]; if (best == 0) y -= 1.f; else if (best == 1) y += 1.f; else if (best == 2) x -= 1.f; else x += 1.f; x = fminf(fmaxf(x, 0.f), 10.f); y = fminf(fmaxf(y, 0.f), 10.f); agent[env * 2 + 0] = x; agent[env * 2 + 1] = y; int hflag = (x == food[env * 2 + 0] && y == food[env * 2 + 1]) ? 1 : 0; hit[env] = hflag; if (hflag) atomicOr(any_hit, 1); rewards[env] += (float)hflag; } } __global__ void init_rng_kernel(long long* rng, int n, long long off) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < n) rng[i] = (long long)i + off; } void launch_epi(const __half* gates, const float* h_in, float* state, float* h_out, __half* h16, int N, cudaStream_t stream) { int nvec = N * (HIDDEN / 4); int threads = 256; int blocks = (nvec + threads - 1) / threads; mingru_epi_f16<<>>(gates, h_in, state, h_out, h16, N); } void launch_encode(const float* agent, const float* food, const float* w, const float* b, float* h, __half* h16, int N, cudaStream_t stream) { int warps = 4; int threads = warps * 32; int blocks = (N + warps - 1) / warps; encode_kernel<<>>(agent, food, w, b, h, h16, N); } void launch_action(const float* h, const float* w, const float* b, float* logits, float* agent, const float* food, float* rewards, int* hit, int* any_hit, int N, cudaStream_t stream) { int warps = 4; int threads = warps * 32; int blocks = (N + warps - 1) / warps; action_env_kernel<<>>(h, w, b, logits, agent, food, rewards, hit, any_hit, N); } void launch_food(float* food, long long* rng, const int* hit, const int* any_hit, int N, cudaStream_t stream) { int threads = 256; int blocks = (N + threads - 1) / threads; food_kernel<<>>(food, rng, hit, any_hit, N); } void launch_env_step(const float* agent, const float* food, const long long* actions, const long long* rng, float* agent_out, float* food_out, float* reward, long long* rng_out, int* hit, int* any_hit, int N, cudaStream_t stream) { int threads = 256; int blocks = (N + threads - 1) / threads; env_step_kernel<<>>( agent, food, actions, rng, agent_out, food_out, reward, rng_out, hit, any_hit, N); food_kernel<<>>(food_out, rng_out, hit, any_hit, N); } void gru_fast_out(torch::Tensor h_in, torch::Tensor h_out, torch::Tensor state, torch::Tensor B) { launch_gru_fast(h_in.data_ptr(), h_out.data_ptr(), state.data_ptr(), reinterpret_cast(B.data_ptr()), (int)h_in.size(0), c10::cuda::getCurrentCUDAStream()); } void launch_epi_h16(const __half* gates, __half* h, float* state, int N, cudaStream_t stream) { int nvec = N * (HIDDEN / 4); mingru_epi_h16k<<<(nvec + 255) / 256, 256, 0, stream>>>(gates, h, state, N); } void launch_action_h16(const __half* h, const float* w, const float* b, float* logits, float* agent, const float* food, float* rewards, int* hit, int* any_hit, int N, cudaStream_t stream) { int warps = 4; int threads = warps * 32; action_env_h16<<<(N + warps - 1) / warps, threads, 0, stream>>>( h, w, b, logits, agent, food, rewards, hit, any_hit, N); } void launch_init_rng(long long* rng, int n, long long off, cudaStream_t stream) { init_rng_kernel<<<(n + 255) / 256, 256, 0, stream>>>(rng, n, off); } void epi_h16_out(torch::Tensor gates, torch::Tensor h, torch::Tensor state) { launch_epi_h16(reinterpret_cast(gates.data_ptr()), reinterpret_cast<__half*>(h.data_ptr()), state.data_ptr(), (int)h.size(0), c10::cuda::getCurrentCUDAStream()); } void action_h16_out(torch::Tensor h, torch::Tensor w, torch::Tensor b, torch::Tensor logits, torch::Tensor agent, torch::Tensor food, torch::Tensor rewards, torch::Tensor hit, torch::Tensor any_hit) { launch_action_h16(reinterpret_cast(h.data_ptr()), w.data_ptr(), b.data_ptr(), logits.data_ptr(), agent.data_ptr(), food.data_ptr(), rewards.data_ptr(), hit.data_ptr(), any_hit.data_ptr(), (int)h.size(0), c10::cuda::getCurrentCUDAStream()); } void init_rng_out(torch::Tensor rng, int64_t off) { launch_init_rng(static_cast(rng.data_ptr()), (int)rng.size(0), (long long)off, c10::cuda::getCurrentCUDAStream()); } void epi_out(torch::Tensor gates, torch::Tensor h_in, torch::Tensor state, torch::Tensor h_out, torch::Tensor h16) { launch_epi(reinterpret_cast(gates.data_ptr()), h_in.data_ptr(), state.data_ptr(), h_out.data_ptr(), h16.defined() ? reinterpret_cast<__half*>(h16.data_ptr()) : nullptr, (int)h_in.size(0), c10::cuda::getCurrentCUDAStream()); } void encode_out(torch::Tensor agent, torch::Tensor food, torch::Tensor w, torch::Tensor b, torch::Tensor h, torch::Tensor h16) { launch_encode(agent.data_ptr(), food.data_ptr(), w.data_ptr(), b.data_ptr(), h.data_ptr(), reinterpret_cast<__half*>(h16.data_ptr()), (int)agent.size(0), c10::cuda::getCurrentCUDAStream()); } void action_out(torch::Tensor h, torch::Tensor w, torch::Tensor b, torch::Tensor logits, torch::Tensor agent, torch::Tensor food, torch::Tensor rewards, torch::Tensor hit, torch::Tensor any_hit) { launch_action(h.data_ptr(), w.data_ptr(), b.data_ptr(), logits.data_ptr(), agent.data_ptr(), food.data_ptr(), rewards.data_ptr(), hit.data_ptr(), any_hit.data_ptr(), (int)h.size(0), c10::cuda::getCurrentCUDAStream()); } void food_out(torch::Tensor food, torch::Tensor rng, torch::Tensor hit, torch::Tensor any_hit) { launch_food(food.data_ptr(), static_cast(rng.data_ptr()), hit.data_ptr(), any_hit.data_ptr(), (int)food.size(0), c10::cuda::getCurrentCUDAStream()); } void env_step_out(torch::Tensor agent, torch::Tensor food, torch::Tensor actions, torch::Tensor rng, torch::Tensor agent_out, torch::Tensor food_out, torch::Tensor reward, torch::Tensor rng_out, torch::Tensor hit, torch::Tensor any_hit) { launch_env_step(agent.data_ptr(), food.data_ptr(), static_cast(actions.data_ptr()), static_cast(rng.data_ptr()), agent_out.data_ptr(), food_out.data_ptr(), reward.data_ptr(), static_cast(rng_out.data_ptr()), hit.data_ptr(), any_hit.data_ptr(), (int)agent.size(0), c10::cuda::getCurrentCUDAStream()); } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("gru_fast", &gru_fast_out, "fused gru"); m.def("epi", &epi_out, "epi"); m.def("epi_h16", &epi_h16_out, "epi h16"); m.def("action_h16", &action_h16_out, "action h16"); m.def("init_rng", &init_rng_out, "init rng"); m.def("encode", &encode_out, "encode"); m.def("action", &action_out, "action"); m.def("food", &food_out, "food"); m.def("env_step", &env_step_out, "env_step"); m.def("smem_fast", &smem_fast, "smem"); }