"""Megakernel PPO training for a vectorized grid-foraging task. The whole training run -- from-scratch policy init, per-iteration rollout (all 32 env steps x 4096 envs + policy forward + sampling + GAE), and the PPO update (4 epochs x 4 minibatches of the clipped loss + grad-norm clip + Adam) -- executes inside ONE persistent cooperative kernel launch. Between phases the kernel synchronizes the grid with cooperative groups; the host only launches once and reads back the per-iteration mean episodic returns. Everything (env, policy forward/backward, GAE, optimizer) is hand-written CUDA -- no RL library is involved. Kernel structure per iteration (17 grid syncs): * rollout phase: warp-per-env; 32 sequential env steps, each computing the policy forward for its env (lanes split the 64 hidden units), categorical sampling, env transition, reward, with values/rewards kept in lane registers by step index; the GAE back-sweep runs in-warp with shuffle broadcasts; per-warp stat sums are block-reduced into global doubles. * 16 x (gradient accumulation phase, grid sync, post phase). The post phase (grad-norm clip + Adam) runs redundantly in every block against shared memory state (bitwise identical), so minibatch k+1 follows minibatch k with NO intervening grid sync. Gradient buffers rotate mod 4 and are zeroed two minibatches ahead, and the advantage stats buffer ping-pongs by iteration parity so no trailing iteration-boundary sync is needed. Parameter layout (645 floats, matching the reference MLP Linear layout): W1 [0..256) 64x4 (out,in) | b1 [256..320) | W2 [320..576) 4x64 (pi) | b2 [576..580) | wv [580..644) (value) | bv [644]. Weights ~ U(-bound, bound), bound = 1/sqrt(fan_in); biases zero. RNG: counter-based Philox4x32-10 keyed by the run seed; disjoint counter tags for env resets, per-step sampling/respawn, weight init, and epoch permutation. Minibatch permutation per epoch is an unbiased bijection over Z_2^17, q = (a*s + b) mod 2^17 (odd a, offset b redrawn per epoch), which partitions the batch into four equal minibatches exactly like a reshuffle. """ import torch from torch.utils.cpp_extension import load_inline _EXT = None _CUDA_SRC = r""" #include #include #include #include namespace cg = cooperative_groups; #define NUM_ENVS 4096 #define HORIZON 32 #define NSAMP 131072 #define MB 32768 #define NPARAM 645 #define NTHREADS 128 #define NWARP (NTHREADS / 32) #define GAMMA 0.99f #define GAE_LAM 0.95f #define CLIP 0.2f #define ENT_COEF 0.01f #define VF_COEF 0.5f #define LR 3.0e-3f #define MAX_GRAD_NORM 0.5f __device__ __forceinline__ uint4 philox_round(uint4 c, uint2 k) { unsigned lo0 = 0xD2511F53u * c.x, hi0 = __umulhi(0xD2511F53u, c.x); unsigned lo1 = 0xCD9E8D57u * c.z, hi1 = __umulhi(0xCD9E8D57u, c.z); return make_uint4(hi1 ^ c.y ^ k.x, lo1, hi0 ^ c.w ^ k.y, lo0); } __device__ __forceinline__ uint4 philox4x32_10(uint4 counter, uint2 key) { #pragma unroll for (int i = 0; i < 10; ++i) { counter = philox_round(counter, key); key.x += 0x9E3779B9u; key.y += 0xBB67AE85u; } return counter; } __device__ __forceinline__ unsigned xnext_(uint4& st) { unsigned r = __funnelshift_l(st.x + st.w, st.x + st.w, 7) + st.x; unsigned t = st.y << 9; st.z ^= st.x; st.w ^= st.y; st.y ^= st.z; st.x ^= st.w; st.z ^= t; st.w = __funnelshift_l(st.w, st.w, 11); return r; } __device__ __forceinline__ float u01(unsigned u) { return (float)(u >> 8) * (1.0f / 16777216.0f); } // policy forward for one sample; lanes j1=lane, j2=lane+32 own hidden units. // returns h (this lane's two units), plus z0..z3, v via warp reduction. __device__ __forceinline__ void fwd( float4 x, const float* __restrict__ th, float& h1, float& h2, float& z0, float& z1, float& z2, float& z3, float& v, int j1, int j2) { h1 = th[256 + j1]; h1 = fmaf(th[j1 * 4 + 0], x.x, h1); h1 = fmaf(th[j1 * 4 + 1], x.y, h1); h1 = fmaf(th[j1 * 4 + 2], x.z, h1); h1 = fmaf(th[j1 * 4 + 3], x.w, h1); h1 = tanhf(h1); h2 = th[256 + j2]; h2 = fmaf(th[j2 * 4 + 0], x.x, h2); h2 = fmaf(th[j2 * 4 + 1], x.y, h2); h2 = fmaf(th[j2 * 4 + 2], x.z, h2); h2 = fmaf(th[j2 * 4 + 3], x.w, h2); h2 = tanhf(h2); float zp0 = th[320 + j1] * h1 + th[320 + j2] * h2; float zp1 = th[384 + j1] * h1 + th[384 + j2] * h2; float zp2 = th[448 + j1] * h1 + th[448 + j2] * h2; float zp3 = th[512 + j1] * h1 + th[512 + j2] * h2; float vp = th[580 + j1] * h1 + th[580 + j2] * h2; #pragma unroll for (int off = 16; off; off >>= 1) { zp0 += __shfl_xor_sync(0xffffffffu, zp0, off); zp1 += __shfl_xor_sync(0xffffffffu, zp1, off); zp2 += __shfl_xor_sync(0xffffffffu, zp2, off); zp3 += __shfl_xor_sync(0xffffffffu, zp3, off); vp += __shfl_xor_sync(0xffffffffu, vp, off); } z0 = zp0 + th[576]; z1 = zp1 + th[577]; z2 = zp2 + th[578]; z3 = zp3 + th[579]; v = vp + th[644]; } __global__ void __launch_bounds__(NTHREADS) train_kernel( float4* __restrict__ rec, // NSAMP x {obs, [logp, actf, adv, ret]} pairs float* __restrict__ grad, float* __restrict__ theta_g, float* __restrict__ curve, double* __restrict__ dstats, // 2 parity x 2 unsigned seed_lo, unsigned seed_hi, int iters) { cg::grid_group grid = cg::this_grid(); // dynamic shared memory carve-out (static limit is 48KB; we need ~49.3KB) extern __shared__ float smem[]; float* s_w2t = smem; // 256 (16B-aligned: holds float4 rows) float* s_theta = s_w2t + 256; // NPARAM (16B-aligned: holds float4 rows) float* s_m = s_theta + NPARAM; // NPARAM float* s_v = s_m + NPARAM; // NPARAM float* s_hbase = s_v + NPARAM + 1; // NWARP * (32*65) (16B-aligned) float* s_dzbase = s_hbase + NWARP * (32 * 65); float* s_xbase = s_dzbase + NWARP * (32 * 8); float* s_red = s_xbase + NWARP * (32 * 4); float* s_scalar = s_red + NWARP * 4; const int tid = threadIdx.x; const int gtid = blockIdx.x * NTHREADS + tid; const int nthr = gridDim.x * NTHREADS; const int nwarps = gridDim.x * NWARP; const int wid = blockIdx.x * NWARP + (tid >> 5); const int lane = tid & 31; const int wib = tid >> 5; const uint2 key = make_uint2(seed_lo, seed_hi); // --- INIT --------------------------------------------------------------- for (int p = gtid; p < NPARAM; p += nthr) { float val = 0.0f; if (p < 256) { float u = u01(philox4x32_10(make_uint4(p, 0xC0DE0001u, 0, 0), key).x); val = (2.0f * u - 1.0f) * 0.5f; } else if ((p >= 320 && p < 576) || (p >= 580 && p < 644)) { float u = u01(philox4x32_10(make_uint4(p, 0xC0DE0001u, 0, 0), key).x); val = (2.0f * u - 1.0f) * 0.125f; } theta_g[p] = val; } for (int p = gtid; p < 4 * NPARAM; p += nthr) grad[p] = 0.0f; if (gtid < 4) dstats[gtid] = 0.0; for (int i = gtid; i < iters; i += nthr) curve[i] = 0.0f; grid.sync(); for (int p = tid; p < NPARAM; p += NTHREADS) { s_theta[p] = theta_g[p]; s_m[p] = 0.0f; s_v[p] = 0.0f; } for (int p = tid; p < 256; p += NTHREADS) s_w2t[p] = s_theta[320 + (p & 3) * 64 + (p >> 2)]; grid.sync(); const int j1 = lane; const int j2 = lane + 32; double pow1 = 0.9, pow2 = 0.999; // beta^t trackers for Adam bias correction for (int iter = 0; iter < iters; ++iter) { const int par = iter & 1; // --- rollout + GAE + stats (warp per env) ------------------------------ float accS1 = 0.0f, accS2 = 0.0f, accRew = 0.0f; for (int ee = wid * 4; ee < NUM_ENVS; ee += nwarps * 4) { int ev[4]; uint4 rr[4]; #pragma unroll for (int i = 0; i < 4; ++i) { ev[i] = ee + i; rr[i] = philox4x32_10(make_uint4(ev[i] < NUM_ENVS ? ev[i] : 0, 0xA000u + (unsigned)iter, 0, 0), key); } int ax[4], ay[4], fx[4], fy[4]; #pragma unroll for (int i = 0; i < 4; ++i) { ax[i] = (int)(u01(rr[i].x) * 11.0f); ay[i] = (int)(u01(rr[i].y) * 11.0f); fx[i] = (int)(u01(rr[i].z) * 11.0f); fy[i] = (int)(u01(rr[i].w) * 11.0f); } uint4 rng[4]; #pragma unroll for (int i = 0; i < 4; ++i) rng[i] = make_uint4(rr[i].x | 1u, rr[i].y ^ 0x9E3779B9u, rr[i].z + 0x6A09E667u, rr[i].w | 0x3C6EF372u); float valA[4], rewA[4]; #pragma unroll for (int i = 0; i < 4; ++i) { valA[i] = 0.0f; rewA[i] = 0.0f; } for (int t = 0; t < HORIZON; ++t) { float zz[4][4]; float vv[4], lpv[4]; int aa[4]; #pragma unroll for (int i = 0; i < 4; ++i) { float h1, h2, z0, z1, z2, z3, v; fwd(make_float4( (float)(fx[i] - ax[i]) * (1.0f / 11.0f), (float)(fy[i] - ay[i]) * (1.0f / 11.0f), (float)ax[i] * 0.1f, (float)ay[i] * 0.1f), s_theta, h1, h2, z0, z1, z2, z3, v, j1, j2); float m = fmaxf(fmaxf(z0, z1), fmaxf(z2, z3)); float e0 = __expf(z0 - m), e1 = __expf(z1 - m), e2 = __expf(z2 - m), e3 = __expf(z3 - m); float se = e0 + e1 + e2 + e3; float u = u01(xnext_(rng[i])) * se; int a = 0; float cum = e0; if (u > cum) { cum += e1; a = 1; } if (u > cum) { cum += e2; a = 2; } if (u > cum) { a = 3; } zz[i][0] = z0; zz[i][1] = z1; zz[i][2] = z2; zz[i][3] = z3; vv[i] = v; aa[i] = a; float za = (a == 0) ? z0 : (a == 1) ? z1 : (a == 2) ? z2 : z3; lpv[i] = za - m - __logf(se); } if (lane == 0) { #pragma unroll for (int i = 0; i < 4; ++i) { if (ev[i] < NUM_ENVS) { int q = t * NUM_ENVS + ev[i]; float4 m4; rec[q * 2] = make_float4( (float)(fx[i] - ax[i]) * (1.0f / 11.0f), (float)(fy[i] - ay[i]) * (1.0f / 11.0f), (float)ax[i] * 0.1f, (float)ay[i] * 0.1f); m4.x = lpv[i]; m4.y = (float)aa[i]; rec[q * 2 + 1] = m4; } } } #pragma unroll for (int i = 0; i < 4; ++i) { int nx = ax[i] + (aa[i] == 3) - (aa[i] == 2); int ny = ay[i] + (aa[i] == 1) - (aa[i] == 0); nx = min(max(nx, 0), 10); ny = min(max(ny, 0), 10); float rew = 0.0f; if (nx == fx[i] && ny == fy[i]) { rew = 1.0f; fx[i] = (int)(u01(xnext_(rng[i])) * 11.0f); fy[i] = (int)(u01(xnext_(rng[i])) * 11.0f); } ax[i] = nx; ay[i] = ny; rewA[i] = lane == t ? rew : rewA[i]; valA[i] = lane == t ? vv[i] : valA[i]; } } #pragma unroll for (int i = 0; i < 4; ++i) { if (ev[i] >= NUM_ENVS) break; float gae = 0.0f, adv_t = 0.0f, ret_t = 0.0f; for (int t = HORIZON - 1; t >= 0; --t) { float ntf = (t < HORIZON - 1) ? 1.0f : 0.0f; float vt = __shfl_sync(0xffffffffu, valA[i], t); float rt = __shfl_sync(0xffffffffu, rewA[i], t); float vt1 = __shfl_sync(0xffffffffu, valA[i], (t < HORIZON - 1) ? t + 1 : t); float delta = rt + GAMMA * vt1 * ntf - vt; gae = delta + GAMMA * GAE_LAM * ntf * gae; if (lane == t) { adv_t = gae; ret_t = gae + vt; } } int q = lane * NUM_ENVS + ev[i]; float4 m4 = rec[q * 2 + 1]; m4.z = adv_t; m4.w = ret_t; rec[q * 2 + 1] = m4; accS1 += adv_t; accS2 += adv_t * adv_t; accRew += rewA[i]; } } // block-reduce the three accumulators #pragma unroll for (int off = 16; off; off >>= 1) { accS1 += __shfl_down_sync(0xffffffffu, accS1, off); accS2 += __shfl_down_sync(0xffffffffu, accS2, off); accRew += __shfl_down_sync(0xffffffffu, accRew, off); } if (lane == 0) { s_red[wib * 4 + 0] = accS1; s_red[wib * 4 + 1] = accS2; s_red[wib * 4 + 2] = accRew; } __syncthreads(); if (tid < 3) { double acc = 0.0; for (int w = 0; w < NWARP; ++w) acc += (double)s_red[w * 4 + tid]; if (tid < 2) atomicAdd(&dstats[par * 2 + tid], acc); else atomicAdd(&curve[iter], (float)acc); } grid.sync(); // --- PPO update --------------------------------------------------------- double* ds = dstats + par * 2; if (tid < 2) s_scalar[tid] = (float)ds[tid]; __syncthreads(); float mean = s_scalar[0] * (1.0f / NSAMP); float var = (float)(((double)s_scalar[1] - (double)s_scalar[0] * (double)mean) / (NSAMP - 1.0)); float invstd = 1.0f / (sqrtf(var) + 1e-8f); for (int epoch = 0; epoch < 4; ++epoch) { uint4 pr = philox4x32_10(make_uint4((unsigned)(iter * 4 + epoch), 0xD000u, 0, 0), key); const unsigned aperm = ((pr.x & 0xFFFFu) << 1) | 1u; const unsigned bperm = pr.y & 0x1FFFFu; for (int k = 0; k < 4; ++k) { const int gmb = epoch * 4 + k; const int gb = gmb & 3; float gW1a[4] = {0.f, 0.f, 0.f, 0.f}, gW1b[4] = {0.f, 0.f, 0.f, 0.f}; float gb1a = 0.f, gb1b = 0.f; float gW2a[4] = {0.f, 0.f, 0.f, 0.f}, gW2b[4] = {0.f, 0.f, 0.f, 0.f}; float gb2 = 0.f, gwva = 0.f, gwvb = 0.f, gbv = 0.f; float* hT = s_hbase + wib * (32 * 65); float* dzT = s_dzbase + wib * (32 * 8); float* xT = s_xbase + wib * (32 * 4); const float4* w1q = reinterpret_cast(s_theta); const float4* w2q = reinterpret_cast(s_w2t); const int mbstart = k * MB; // tile = 32 samples per warp iteration; thread-per-sample inside for (int slot0 = mbstart + wid * 32; slot0 < mbstart + MB; slot0 += nwarps * 32) { const int slot = slot0 + lane; const unsigned q = (aperm * (unsigned)slot + bperm) & 0x1FFFFu; float4 x = rec[q * 2]; float4 m4 = rec[q * 2 + 1]; const int a = (int)m4.y; const float lpo = m4.x; const float ad = (m4.z - mean) * invstd; const float rt = m4.w; xT[lane * 4 + 0] = x.x; xT[lane * 4 + 1] = x.y; xT[lane * 4 + 2] = x.z; xT[lane * 4 + 3] = x.w; float z0 = s_theta[576], z1 = s_theta[577], z2 = s_theta[578], z3 = s_theta[579]; float v = s_theta[644]; #pragma unroll 4 for (int j = 0; j < 64; ++j) { float4 w1 = w1q[j]; float acc = s_theta[256 + j]; acc = fmaf(w1.x, x.x, acc); acc = fmaf(w1.y, x.y, acc); acc = fmaf(w1.z, x.z, acc); acc = fmaf(w1.w, x.w, acc); float hj = tanhf(acc); hT[lane * 65 + j] = hj; float4 w2 = w2q[j]; z0 = fmaf(w2.x, hj, z0); z1 = fmaf(w2.y, hj, z1); z2 = fmaf(w2.z, hj, z2); z3 = fmaf(w2.w, hj, z3); v = fmaf(s_theta[580 + j], hj, v); } float m = fmaxf(fmaxf(z0, z1), fmaxf(z2, z3)); float e0 = __expf(z0 - m), e1 = __expf(z1 - m), e2 = __expf(z2 - m), e3 = __expf(z3 - m); float se = e0 + e1 + e2 + e3; float inv = 1.0f / se; float p0 = e0 * inv, p1 = e1 * inv, p2 = e2 * inv, p3 = e3 * inv; float lg = __logf(se); float lp_new = ((a == 0) ? z0 : (a == 1) ? z1 : (a == 2) ? z2 : z3) - m - lg; float r = __expf(lp_new - lpo); bool active = (ad > 0.f) ? (r < (1.0f + CLIP)) : (r > (1.0f - CLIP)); float g = active ? -ad * r : 0.0f; float ls0 = z0 - m - lg, ls1 = z1 - m - lg, ls2 = z2 - m - lg, ls3 = z3 - m - lg; float H = -(p0 * ls0 + p1 * ls1 + p2 * ls2 + p3 * ls3); const float sc = 1.0f / (float)MB; float dz0 = sc * (g * ((a == 0) - p0) + ENT_COEF * p0 * (ls0 + H)); float dz1 = sc * (g * ((a == 1) - p1) + ENT_COEF * p1 * (ls1 + H)); float dz2 = sc * (g * ((a == 2) - p2) + ENT_COEF * p2 * (ls2 + H)); float dz3 = sc * (g * ((a == 3) - p3) + ENT_COEF * p3 * (ls3 + H)); float dv = sc * VF_COEF * (v - rt); dzT[lane * 8 + 0] = dz0; dzT[lane * 8 + 1] = dz1; dzT[lane * 8 + 2] = dz2; dzT[lane * 8 + 3] = dz3; dzT[lane * 8 + 4] = dv; __syncwarp(); #pragma unroll for (int t = 0; t < 32; ++t) { float h1v = hT[t * 65 + j1]; float h2v = hT[t * 65 + j2]; float4 dz4 = ((const float4*)dzT)[t * 2]; float d0 = dz4.x, d1 = dz4.y, d2 = dz4.z, d3 = dz4.w; float dvv = dzT[t * 8 + 4]; gW2a[0] = fmaf(d0, h1v, gW2a[0]); gW2a[1] = fmaf(d1, h1v, gW2a[1]); gW2a[2] = fmaf(d2, h1v, gW2a[2]); gW2a[3] = fmaf(d3, h1v, gW2a[3]); gW2b[0] = fmaf(d0, h2v, gW2b[0]); gW2b[1] = fmaf(d1, h2v, gW2b[1]); gW2b[2] = fmaf(d2, h2v, gW2b[2]); gW2b[3] = fmaf(d3, h2v, gW2b[3]); if (lane < 4) gb2 += dzT[t * 8 + lane]; gwva = fmaf(dvv, h1v, gwva); gwvb = fmaf(dvv, h2v, gwvb); if (lane == 0) gbv += dvv; } #pragma unroll 4 for (int j = 0; j < 64; ++j) { float hj = hT[lane * 65 + j]; float4 w2 = w2q[j]; float dh = w2.x * dz0 + w2.y * dz1 + w2.z * dz2 + w2.w * dz3 + s_theta[580 + j] * dv; hT[lane * 65 + j] = dh * (1.0f - hj * hj); } __syncwarp(); #pragma unroll for (int t = 0; t < 32; ++t) { float d1v = hT[t * 65 + j1]; float d2v = hT[t * 65 + j2]; float4 xv = ((const float4*)xT)[t]; gW1a[0] = fmaf(d1v, xv.x, gW1a[0]); gW1a[1] = fmaf(d1v, xv.y, gW1a[1]); gW1a[2] = fmaf(d1v, xv.z, gW1a[2]); gW1a[3] = fmaf(d1v, xv.w, gW1a[3]); gW1b[0] = fmaf(d2v, xv.x, gW1b[0]); gW1b[1] = fmaf(d2v, xv.y, gW1b[1]); gW1b[2] = fmaf(d2v, xv.z, gW1b[2]); gW1b[3] = fmaf(d2v, xv.w, gW1b[3]); gb1a += d1v; gb1b += d2v; } __syncwarp(); } { // reuse the (now consumed) h tile as this warp's private grad staging: // plain smem stores, no atomic spin, then block-reduce the 4 copies float* wg = hT; #pragma unroll for (int i = 0; i < 21; ++i) wg[lane + i * 32] = 0.0f; __syncwarp(); wg[j1 * 4 + 0] = gW1a[0]; wg[j1 * 4 + 1] = gW1a[1]; wg[j1 * 4 + 2] = gW1a[2]; wg[j1 * 4 + 3] = gW1a[3]; wg[j2 * 4 + 0] = gW1b[0]; wg[j2 * 4 + 1] = gW1b[1]; wg[j2 * 4 + 2] = gW1b[2]; wg[j2 * 4 + 3] = gW1b[3]; wg[256 + j1] = gb1a; wg[256 + j2] = gb1b; wg[320 + j1] = gW2a[0]; wg[384 + j1] = gW2a[1]; wg[448 + j1] = gW2a[2]; wg[512 + j1] = gW2a[3]; wg[320 + j2] = gW2b[0]; wg[384 + j2] = gW2b[1]; wg[448 + j2] = gW2b[2]; wg[512 + j2] = gW2b[3]; wg[580 + j1] = gwva; wg[580 + j2] = gwvb; __syncwarp(); if (lane < 4) wg[576 + lane] = gb2; if (lane == 0) wg[644] = gbv; } __syncthreads(); { #pragma unroll for (int i = 0; i < 6; ++i) { int p = tid + i * NTHREADS; if (p < NPARAM) { float acc = s_hbase[0 * (32 * 65) + p]; #pragma unroll for (int w = 1; w < NWARP; ++w) acc += s_hbase[w * (32 * 65) + p]; atomicAdd(&grad[gb * NPARAM + p], acc); } } } grid.sync(); // ---- post: redundant clipped-Adam in every block ---- float mygs[6]; float ssq = 0.0f; #pragma unroll for (int i = 0; i < 6; ++i) { int p = tid + i * NTHREADS; if (p < NPARAM) { mygs[i] = grad[gb * NPARAM + p]; ssq = fmaf(mygs[i], mygs[i], ssq); } } s_red[wib * 4 + 3] = ssq; __syncthreads(); if (tid == 0) { float acc = 0.0f; for (int w = 0; w < NWARP; ++w) acc += s_red[w * 4 + 3]; s_scalar[2] = acc; } __syncthreads(); float total_norm = sqrtf(s_scalar[2]); float clip_coef = MAX_GRAD_NORM / (total_norm + 1e-6f); float cscale = clip_coef < 1.0f ? clip_coef : 1.0f; float bc1 = (float)(1.0 - pow1); float inv_bc2_sqrt = (float)(1.0 / sqrt(pow2 > 0.0 ? (1.0 - pow2) : 1.0)); float lr_t = LR / bc1; #pragma unroll for (int i = 0; i < 6; ++i) { int p = tid + i * NTHREADS; if (p < NPARAM) { float gg = mygs[i] * cscale; float mm = 0.9f * s_m[p] + 0.1f * gg; float vv = 0.999f * s_v[p] + 0.001f * gg * gg; s_m[p] = mm; s_v[p] = vv; float denom = sqrtf(vv) * inv_bc2_sqrt + 1e-8f; s_theta[p] -= lr_t * mm / denom; if (blockIdx.x == 0) theta_g[p] = s_theta[p]; } } pow1 *= 0.9; pow2 *= 0.999; __syncthreads(); for (int p = tid; p < 256; p += NTHREADS) s_w2t[p] = s_theta[320 + (p & 3) * 64 + (p >> 2)]; const int zb = (gmb + 2) & 3; #pragma unroll for (int i = 0; i < 6; ++i) { int p = tid + i * NTHREADS; if (p < NPARAM) grad[zb * NPARAM + p] = 0.0f; } if (gmb == 15 && gtid < 2) ds[gtid] = 0.0; } } } } void train_run(at::Tensor rec, at::Tensor grad, at::Tensor theta, at::Tensor curve, at::Tensor dstats, int64_t seed, int64_t iters) { const int SMEM_FLOATS = 3 * NPARAM + 256 + 1 + NWARP * (32 * 65) + NWARP * (32 * 8) + NWARP * (32 * 4) + NWARP * 4 + 8; const int SMEM_BYTES = SMEM_FLOATS * 4; static int num_blocks = -1; if (num_blocks < 0) { int dev; cudaGetDevice(&dev); cudaDeviceProp prop; cudaGetDeviceProperties(&prop, dev); cudaError_t err = cudaFuncSetAttribute(train_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES); TORCH_CHECK(err == cudaSuccess, "smem attr: ", cudaGetErrorString(err)); int bpm = 0; err = cudaOccupancyMaxActiveBlocksPerMultiprocessor( &bpm, train_kernel, NTHREADS, SMEM_BYTES); TORCH_CHECK(err == cudaSuccess, "occupancy: ", cudaGetErrorString(err)); TORCH_CHECK(bpm >= 1, "kernel not co-resident"); num_blocks = prop.multiProcessorCount * bpm; } float4* rec_p = reinterpret_cast(rec.data_ptr()); float* grad_p = grad.data_ptr(); float* theta_p = theta.data_ptr(); float* curve_p = curve.data_ptr(); double* dstats_p = dstats.data_ptr(); unsigned seed_lo = (unsigned)(seed & 0xffffffff); unsigned seed_hi = (unsigned)((seed >> 32) & 0xffffffff) ^ 0x5EED5EEDu; int iters_i = (int)iters; void* args[] = {&rec_p, &grad_p, &theta_p, &curve_p, &dstats_p, &seed_lo, &seed_hi, &iters_i}; cudaError_t err = cudaLaunchCooperativeKernel( (void*)train_kernel, dim3(num_blocks), dim3(NTHREADS), args, SMEM_BYTES, at::cuda::getCurrentCUDAStream()); TORCH_CHECK(err == cudaSuccess, "coop launch: ", cudaGetErrorString(err)); } """ def _get_ext(): global _EXT if _EXT is None: _EXT = load_inline( name="grid_ppo_megakernel_v2", cpp_sources=["void train_run(at::Tensor rec, at::Tensor grad, at::Tensor theta, at::Tensor curve, at::Tensor dstats, int64_t seed, int64_t iters);"], cuda_sources=[_CUDA_SRC], functions=["train_run"], extra_cuda_cflags=[ "-O3", "--use_fast_math", "-gencode=arch=compute_120,code=sm_120", ], verbose=False, ) return _EXT _BUFS = {} def _get_bufs(): if _BUFS: return _BUFS dev = torch.device("cuda:0") nsamp = 32 * 4096 _BUFS["rec"] = torch.zeros(nsamp * 8, device=dev, dtype=torch.float32) _BUFS["grad"] = torch.zeros(4 * 645, device=dev, dtype=torch.float32) _BUFS["theta"] = torch.zeros(645, device=dev, dtype=torch.float32) _BUFS["curve"] = torch.zeros(1024, device=dev, dtype=torch.float32) _BUFS["dstats"] = torch.zeros(4, device=dev, dtype=torch.float64) return _BUFS def train(total_env_steps: int, seed: int) -> list[float]: ext = _get_ext() b = _get_bufs() rollout, num_envs = 32, 4096 iters = max(1, total_env_steps // (rollout * num_envs)) ext.train_run( b["rec"], b["grad"], b["theta"], b["curve"], b["dstats"], int(seed), iters, ) torch.cuda.synchronize() c = b["curve"][:iters].cpu() return (c / num_envs).tolist()