KernelBench mega · H100
rl grid ppo Kimi K3 (256k)
manually audited: clean
H100 kinetic-0715 RL grid PPO MEGA cell, graded at 5.0604x. Manual local audit read all 650 lines of solution.py, the reference/check/benchmark contract, check.log, benchmark.log, result provenance, and the complete transcript evidence relevant to contamination and grader mutation. The timed path is a real fused from-scratch trainer, not a returned constant or cached curve: train() allocates fresh rollout, metadata, curve, statistics, gradient-slot, minibatch-list, and count tensors on every invocation, then makes one cudaLaunchCooperativeKernel call. Inside that launch, the kernel initializes all 645 actor-critic parameters and Adam state from the live seed, simulates the 11x11 grid-foraging MDP for all 4096 environments and 32 rollout steps, samples actions, respawns food, accumulates rewards, computes GAE and normalized advantages, and executes all 4 epochs x 4 PPO minibatch updates with clipped policy loss, value loss, entropy term, global grad-norm clipping, and Adam. Grid-wide barriers separate the phases. There is no CUDA graph, torch.compile, per-step host launch loop, prebuilt RL dependency, reference or baseline import, output memoization, data_ptr/input-identity lookup, or sleep. The only process cache is _EXT for the compiled extension and choose_nt's cached occupancy choice; neither stores outputs or training state. Every train() call uses fresh buffers and recomputes the curve from live seed-driven environment and policy state. The official check log shows solution finals 3.964/3.966/3.960 versus reference 3.956/3.963/4.000, early returns 0.226/0.233/0.215, then PASS. The official benchmark's three fresh random seeds all learn to 3.956-3.971 at 126.29-126.51M env-steps/s; the logged best 126,511,028 divided by the 25,000,000 anchor is exactly 5.0604, so the score's magnitude is internally consistent. No GPU rerun was performed for this audit; code structure makes a stale-result path impossible.
Kernel source (redacted)
"""Grid-foraging PPO training megakernel.
The entire PPO training run -- rollout (all HORIZON steps x all envs, policy
forward, action sampling, env step, reward, GAE), advantage normalization, and
the full PPO update (4 epochs x 4 minibatches of backprop + grad-norm clip +
Adam) -- executes inside ONE persistent cooperative CUDA kernel launch per
training run. Grid-wide phase synchronization uses cooperative-groups
grid.sync(); kernel-launch count is O(1) per run and never scales with env
steps, horizon, epochs, or minibatches. No CUDA graphs, no torch.compile, no
prebuilt RL libraries: the environment, policy, and PPO update are
hand-written fused CUDA. There is no Python-side per-step or per-minibatch
kernel loop anywhere: train() issues exactly one cooperative launch and one
device->host copy of the 40-iteration return curve.
Parallel decomposition (the kernel body is a template instantiated at both
NT=512 and NT=256; the host picks the fattest variant that actually fits with
a cooperative co-resident grid on the current arch):
* NT threads (NT/32 warps) per block, one block per SM resident (grid capped
by co-residency so the cooperative launch is always legal).
* Rollout: each warp owns whole environments (balanced contiguous ranges).
The 32 steps of an env run lane-parallel inside the warp (lane l computes
hidden units l and l+32; a butterfly shuffle reduction yields the 4 actor
logits + value), while every lane redundantly - hence bit-identically -
tracks the env state and PCG RNG. GAE is a per-warp backward recurrence
over a shared-memory reward/value stash.
* Advantage mean/std: per-warp partials -> block -> global double atomics;
normalization is applied on the fly when the update consumes samples.
* PPO update: 16 minibatch barriers per iteration. For perfect balance, each
iteration first scatters rollout slots into 16 per-(epoch,minibatch) index
lists (membership = hash of the slot index; uniform partition, each sample
used exactly once per epoch), built between rollout and the update while
rollout stragglers finish. Warps then take contiguous, count-balanced list
slices and process samples warp-collectively two-at-a-time with software
pipelined gathers. Analytic PPO gradients (exactly matching autograd of
loss = pg + 0.5*vloss - 0.01*entropy) accumulate in per-lane registers
(each lane owns its own hidden units' gradient slots), then per-block
shared-memory sums, then pre-zeroed global atomic slots per minibatch.
After grid.sync every block performs the identical grad-norm clip + Adam
update on its own shared-memory copy of the 645 parameters - redundantly
deterministic, so all replicas stay bitwise identical with no weight
broadcast.
"""
from __future__ import annotations
import os
import torch
OP_TYPE = "rl_grid_ppo"
HARDWARE_REQUIRED = ["RTX_PRO_6000"]
# --- Task constants (must match reference.py exactly) ------------------------
GRID = 11
NUM_ENVS = 4096
HORIZON = 32
ROLLOUT = 32
OBS_DIM = 4
N_ACT = 4
HIDDEN = 64
GAMMA = 0.99
LAM = 0.95
CLIP = 0.2
EPOCHS = 4
MINIBATCHES = 4
LR = 3.0e-3
ENT_COEF = 0.01
VF_COEF = 0.5
MAX_GRAD_NORM = 0.5
N_SAMPLES = ROLLOUT * NUM_ENVS
N_MINIBATCHES = EPOCHS * MINIBATCHES
NPARAM = 645 # 4*64 + 64 + 4*64 + 4 + 64 + 1
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cooperative_groups.h>
#include <math.h>
#include <stdint.h>
namespace cg = cooperative_groups;
#define NUM_ENVS 4096
#define T_STEPS 32
#define N_SAMPLES 131072
#define NPARAM 645
#define N_MB 16
#define INV_MB (1.0f / 32768.0f)
#define GAMMA 0.99f
#define GL (0.99f * 0.95f)
#define LRF 3.0e-3f
#define ENT_COEF 0.01f
#define MAXN 0.5f
#define A_EPS 1e-8f
#define BETA1 0.9f
#define BETA2 0.999f
#define INV_N (1.0 / 131072.0)
// Parameter layout (645 floats, shared-memory master copy per block):
// [0,256) W1T : k*64 + j (obs k -> hidden j), init U(-1/2,1/2)
// [256,320) B1 : j zero
// [320,576) W2 : m*64 + j (hidden j -> logit m), init U(-1/8,1/8)
// [576,580) B2 : m zero
// [580,644) W3 : j (hidden j -> value), init U(-1/8,1/8)
// [644] B3 zero
template<int NWARP>
struct SharedMem {
float w[NPARAM];
float am[NPARAM];
float av[NPARAM];
float gacc[NPARAM];
float rew_stash[NWARP][T_STEPS];
float val_stash[NWARP][T_STEPS];
double stats_blk[2];
float curve_blk;
float red[NWARP];
};
__device__ __forceinline__ uint64_t splitmix64(uint64_t x) {
x += 0x9E3779B97F4A7C15ULL;
x = (x ^ (x >> 30)) * 0xBF58476D1CE4E5B9ULL;
x = (x ^ (x >> 27)) * 0x94D049BB133111EBULL;
return x ^ (x >> 31);
}
__device__ __forceinline__ float u01(uint64_t h) {
return (float)(h >> 40) * (1.0f / 16777216.0f);
}
// PCG32 (XSH-RR) per-env RNG; advanced redundantly & identically by all lanes.
struct Rng { uint64_t s; };
__device__ __forceinline__ uint32_t rng_next(Rng* r) {
uint64_t old = r->s;
r->s = old * 6364136223846793005ULL + 1442695040888963407ULL;
uint32_t x = (uint32_t)(((old >> 18) ^ old) >> 27);
uint32_t rot = (uint32_t)(old >> 59);
return (x >> rot) | (x << ((32u - rot) & 31u));
}
__device__ __forceinline__ float rng_u01(Rng* r) {
return (float)(rng_next(r) >> 8) * (1.0f / 16777216.0f);
}
__device__ __forceinline__ int rng_cell(Rng* r) {
return (int)(((uint64_t)rng_next(r) * 11ULL) >> 32); // uniform in [0,10]
}
// Lane-parallel policy forward. Lane l owns hidden units l and l+32.
// Fills out5[0..3]=logits, out5[4]=value in every lane; returns h1,h2.
__device__ __forceinline__ void policy_fwd(const float* w, int lane, const float obs[4],
float* h1o, float* h2o, float* out5) {
float a1 = w[256 + lane], a2 = w[256 + 32 + lane];
#pragma unroll
for (int k = 0; k < 4; ++k) {
a1 += w[k * 64 + lane] * obs[k];
a2 += w[k * 64 + 32 + lane] * obs[k];
}
float h1 = __tanhf(a1);
float h2 = __tanhf(a2);
*h1o = h1; *h2o = h2;
#pragma unroll
for (int m = 0; m < 4; ++m) {
out5[m] = w[320 + m * 64 + lane] * h1 + w[320 + m * 64 + 32 + lane] * h2;
}
out5[4] = w[580 + lane] * h1 + w[580 + 32 + lane] * h2;
#pragma unroll
for (int off = 16; off >= 1; off >>= 1) {
#pragma unroll
for (int m = 0; m < 5; ++m) {
out5[m] += __shfl_xor_sync(0xffffffffu, out5[m], off);
}
}
#pragma unroll
for (int m = 0; m < 4; ++m) out5[m] += w[576 + m];
out5[4] += w[644];
}
// grad wrt hidden unit j given head weights and upstream head grads (dz, dv)
__device__ __forceinline__ float w_dot_head(const float* w, int j, const float dz[4], float dv) {
return w[320 + j] * dz[0] + w[320 + 64 + j] * dz[1] + w[320 + 128 + j] * dz[2]
+ w[320 + 192 + j] * dz[3] + w[580 + j] * dv;
}
template<int NT, int NWARP>
__device__ __forceinline__ void
ppo_body(float4* __restrict__ obsQ, // [N_SAMPLES] observations
float4* __restrict__ metaQ, // [N_SAMPLES] (logp_old, act, adv, ret)
float* __restrict__ curveG, // [iters] reward sums (unscaled)
double* __restrict__ statsG, // [2] sum, sumsq of advantages
float* __restrict__ gslotG, // [iters*16*645] grad slots
int* __restrict__ listG, // [16*40960] minibatch member lists
int* __restrict__ cntG, // [iters*16] per-minibatch counts
int iters, uint64_t seed) {
cg::grid_group grid = cg::this_grid();
__shared__ SharedMem<NWARP> smobj;
SharedMem<NWARP>* sm = &smobj;
__shared__ float gvals[NPARAM];
const int tid = threadIdx.x;
const int lane = tid & 31;
const int wid = tid >> 5;
const int nb = gridDim.x;
const int gw = blockIdx.x * NWARP + wid; // global warp id
const int TW = nb * NWARP; // total warps
// ---- weight + Adam-state init (identical in every block) ----------------
for (int k = tid; k < NPARAM; k += NT) {
float val;
if (k < 256) {
val = u01(splitmix64(seed ^ (0x1A2B3C4D00000000ULL + (uint64_t)k))) - 0.5f;
} else if (k < 320) {
val = 0.0f;
} else if (k < 576) {
val = (u01(splitmix64(seed ^ (0x5E6F7A8B00000000ULL + (uint64_t)k))) - 0.5f) * 0.25f;
} else if (k < 580) {
val = 0.0f;
} else if (k < 644) {
val = (u01(splitmix64(seed ^ (0x9C0D1E2F00000000ULL + (uint64_t)k))) - 0.5f) * 0.25f;
} else {
val = 0.0f;
}
sm->w[k] = val;
sm->am[k] = 0.0f;
sm->av[k] = 0.0f;
}
__syncthreads();
for (int it = 0; it < iters; ++it) {
// ========================= ROLLOUT + GAE ===========================
if (tid == 0) { sm->stats_blk[0] = 0.0; sm->stats_blk[1] = 0.0; sm->curve_blk = 0.0f; }
__syncthreads();
const int e_lo = (int)(((long long)gw * NUM_ENVS) / TW);
const int e_hi = (int)(((long long)(gw + 1) * NUM_ENVS) / TW);
for (int e = e_lo; e < e_hi; ++e) {
Rng rng; rng.s = splitmix64((seed * 0x6A09E667F3BCC909ULL) ^
((uint64_t)(e + 1) * 0x9E3779B97F4A7C15ULL) ^
((uint64_t)(it + 1) * 0x165667B19E3779F9ULL));
int ax = rng_cell(&rng), ay = rng_cell(&rng);
int fx = rng_cell(&rng), fy = rng_cell(&rng);
float env_ret = 0.0f;
#pragma unroll 1
for (int t = 0; t < T_STEPS; ++t) {
float obs[4] = { (float)(fx - ax) * (1.0f / 11.0f),
(float)(fy - ay) * (1.0f / 11.0f),
(float)ax * 0.1f, (float)ay * 0.1f };
float h1, h2, out5[5];
policy_fwd(sm->w, lane, obs, &h1, &h2, out5);
// categorical sample from softmax(logits)
float mx = fmaxf(fmaxf(out5[0], out5[1]), fmaxf(out5[2], out5[3]));
float ex0 = __expf(out5[0] - mx), ex1 = __expf(out5[1] - mx);
float ex2 = __expf(out5[2] - mx), ex3 = __expf(out5[3] - mx);
float Z = ex0 + ex1 + ex2 + ex3;
float invZ = __frcp_rn(Z);
float q0 = ex0 * invZ, q1 = ex1 * invZ, q2 = ex2 * invZ;
float u = rng_u01(&rng);
int act = (u < q0) ? 0 : (u < q0 + q1) ? 1 : (u < q0 + q1 + q2) ? 2 : 3;
float lz = __logf(Z);
float sel = (act == 0) ? out5[0] : (act == 1) ? out5[1] : (act == 2) ? out5[2] : out5[3];
float logp = (sel - mx) - lz;
const int s = t * NUM_ENVS + e;
if (lane == 0) obsQ[s] = make_float4(obs[0], obs[1], obs[2], obs[3]);
else if (lane == 1) metaQ[s].x = logp;
else if (lane == 2) metaQ[s].y = (float)act;
// env step: 0=up,1=down,2=left,3=right, clamped to [0,10]
if (act == 0) ay = (ay > 0) ? ay - 1 : 0;
else if (act == 1) ay = (ay < 10) ? ay + 1 : 10;
else if (act == 2) ax = (ax > 0) ? ax - 1 : 0;
else ax = (ax < 10) ? ax + 1 : 10;
float r = (ax == fx && ay == fy) ? 1.0f : 0.0f;
if (r > 0.0f) { fx = rng_cell(&rng); fy = rng_cell(&rng); }
env_ret += r;
__syncwarp();
if (lane == 3) sm->rew_stash[wid][t] = r;
if (lane == 4) sm->val_stash[wid][t] = out5[4];
}
__syncwarp();
// ---- GAE backward recurrence (all lanes redundantly) -----------
float gae = 0.0f;
float adv_sum = 0.0f, adv_sumsq = 0.0f;
#pragma unroll 1
for (int t = T_STEPS - 1; t >= 0; --t) {
float v_t = sm->val_stash[wid][t];
float r_t = sm->rew_stash[wid][t];
float delta;
if (t == T_STEPS - 1) {
delta = r_t - v_t; // no bootstrap past horizon
gae = delta;
} else {
delta = r_t + GAMMA * sm->val_stash[wid][t + 1] - v_t;
gae = delta + GL * gae;
}
const int s = t * NUM_ENVS + e;
if ((t & 31) == lane) { metaQ[s].z = gae; metaQ[s].w = gae + v_t; }
adv_sum += gae;
adv_sumsq += gae * gae;
}
if (lane == 0) {
atomicAdd(&sm->stats_blk[0], (double)adv_sum);
atomicAdd(&sm->stats_blk[1], (double)adv_sumsq);
atomicAdd(&sm->curve_blk, env_ret);
}
__syncwarp();
}
// ---- build per-minibatch member index lists (balance for update) ----
{
__shared__ int cnt16[16], cur16[16];
for (int i = tid; i < 16; i += NT) cnt16[i] = 0;
__syncthreads();
const int spb = (N_SAMPLES + nb - 1) / nb;
const int s_lo = blockIdx.x * spb;
const int s_hi = min(N_SAMPLES, s_lo + spb);
const uint64_t lkey = splitmix64(seed ^ (0xC6A4A7935BD1E995ULL + (uint64_t)it * 0x9E3779B97F4A7C15ULL));
for (int s0 = s_lo + (wid << 5); s0 < s_hi; s0 += (NWARP << 5)) {
const int ss = s0 + lane;
const uint64_t h = (ss < s_hi) ? splitmix64((uint64_t)ss ^ lkey) : 0ULL;
const bool ok = (ss < s_hi);
#pragma unroll
for (int ep = 0; ep < 4; ++ep) {
const int mbv = (int)((h >> (56 + 2 * ep)) & 3ULL);
#pragma unroll
for (int m = 0; m < 4; ++m) {
const unsigned b = __ballot_sync(0xffffffffu, ok && mbv == m);
if (lane == 0 && b) atomicAdd(&cnt16[ep * 4 + m], __popc(b));
}
}
}
__syncthreads();
if (tid < 16) {
cur16[tid] = atomicAdd(&cntG[it * 16 + tid], cnt16[tid]);
}
__syncthreads();
for (int s0 = s_lo + (wid << 5); s0 < s_hi; s0 += (NWARP << 5)) {
const int ss = s0 + lane;
const uint64_t h = (ss < s_hi) ? splitmix64((uint64_t)ss ^ lkey) : 0ULL;
const bool ok = (ss < s_hi);
#pragma unroll
for (int ep = 0; ep < 4; ++ep) {
const int mbv = (int)((h >> (56 + 2 * ep)) & 3ULL);
#pragma unroll
for (int m = 0; m < 4; ++m) {
const unsigned b = __ballot_sync(0xffffffffu, ok && mbv == m);
const int pc = __popc(b);
if (pc) {
int pos0 = 0;
if (lane == 0) pos0 = atomicAdd(&cur16[ep * 4 + m], pc);
pos0 = __shfl_sync(0xffffffffu, pos0, 0);
if (ok && mbv == m) {
const int r = __popc(b & ((1u << lane) - 1u));
listG[(ep * 4 + m) * 40960 + pos0 + r] = ss;
}
}
}
}
}
__syncthreads();
}
__syncthreads();
if (tid == 0) {
double* st = statsG + 2 * it;
atomicAdd(&st[0], sm->stats_blk[0]);
atomicAdd(&st[1], sm->stats_blk[1]);
atomicAdd(&curveG[it], sm->curve_blk);
}
grid.sync();
const double mean = statsG[2 * it] * INV_N;
double var = statsG[2 * it + 1] * INV_N - mean * mean;
var = var * ((double)N_SAMPLES / (double)(N_SAMPLES - 1)); // unbiased std
const float inv_std = (float)(1.0 / (sqrt(var) + 1e-8));
const float fmean = (float)mean;
// ========================= PPO UPDATE ==============================
for (int epc = 0; epc < 4; ++epc) {
for (int mb = 0; mb < 4; ++mb) {
const int midx = epc * 4 + mb;
const int slot0 = (it * N_MB + midx) * NPARAM;
for (int k = tid; k < NPARAM; k += NT) sm->gacc[k] = 0.0f;
__syncthreads();
// per-lane gradient accumulators (lane owns its units' slots)
float a_w1[8], a_w2[8], a_b1[2] = {0.0f, 0.0f};
float a_b2 = 0.0f, a_w3[2] = {0.0f, 0.0f}, a_b3 = 0.0f;
#pragma unroll
for (int i = 0; i < 8; ++i) { a_w1[i] = 0.0f; a_w2[i] = 0.0f; }
// balanced list-driven processing; 2 samples in flight
const int mcnt = cntG[it * 16 + midx];
const int* mlist = listG + midx * 40960;
const int ilo = (int)(((long long)mcnt * gw) / TW);
const int ihi = (int)(((long long)mcnt * (gw + 1)) / TW);
int ps0 = (ilo < ihi) ? mlist[ilo] : -1;
int ps1 = (ilo + 1 < ihi) ? mlist[ilo + 1] : -1;
#pragma unroll 1
for (int i0 = ilo; i0 < ihi; i0 += 2) {
const int ns0 = (i0 + 2 < ihi) ? mlist[i0 + 2] : -1;
const int ns1 = (i0 + 3 < ihi) ? mlist[i0 + 3] : -1;
float4 ob[2], mt[2];
float valid[2] = {0.0f, 0.0f};
#pragma unroll
for (int u = 0; u < 2; ++u) {
const int s = u ? ps1 : ps0;
if (s >= 0) {
ob[u] = obsQ[s];
mt[u] = metaQ[s];
valid[u] = 1.0f;
} else {
ob[u] = make_float4(0.f, 0.f, 0.f, 0.f);
mt[u] = make_float4(0.f, 0.f, 0.f, 0.f);
}
}
ps0 = ns0; ps1 = ns1;
float obs[2][4], out5[2][5], h1[2], h2[2];
#pragma unroll
for (int u = 0; u < 2; ++u) {
obs[u][0] = ob[u].x; obs[u][1] = ob[u].y; obs[u][2] = ob[u].z; obs[u][3] = ob[u].w;
policy_fwd(sm->w, lane, obs[u], &h1[u], &h2[u], out5[u]);
}
#pragma unroll
for (int u = 0; u < 2; ++u) {
const int act = (int)mt[u].y;
const float advn = ((mt[u].z - fmean) * inv_std) * valid[u];
float mx = fmaxf(fmaxf(out5[u][0], out5[u][1]), fmaxf(out5[u][2], out5[u][3]));
float ex0 = __expf(out5[u][0] - mx), ex1 = __expf(out5[u][1] - mx);
float ex2 = __expf(out5[u][2] - mx), ex3 = __expf(out5[u][3] - mx);
float Z = ex0 + ex1 + ex2 + ex3;
float invZ = __frcp_rn(Z);
float q[4] = { ex0 * invZ, ex1 * invZ, ex2 * invZ, ex3 * invZ };
float lZ = __logf(Z);
float lsm[4] = { (out5[u][0] - mx) - lZ, (out5[u][1] - mx) - lZ,
(out5[u][2] - mx) - lZ, (out5[u][3] - mx) - lZ };
float ratio = __expf(lsm[act] - mt[u].x);
float H = -(q[0] * lsm[0] + q[1] * lsm[1] + q[2] * lsm[2] + q[3] * lsm[3]);
float gate = ((advn > 0.0f && ratio > 1.2f) ||
(advn < 0.0f && ratio < 0.8f)) ? 0.0f : 1.0f;
float pgw = -gate * advn * ratio * INV_MB;
float dz[4];
#pragma unroll
for (int m = 0; m < 4; ++m) {
dz[m] = (pgw * ((m == act) ? 1.0f - q[m] : -q[m])
+ (ENT_COEF * INV_MB) * q[m] * (H + lsm[m])) * valid[u];
}
float dv = INV_MB * (out5[u][4] - mt[u].w) * valid[u];
float dh1 = w_dot_head(sm->w, lane, dz, dv);
float dh2 = w_dot_head(sm->w, lane + 32, dz, dv);
float dp1 = dh1 * (1.0f - h1[u] * h1[u]);
float dp2 = dh2 * (1.0f - h2[u] * h2[u]);
#pragma unroll
for (int k = 0; k < 4; ++k) {
a_w1[k] += dp1 * obs[u][k];
a_w1[4 + k] += dp2 * obs[u][k];
}
a_b1[0] += dp1;
a_b1[1] += dp2;
#pragma unroll
for (int m = 0; m < 4; ++m) {
a_w2[m] += dz[m] * h1[u];
a_w2[4 + m] += dz[m] * h2[u];
}
if (lane < 4) a_b2 += dz[lane];
a_w3[0] += dv * h1[u];
a_w3[1] += dv * h2[u];
if (lane == 0) a_b3 += dv;
}
}
// commit per-lane accumulators -> shared-memory grad slots
#pragma unroll
for (int k = 0; k < 4; ++k) {
atomicAdd(&sm->gacc[k * 64 + lane], a_w1[k]);
atomicAdd(&sm->gacc[k * 64 + 32 + lane], a_w1[4 + k]);
}
atomicAdd(&sm->gacc[256 + lane], a_b1[0]);
atomicAdd(&sm->gacc[256 + 32 + lane], a_b1[1]);
#pragma unroll
for (int m = 0; m < 4; ++m) {
atomicAdd(&sm->gacc[320 + m * 64 + lane], a_w2[m]);
atomicAdd(&sm->gacc[320 + m * 64 + 32 + lane], a_w2[4 + m]);
}
if (lane < 4) atomicAdd(&sm->gacc[576 + lane], a_b2);
atomicAdd(&sm->gacc[580 + lane], a_w3[0]);
atomicAdd(&sm->gacc[580 + 32 + lane], a_w3[1]);
if (lane == 0) atomicAdd(&sm->gacc[644], a_b3);
__syncthreads();
for (int k = tid; k < NPARAM; k += NT) {
atomicAdd(&gslotG[slot0 + k], sm->gacc[k]);
}
grid.sync();
// ---- identical grad-clip + Adam in every block -------------
float part = 0.0f;
for (int k = tid; k < NPARAM; k += NT) {
float g = gslotG[slot0 + k];
gvals[k] = g;
part += g * g;
}
#pragma unroll
for (int off = 16; off >= 1; off >>= 1) part += __shfl_xor_sync(0xffffffffu, part, off);
if (lane == 0) sm->red[wid] = part;
__syncthreads();
float gtot = 0.0f;
#pragma unroll
for (int i = 0; i < NWARP; ++i) gtot += sm->red[i];
const float clip_coef = (gtot > 0.0f)
? fminf(1.0f, MAXN / (__fsqrt_rn(gtot) + 1e-6f)) : 1.0f;
const int t_step = it * N_MB + midx + 1;
const float bc1 = 1.0f - __powf(BETA1, (float)t_step);
const float bc2 = 1.0f - __powf(BETA2, (float)t_step);
for (int k = tid; k < NPARAM; k += NT) {
float g = gvals[k] * clip_coef;
float m1 = BETA1 * sm->am[k] + (1.0f - BETA1) * g;
float v1 = BETA2 * sm->av[k] + (1.0f - BETA2) * g * g;
sm->am[k] = m1;
sm->av[k] = v1;
float mhat = m1 / bc1;
float vhat = v1 / bc2;
sm->w[k] -= LRF * mhat / (__fsqrt_rn(vhat) + A_EPS);
}
__syncthreads();
}
}
}
}
// --------------------------------------------------------------------------
// two instantiations: NT=512 (preferred: fewest barrier arrivals) and NT=256
// (fallback if the 512-thread variant does not fit on the grading arch)
// --------------------------------------------------------------------------
extern "C" __global__ void __launch_bounds__(512)
ppo_k512(float4* o, float4* m, float* c, double* st, float* g, int* l, int* cn, int it, uint64_t sd) {
ppo_body<512, 16>(o, m, c, st, g, l, cn, it, sd);
}
extern "C" __global__ void __launch_bounds__(256)
ppo_k256(float4* o, float4* m, float* c, double* st, float* g, int* l, int* cn, int it, uint64_t sd) {
ppo_body<256, 8>(o, m, c, st, g, l, cn, it, sd);
}
static int choose_nt() {
static int cached = -1;
if (cached > 0) return cached;
int dev = 0; cudaGetDevice(&dev);
int coop = 0;
cudaDeviceGetAttribute(&coop, cudaDevAttrCooperativeLaunch, dev);
TORCH_CHECK(coop, "device does not support cooperative launch");
int per_sm = 0;
cudaError_t oerr = cudaOccupancyMaxActiveBlocksPerMultiprocessor(&per_sm, (const void*)ppo_k512, 512, 0);
int nt = 256;
if (oerr == cudaSuccess && per_sm > 0) nt = 512;
else {
cudaError_t o2 = cudaOccupancyMaxActiveBlocksPerMultiprocessor(&per_sm, (const void*)ppo_k256, 256, 0);
TORCH_CHECK(o2 == cudaSuccess && per_sm > 0, "kernel does not fit on any SM");
}
cached = nt;
return nt;
}
int64_t query_nt() { return choose_nt(); }
void run(torch::Tensor obsQ, torch::Tensor metaQ, torch::Tensor curve,
torch::Tensor stats, torch::Tensor gslot, torch::Tensor listT, torch::Tensor cntT, int64_t iters, int64_t seed) {
const int nt = choose_nt();
const void* func = (nt == 512) ? (const void*)ppo_k512 : (const void*)ppo_k256;
int dev = 0; cudaGetDevice(&dev);
cudaDeviceProp prop; cudaGetDeviceProperties(&prop, dev);
int per_sm = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&per_sm, func, nt, 0);
int nb = per_sm * prop.multiProcessorCount;
if (nb > 512) nb = 512;
float4* obsP = (float4*)obsQ.data_ptr<float>();
float4* metaP = (float4*)metaQ.data_ptr<float>();
float* curveP = curve.data_ptr<float>();
double* statsP = stats.data_ptr<double>();
float* gslotP = gslot.data_ptr<float>();
int* listP = (int*)listT.data_ptr<int>();
int* cntP = (int*)cntT.data_ptr<int>();
int it = (int)iters;
uint64_t sd = (uint64_t)seed;
void* args[] = { (void*)&obsP, (void*)&metaP, (void*)&curveP,
(void*)&statsP, (void*)&gslotP, (void*)&listP, (void*)&cntP, (void*)&it, (void*)&sd };
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
cudaError_t err = cudaLaunchCooperativeKernel(func, dim3(nb), dim3(nt), args, 0, stream);
TORCH_CHECK(err == cudaSuccess, "cooperative launch failed: ", cudaGetErrorString(err));
cudaError_t lerr = cudaGetLastError();
TORCH_CHECK(lerr == cudaSuccess, "kernel launch error: ", cudaGetErrorString(lerr));
}
"""
def _build():
from torch.utils.cpp_extension import load_inline
p = torch.cuda.get_device_properties(0)
os.environ["TORCH_CUDA_ARCH_LIST"] = f"{p.major}.{p.minor}"
return load_inline(
name="ppo_megakernel_v11",
cpp_sources=(
"void run(torch::Tensor obsQ, torch::Tensor metaQ, torch::Tensor curve,"
" torch::Tensor stats, torch::Tensor gslot, torch::Tensor listT, torch::Tensor cntT,"
" int64_t iters, int64_t seed);"
"int64_t query_nt();"
),
cuda_sources=_CUDA_SRC,
functions=["run", "query_nt"],
extra_cuda_cflags=["-O3", "--restrict"],
verbose=False,
)
_EXT = None
def _ext():
global _EXT
if _EXT is None:
_EXT = _build()
return _EXT
def train(total_env_steps: int, seed: int) -> list[float]:
iters = max(1, int(total_env_steps) // (ROLLOUT * NUM_ENVS))
dev = torch.device("cuda:0")
ext = _ext()
obsQ = torch.empty(N_SAMPLES, 4, device=dev, dtype=torch.float32)
metaQ = torch.empty(N_SAMPLES, 4, device=dev, dtype=torch.float32)
curve = torch.zeros(iters, device=dev, dtype=torch.float32)
stats = torch.zeros(2 * iters, device=dev, dtype=torch.float64)
gslot = torch.zeros(iters * N_MINIBATCHES * NPARAM, device=dev, dtype=torch.float32)
listT = torch.empty(N_MINIBATCHES * 40960, device=dev, dtype=torch.int32)
cntT = torch.zeros(iters * N_MINIBATCHES, device=dev, dtype=torch.int32)
ext.run(obsQ, metaQ, curve, stats, gslot, listT, cntT, iters, int(seed))
torch.cuda.synchronize()
return (curve * (1.0 / NUM_ENVS)).cpu().tolist()
if __name__ == "__main__":
import time
t0 = time.perf_counter()
c = train(ROLLOUT * NUM_ENVS * 40, seed=0)
torch.cuda.synchronize()
dt = time.perf_counter() - t0
print(f"iters={len(c)} first={c[0]:.3f} last={c[-1]:.3f}")
print(f"curve_head={[round(x,2) for x in c[:5]]} tail={[round(x,2) for x in c[-5:]]}")
print(f"elapsed={dt:.3f}s sps={ROLLOUT * NUM_ENVS * len(c) / dt:,.0f}")
20260716_091652_kinetic-claude_kinetic-0715_01_rl_grid_ppo