KernelBench mega · H100

rl grid ppo Kimi K3 (1M)

cleandid not score

manually audited: clean

Clean and megakernel-authentic at 23.1370x on H100 (578,425,950 SPS vs the 25,000,000 SPS anchor). The entire PPO training run — 32-step x 4096-env rollout with cooperative tiny-MLP forward, categorical sampling, fused GAE, global advantage stats, and 4 epochs x 4 minibatches of exact clipped-PPO gradient + grad-norm clip + Adam — runs inside ONE cooperative kernel launch per train() call, with parameters and Adam state resident in block shared memory (all block copies kept bit-identical by redundant Adam). Seed-keyed Philox4x32-10 drives env resets, sampling, and respawns; an 18-bit Feistel bijection with cycle-walking replaces randperm for exact minibatch permutation. No torch.cuda.CUDAGraph or graph capture anywhere, no per-step Python loop, no caching keyed on inputs, no check.py sniffing, no getattr/exec obfuscation, and all grader templates byte-match the snapshot. Independent design from the audited-clean RTX PRO 6000 sibling (28.7578x, run 20260716_233413): different solution md5 (89deb86b72d91bc85318aaf09f261f05 vs a63618ba936472df072c942b6e783e69), different env-group width (8 lanes/env here vs half-warp there), and zero transcript references to the sibling run.

harnesskinetic-claude
Kernel source (redacted)
"""Grid-foraging PPO training megakernel (fused, single persistent launch).

The whole training run -- environment rollout, policy forward, action sampling,
GAE, advantage normalization, and all 16 minibatch PPO updates per iteration --
executes inside ONE cooperative CUDA kernel launch per train() call. The launch
count is therefore independent of env-steps (32/iter), horizon, and minibatches
(16/iter): exactly one launch per run no matter the iteration count.

Task semantics match reference.py exactly (see PROMPT): same MDP (11x11 board,
clamped moves, +1 reward for stepping onto the food, uniform respawn, board
re-randomized every 32 steps), same network (Linear(4,64)-tanh -> Linear(64,4)
and Linear(64,1) heads), same PPO (GAE gamma 0.99 / lambda 0.95 terminated at
the horizon, clip 0.2, 4 epochs x 4 minibatches, advantage standardized with
the unbiased std, Adam lr 3e-3 betas (0.9, 0.999) eps 1e-8, entropy coef 0.01,
value coef 0.5, grad-norm clip 0.5 with torch's convention), same per-iteration
mean-episodic-return curve.

Kernel organization (grid = NB blocks x TPB threads, chosen by occupancy):
  P1 rollout + GAE (per iteration): each of the 4096 envs is owned by an
      8-lane group (4 envs per warp). The tiny MLP is evaluated cooperatively
      with the 645 parameters broadcast from block shared memory; categorical
      sampling and the env update run redundantly-but-identically across the
      group (so they need no cross-lane communication beyond the logits
      xor-reduction). RNG is counter-based Philox4x32-10 keyed by (seed; env,
      iteration, draw), pre-generated per env into shared memory. The (value,
      reward) trace for GAE stays in shared memory; GAE itself runs in the
      same phase, followed by advantage / return statistics accumulation.
  P3: 16 sequential minibatch updates. Advantage normalization scalars are
      computed inline at the first minibatch. Each minibatch:
      A) per-thread forward + closed-form per-sample loss gradients wrt the
         logits/value (exactly the gradient of the reference's clipped-PPO +
         entropy + value loss; verified against autograd elementwise), staged
         in block shared memory. Batch order per epoch is a genuine pseudo-
         random permutation: an 18-bit 4-round Feistel cipher with cycle-
         walking (bijective, computed on the fly, keyed per (seed, epoch)).
      B) parameter-gradient accumulation: each thread owns a PAIR of hidden
         units (so each staged (obs, dlogits, dvalue) record feeds 20
         accumulators), hidden activation & delta recomputed per sample --
         bit-identical to A's forward path, same instruction stream and same
         smem parameters. Sample/chunk assignment is interleaved so warps hit
         distinct smem banks. Warp-shuffle reduction over chunks, then one
         striped global red.add per parameter per block.
      grid.sync();
      Adam: every block redundantly reads the fence-visible total gradient,
         applies grad-norm clipping and the Adam update to its *shared-memory*
         copy of parameters and optimizer state with identical arithmetic, so
         all block copies stay bit-identical -- this removes the second barrier
         that a single-writer Adam would need on the minibatch path. Striped
         partial-gradient slots are re-zeroed here (double-buffered by parity
         so no zero/add race exists; every add is barrier-separated from the
         zeroing of its slot).
Layout note: the policy head weight W2 (nn.Linear(HIDDEN, 4) weight) is stored
TRANSPOSED (W2T[j][k]) inside the kernel so hidden-unit slices are 16B vector
loads in every phase; gradients are accumulated in the same transposed layout,
so the layout is invisible to the semantics of the update. Full parameter
vector layout: W1(64x4) | b1(64) | W2T(64x4) | W3(64) | b2(4) | b3(1).
"""
from __future__ import annotations

import torch
from torch.utils.cpp_extension import load_inline

_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cooperative_groups.h>
#include <stdint.h>

namespace cg = cooperative_groups;

// ----------------------------- task constants ------------------------------
#define GRID        11
#define NUM_ENVS    4096
#define HORIZON     32
#define BATCH       (HORIZON * NUM_ENVS)          // 131072
#define MB_SIZE     (BATCH / 4)                   // 32768
#define GAMMA       0.99f
#define LAMBDA      0.95f
#define CLIP        0.2f
#define EPOCHS      4
#define MINIBATCHES 4
#define LR          3.0e-3f
#define ENT_COEF    0.01f
#define MAXGN       0.5f

#define NPARAM      645   // W1 256 | b1 64 | W2T 256 | W3 64 | b2 4 | b3 1
#define NPAD        648   // NPARAM rounded up for aligned float4 grad slots
#define OFF_W1      0
#define OFF_B1      256
#define OFF_W2      320
#define OFF_W3      576
#define OFF_B2      640
#define OFF_B3      644

#define NSLOT       2     // striped gradient slots per parity (contention control)

// ------------------------------ math helpers -------------------------------
__device__ __forceinline__ float ftanh(float x) {
    float r;
    asm("tanh.approx.f32 %0, %1;" : "=f"(r) : "f"(x));   // HW SFU tanh (~5e-4 abs err)
    return r;
}

__device__ __forceinline__ uint4 philox4x32_10(uint4 ctr, uint2 key) {
    #pragma unroll
    for (int i = 0; i < 10; i++) {
        unsigned hi0 = __umulhi(ctr.x, 0xD2511F53u), lo0 = ctr.x * 0xD2511F53u;
        unsigned hi1 = __umulhi(ctr.z, 0xCD9E8D57u), lo1 = ctr.z * 0xCD9E8D57u;
        ctr.x = hi1 ^ ctr.y ^ key.x;  ctr.y = lo1;
        ctr.z = hi0 ^ ctr.w ^ key.y;  ctr.w = lo0;
        key.x += 0x9E3779B9u;         key.y += 0xBB67AE85u;
    }
    return ctr;
}

__device__ __forceinline__ float u01(unsigned x) {
    return (float)(x >> 8) * (1.0f / 16777216.0f);       // [0,1)
}
__device__ __forceinline__ int ucell(unsigned x) {
    return (int)(((unsigned long long)x * GRID) >> 32);  // uniform [0,10]
}

// Bijective permutation of [0, 2^17): 18-bit Feistel (4 rounds) with
// cycle-walking back into 17 bits. Keyed per (seed, epoch).
__device__ __forceinline__ unsigned perm17(unsigned x, unsigned k) {
    unsigned l = x & 511u, r = (x >> 9) & 511u;
    do {
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            unsigned f = r * (k + 0x9E3779B9u * (unsigned)i);
            f = (f ^ (f >> 16)) * 0x85EBCA6Bu;
            unsigned t = r; r = l ^ ((f ^ (f >> 13)) & 511u); l = t;
        }
        x = (r << 9) | l;
    } while (x >= (unsigned)BATCH);
    return x;
}

__device__ __forceinline__ unsigned fmix(unsigned h) {
    h ^= h >> 16; h *= 0x85EBCA6Bu; h ^= h >> 13; h *= 0xC2B2AE35u; h ^= h >> 16;
    return h;
}

#define GSYNC() grid.sync()

__device__ __forceinline__ void block_reduce_sum3(float& x, float& y, float& z, float* red) {
    #pragma unroll
    for (int off = 16; off; off >>= 1) {
        x += __shfl_down_sync(~0u, x, off);
        y += __shfl_down_sync(~0u, y, off);
        z += __shfl_down_sync(~0u, z, off);
    }
    int wid = threadIdx.x >> 5, lane = threadIdx.x & 31;
    int nw = blockDim.x >> 5;
    if (!lane) { red[wid] = x; red[nw + wid] = y; red[2 * nw + wid] = z; }
    __syncthreads();
    if (!wid) {
        float vx = (lane < nw) ? red[lane] : 0.0f;
        float vy = (lane < nw) ? red[nw + lane] : 0.0f;
        float vz = (lane < nw) ? red[2 * nw + lane] : 0.0f;
        #pragma unroll
        for (int off = 16; off; off >>= 1) {
            vx += __shfl_down_sync(~0u, vx, off);
            vy += __shfl_down_sync(~0u, vy, off);
            vz += __shfl_down_sync(~0u, vz, off);
        }
        if (!lane) { red[0] = vx; red[1] = vy; red[2] = vz; }
    }
    __syncthreads();
    x = red[0]; y = red[1]; z = red[2];
    __syncthreads();
}

// ------------------------------ the kernel ---------------------------------
template <int NB, int TPB>
__global__ void __launch_bounds__(TPB)
ppo_fused(float4* __restrict__ samp_buf,  // [BATCH][2]: [obs4, act|logp|adv|ret]
          float* __restrict__ gp,         // [2][NSLOT][NPARAM] partial grads
          float* __restrict__ stats,      // [3] {sum_rew, sum_adv, sum_adv2}
          float* __restrict__ curve,      // [iters]
          int iters, unsigned long long seed) {
    constexpr int EPG    = TPB >> 3;              // env groups per block (8 lanes/env)
    constexpr int NGRP   = NB * EPG;              // total env groups
    constexpr int EPG_R  = (EPG < NUM_ENVS / NB) ? EPG : (NUM_ENVS / NB);
    constexpr int SPB    = MB_SIZE / NB;          // staged minibatch samples/block
    constexpr int LNC    = (TPB == 512) ? 4 : ((TPB == 256) ? 3 : 2);  // log2(sample chunks)
    constexpr int NCHUNK = 1 << LNC;              // sample chunks per block in B
    constexpr int SPC    = SPB / NCHUNK;          // samples per (j-pair, chunk)
    static_assert(MB_SIZE % NB == 0, "slice");
    static_assert(SPB / 2 <= TPB, "A 2-sample mapping fits");
    static_assert((1 << LNC) * 32 == TPB, "B pair mapping");
    static_assert(NUM_ENVS % NGRP == 0 || NGRP % NUM_ENVS == 0, "env groups");
    static_assert((TPB & 7) == 0, "TPE8");
    __shared__ __align__(16) float sp[NPARAM];    // parameters (block-local copy)
    __shared__ float sm[NPARAM];    // Adam m
    __shared__ float sv[NPARAM];    // Adam v
    __shared__ __align__(16) float s_obs[SPB * 4];
    __shared__ __align__(16) float s_dl[SPB * 4];
    __shared__ float s_dv[SPB];
    __shared__ __align__(16) float s_acc[NPAD];
    __shared__ float s_bias[5];
    __shared__ __align__(8) float s_b1w3[128];
    __shared__ __align__(8) float2 s_vr[EPG_R * HORIZON];   // per-env (value, reward) trace
    __shared__ float s_red[3 * TPB / 32 + 8];
    __shared__ unsigned s_rng[EPG_R * 33 * 4];

    const int tid = threadIdx.x;
    const int bid = blockIdx.x;
    cg::grid_group grid = cg::this_grid();

    // ============================== init (once) =============================
    {
        for (int p = tid; p < NPARAM; p += TPB) {
            unsigned h = fmix((unsigned)seed ^ (unsigned)p * 0x9E3779B9u);
            float u = u01(h) * 2.0f - 1.0f;
            float bound = (p < OFF_B1) ? 0.5f : ((p < OFF_B2) ? 0.125f : 0.0f);
            sp[p] = u * bound;
            sm[p] = 0.0f;
            sv[p] = 0.0f;
        }
        for (int j = tid; j < 64; j += TPB) {
            s_b1w3[j * 2] = sp[OFF_B1 + j];
            s_b1w3[j * 2 + 1] = sp[OFF_W3 + j];
        }
        if (tid < 3) stats[tid] = 0.0f;
        if (tid < 5) s_bias[tid] = 0.0f;      // used by the first minibatch's A
        for (int p = tid; p < 2 * NSLOT * NPAD; p += TPB) gp[p] = 0.0f;
    }
    GSYNC();

    for (int it = 0; it < iters; it++) {
        // ======================= P1: rollout + GAE ==========================
        float s_rew = 0.0f, s_adv = 0.0f, s_adv2 = 0.0f;
        for (int env = bid * EPG + (tid >> 3); env < NUM_ENVS; env += NGRP) {
            const int L = tid & 7;
            const int env_loc = tid >> 3;               // env slot within block
            // -- pre-generate RNG: 33 philox draws per env into smem
            {
                unsigned* dst = &s_rng[env_loc * 33 * 4];
                uint2 key = make_uint2((unsigned)seed, (unsigned)(seed >> 32));
                for (int d = L; d < 33; d += 8) {
                    uint4 ctr = make_uint4((unsigned)env, (unsigned)it, (unsigned)d, 0u);
                    uint4 r = philox4x32_10(ctr, key);
                    dst[d * 4 + 0] = r.x; dst[d * 4 + 1] = r.y;
                    dst[d * 4 + 2] = r.z; dst[d * 4 + 3] = r.w;
                }
            }
            __syncthreads();
            const float b2_0 = sp[OFF_B2 + 0], b2_1 = sp[OFF_B2 + 1];
            const float b2_2 = sp[OFF_B2 + 2], b2_3 = sp[OFF_B2 + 3];
            const float b3 = sp[OFF_B3];

            const unsigned* rng = &s_rng[env_loc * 33 * 4];
            int ax = ucell(rng[0]), ay = ucell(rng[1]);
            int fx = ucell(rng[2]), fy = ucell(rng[3]);

            for (int t = 0; t < HORIZON; t++) {
                const int gi = t * NUM_ENVS + env;
                float o0 = (float)(fx - ax) * (1.0f / GRID);
                float o1 = (float)(fy - ay) * (1.0f / GRID);
                float o2 = (float)ax * (1.0f / (GRID - 1));
                float o3 = (float)ay * (1.0f / (GRID - 1));
                if (L < 4) ((float*)&samp_buf[gi * 2])[L] = (L == 0) ? o0 : (L == 1) ? o1 : (L == 2) ? o2 : o3;
                float z0 = 0.0f, z1 = 0.0f, z2 = 0.0f, z3 = 0.0f, vv = 0.0f;
                #pragma unroll
                for (int q = 0; q < 8; q++) {
                    const int j = L + 8 * q;
                    float4 w1v = *(float4*)&sp[OFF_W1 + j * 4];
                    float4 w2v = *(float4*)&sp[OFF_W2 + j * 4];
                    float2 bw = *(float2*)&s_b1w3[j * 2];
                    float hq = ftanh(bw.x + w1v.x * o0 + w1v.y * o1 + w1v.z * o2 + w1v.w * o3);
                    z0 += hq * w2v.x; z1 += hq * w2v.y;
                    z2 += hq * w2v.z; z3 += hq * w2v.w;
                    vv += hq * bw.y;
                }
                #pragma unroll
                for (int off = 4; off; off >>= 1) {
                    z0 += __shfl_xor_sync(~0u, z0, off);
                    z1 += __shfl_xor_sync(~0u, z1, off);
                    z2 += __shfl_xor_sync(~0u, z2, off);
                    z3 += __shfl_xor_sync(~0u, z3, off);
                    vv += __shfl_xor_sync(~0u, vv, off);
                }
                z0 += b2_0; z1 += b2_1; z2 += b2_2; z3 += b2_3;
                float val = vv + b3;
                float mx = fmaxf(fmaxf(z0, z1), fmaxf(z2, z3));
                float e0 = __expf(z0 - mx), e1 = __expf(z1 - mx);
                float e2 = __expf(z2 - mx), e3 = __expf(z3 - mx);
                float es = e0 + e1 + e2 + e3;
                float u = u01(rng[(t + 1) * 4]);
                float tgt = u * es;
                int act = 0; float cum = e0;
                if (tgt > cum) { act = 1; cum += e1; }
                if (tgt > cum) { act = 2; cum += e2; }
                if (tgt > cum) { act = 3; }
                float zsel = (act == 0) ? z0 : (act == 1) ? z1 : (act == 2) ? z2 : z3;
                float logp = zsel - mx - __logf(es);
                if (L == 0) ((float*)&samp_buf[gi * 2 + 1])[0] = (float)act;
                if (L == 1) ((float*)&samp_buf[gi * 2 + 1])[1] = logp;
                int dx = (act == 2) ? -1 : (act == 3) ? 1 : 0;
                int dy = (act == 0) ? -1 : (act == 1) ? 1 : 0;
                ax = min(max(ax + dx, 0), GRID - 1);
                ay = min(max(ay + dy, 0), GRID - 1);
                float rew = (ax == fx && ay == fy) ? 1.0f : 0.0f;
                if (rew != 0.0f) {
                    fx = ucell(rng[(t + 1) * 4 + 1]);
                    fy = ucell(rng[(t + 1) * 4 + 2]);
                }
                if (L == 1) s_vr[env_loc * HORIZON + t] = make_float2(val, rew);
            }
            // ---- last value (value-only forward) + fused GAE ----
            __syncwarp();   // s_vr trace written by lane 1 of this env group
            {
                float o0 = (float)(fx - ax) * (1.0f / GRID);
                float o1 = (float)(fy - ay) * (1.0f / GRID);
                float o2 = (float)ax * (1.0f / (GRID - 1));
                float o3 = (float)ay * (1.0f / (GRID - 1));
                float vv = 0.0f;
                #pragma unroll
                for (int q = 0; q < 8; q++) {
                    const int j = L + 8 * q;
                    float4 w1v = *(float4*)&sp[OFF_W1 + j * 4];
                    float2 bw = *(float2*)&s_b1w3[j * 2];
                    float hq = ftanh(bw.x + w1v.x * o0 + w1v.y * o1 + w1v.z * o2 + w1v.w * o3);
                    vv += hq * bw.y;
                }
                #pragma unroll
                for (int off = 4; off; off >>= 1) vv += __shfl_xor_sync(~0u, vv, off);
                float gae = 0.0f, nextv = vv + b3;
                #pragma unroll 4
                for (int t = HORIZON - 1; t >= 0; t--) {
                    const int gi = t * NUM_ENVS + env;
                    float2 vr = s_vr[env_loc * HORIZON + t];
                    float nt = (t == HORIZON - 1) ? 0.0f : 1.0f;
                    float delta = vr.y + GAMMA * nextv * nt - vr.x;
                    gae = delta + GAMMA * LAMBDA * nt * gae;
                    nextv = vr.x;
                    s_rew += vr.y;
                    s_adv += gae;
                    s_adv2 += gae * gae;
                    if (L == 2) ((float*)&samp_buf[gi * 2 + 1])[2] = gae;
                    if (L == 3) ((float*)&samp_buf[gi * 2 + 1])[3] = gae + vr.x;
                }
            }
            __syncthreads();    // protect s_rng overwrite on the next env pass
        }
        {
            const bool lane0 = ((tid & 7) == 0);
            s_rew = lane0 ? s_rew : 0.0f;
            s_adv = lane0 ? s_adv : 0.0f;
            s_adv2 = lane0 ? s_adv2 : 0.0f;
            block_reduce_sum3(s_rew, s_adv, s_adv2, s_red);
            if (tid == 0) {
                atomicAdd(&stats[0], s_rew);
                atomicAdd(&stats[1], s_adv);
                atomicAdd(&stats[2], s_adv2);
            }
        }
        GSYNC();

        // ================= P3: PPO update (16 minibatches) ==================
        float adv_mean = 0.0f, adv_invstd = 1.0f;
        for (int e = 0; e < EPOCHS; e++) {
            const unsigned ekey = fmix((unsigned)seed ^ ((unsigned)(it * EPOCHS + e) * 0x85EBCA6Bu));
            for (int m = 0; m < MINIBATCHES; m++) {
                const int mb = e * MINIBATCHES + m;
                const int parity = mb & 1;
                if (mb == 0) {
                    // adv stats from the rollout (all blocks read redundantly)
                    float s1 = stats[1], s2 = stats[2];
                    adv_mean = s1 * (1.0f / BATCH);
                    float var = (s2 - BATCH * adv_mean * adv_mean) * (1.0f / (BATCH - 1));
                    adv_invstd = 1.0f / (sqrtf(var) + 1e-8f);
                    if (bid == 0 && tid == 0) curve[it] = stats[0] * (1.0f / NUM_ENVS);
                }
                // ---------------- A: per-sample loss gradients --------------
                {
                    // one sample per thread: block covers [bid*SPB, (bid+1)*SPB)
                    const bool activeA = (tid < SPB);
                    const unsigned pos = (unsigned)(m * MB_SIZE + bid * SPB + (unsigned)tid);
                    float b0 = 0.0f, b1 = 0.0f, b2 = 0.0f, b3 = 0.0f, bv = 0.0f;
                    if (activeA) {
                        const unsigned idx = perm17(pos, ekey);
                        float4 o = samp_buf[idx * 2];
                        float4 mt = samp_buf[idx * 2 + 1];
                        *(float4*)&s_obs[tid * 4] = o;
                        float z0 = sp[OFF_B2 + 0], z1 = sp[OFF_B2 + 1];
                        float z2 = sp[OFF_B2 + 2], z3 = sp[OFF_B2 + 3];
                        float vv = sp[OFF_B3];
                        #pragma unroll 8
                        for (int j = 0; j < 64; j++) {
                            float4 w1v = *(float4*)&sp[OFF_W1 + j * 4];
                            float4 w2v = *(float4*)&sp[OFF_W2 + j * 4];
                            float2 bw = *(float2*)&s_b1w3[j * 2];
                            float pq = bw.x + w1v.x * o.x + w1v.y * o.y + w1v.z * o.z + w1v.w * o.w;
                            float hq = ftanh(pq);
                            z0 += hq * w2v.x; z1 += hq * w2v.y; z2 += hq * w2v.z; z3 += hq * w2v.w;
                            vv += hq * bw.y;
                        }
                        const float invS = 1.0f / MB_SIZE;
                        {
                            const int a_ = (int)mt.x;
                            const float olp = mt.y;
                            const float av_ = (mt.z - adv_mean) * adv_invstd;
                            const float rt_ = mt.w;
                            float mx = fmaxf(fmaxf(z0, z1), fmaxf(z2, z3));
                            float e0 = __expf(z0 - mx), e1 = __expf(z1 - mx);
                            float e2 = __expf(z2 - mx), e3 = __expf(z3 - mx);
                            float es = e0 + e1 + e2 + e3;
                            float inv_es = __fdividef(1.0f, es);
                            float p0 = e0 * inv_es, p1 = e1 * inv_es, p2 = e2 * inv_es, p3 = e3 * inv_es;
                            float lse = mx + __logf(es);
                            float l0 = z0 - lse, l1 = z1 - lse, l2 = z2 - lse, l3 = z3 - lse;
                            float logpn = (a_ == 0) ? l0 : (a_ == 1) ? l1 : (a_ == 2) ? l2 : l3;
                            float H = -(p0 * l0 + p1 * l1 + p2 * l2 + p3 * l3);
                            float r = __expf(logpn - olp);
                            bool cactive = (av_ > 0.0f) ? (r < 1.0f + CLIP) : (av_ < 0.0f && r > 1.0f - CLIP);
                            float c = cactive ? av_ * r : 0.0f;
                            float d0 = -c * invS * ((a_ == 0 ? 1.0f : 0.0f) - p0) + ENT_COEF * invS * p0 * (l0 + H);
                            float d1 = -c * invS * ((a_ == 1 ? 1.0f : 0.0f) - p1) + ENT_COEF * invS * p1 * (l1 + H);
                            float d2 = -c * invS * ((a_ == 2 ? 1.0f : 0.0f) - p2) + ENT_COEF * invS * p2 * (l2 + H);
                            float d3 = -c * invS * ((a_ == 3 ? 1.0f : 0.0f) - p3) + ENT_COEF * invS * p3 * (l3 + H);
                            float dv = (vv - rt_) * invS;
                            *(float4*)&s_dl[tid * 4] = make_float4(d0, d1, d2, d3);
                            s_dv[tid] = dv;
                            b0 = d0; b1 = d1; b2 = d2; b3 = d3; bv = dv;
                        }
                    }
                    // bias grads: warp-reduce (all 32 lanes take part), then
                    // one shared-memory atomic per warp.
                    #pragma unroll
                    for (int off = 16; off; off >>= 1) {
                        b0 += __shfl_down_sync(~0u, b0, off);
                        b1 += __shfl_down_sync(~0u, b1, off);
                        b2 += __shfl_down_sync(~0u, b2, off);
                        b3 += __shfl_down_sync(~0u, b3, off);
                        bv += __shfl_down_sync(~0u, bv, off);
                    }
                    if ((tid & 31) == 0) {
                        atomicAdd(&s_bias[0], b0);
                        atomicAdd(&s_bias[1], b1);
                        atomicAdd(&s_bias[2], b2);
                        atomicAdd(&s_bias[3], b3);
                        atomicAdd(&s_bias[4], bv);
                    }
                }
                __syncthreads();
                // ---------------- B: parameter-gradient accumulation --------
                // thread owns a pair of hidden units (j, j+32), so each staged
                // (obs, dlogits, dvalue) load feeds 20 accumulators - halves the
                // shared-memory instruction count, the phase's bottleneck.
                {
                    const int jp = tid >> LNC;                  // hidden-unit pair 0..31
                    const int c = tid & (NCHUNK - 1);           // sample chunk
                    const int j0 = jp, j1 = jp + 32;
                    const float4 w1a = *(float4*)&sp[OFF_W1 + j0 * 4];
                    const float4 w1b = *(float4*)&sp[OFF_W1 + j1 * 4];
                    const float4 w2a = *(float4*)&sp[OFF_W2 + j0 * 4];
                    const float4 w2b = *(float4*)&sp[OFF_W2 + j1 * 4];
                    const float b1a = s_b1w3[j0 * 2], w3a = s_b1w3[j0 * 2 + 1];
                    const float b1b = s_b1w3[j1 * 2], w3b = s_b1w3[j1 * 2 + 1];
                    float aW1a[4] = {0, 0, 0, 0}, aW2a[4] = {0, 0, 0, 0};
                    float aW1b[4] = {0, 0, 0, 0}, aW2b[4] = {0, 0, 0, 0};
                    float aB1a = 0, aW3a = 0, aB1b = 0, aW3b = 0;
                    // samples interleaved across chunk-lanes so a warp's
                    // concurrent smem reads hit distinct banks (no conflicts)
                    #pragma unroll 4
                    for (int i = 0; i < SPC; i++) {
                        const int s = c + i * NCHUNK;
                        float4 o = *(float4*)&s_obs[s * 4];
                        float4 d = *(float4*)&s_dl[s * 4];
                        float dvv = s_dv[s];
                        float pa = b1a + w1a.x * o.x + w1a.y * o.y + w1a.z * o.z + w1a.w * o.w;
                        float pb = b1b + w1b.x * o.x + w1b.y * o.y + w1b.z * o.z + w1b.w * o.w;
                        float ha = ftanh(pa);
                        float hb = ftanh(pb);
                        float dha = (1.0f - ha * ha) * (d.x * w2a.x + d.y * w2a.y + d.z * w2a.z + d.w * w2a.w + dvv * w3a);
                        float dhb = (1.0f - hb * hb) * (d.x * w2b.x + d.y * w2b.y + d.z * w2b.z + d.w * w2b.w + dvv * w3b);
                        aW1a[0] += dha * o.x; aW1a[1] += dha * o.y;
                        aW1a[2] += dha * o.z; aW1a[3] += dha * o.w;
                        aW1b[0] += dhb * o.x; aW1b[1] += dhb * o.y;
                        aW1b[2] += dhb * o.z; aW1b[3] += dhb * o.w;
                        aB1a += dha; aB1b += dhb;
                        aW2a[0] += d.x * ha; aW2a[1] += d.y * ha;
                        aW2a[2] += d.z * ha; aW2a[3] += d.w * ha;
                        aW2b[0] += d.x * hb; aW2b[1] += d.y * hb;
                        aW2b[2] += d.z * hb; aW2b[3] += d.w * hb;
                        aW3a += dvv * ha; aW3b += dvv * hb;
                    }
                    // reduce across chunk threads (same pair: tid = jp*NCHUNK + c)
                    #pragma unroll
                    for (int off = 1; off < NCHUNK; off <<= 1) {
                        #pragma unroll
                        for (int q = 0; q < 4; q++) {
                            aW1a[q] += __shfl_xor_sync(~0u, aW1a[q], off);
                            aW2a[q] += __shfl_xor_sync(~0u, aW2a[q], off);
                            aW1b[q] += __shfl_xor_sync(~0u, aW1b[q], off);
                            aW2b[q] += __shfl_xor_sync(~0u, aW2b[q], off);
                        }
                        aB1a += __shfl_xor_sync(~0u, aB1a, off);
                        aW3a += __shfl_xor_sync(~0u, aW3a, off);
                        aB1b += __shfl_xor_sync(~0u, aB1b, off);
                        aW3b += __shfl_xor_sync(~0u, aW3b, off);
                    }
                    if (c == 0) {
                        *(float4*)&s_acc[OFF_W1 + j0 * 4] = make_float4(aW1a[0], aW1a[1], aW1a[2], aW1a[3]);
                        s_acc[OFF_B1 + j0] = aB1a;
                        *(float4*)&s_acc[OFF_W2 + j0 * 4] = make_float4(aW2a[0], aW2a[1], aW2a[2], aW2a[3]);
                        s_acc[OFF_W3 + j0] = aW3a;
                        *(float4*)&s_acc[OFF_W1 + j1 * 4] = make_float4(aW1b[0], aW1b[1], aW1b[2], aW1b[3]);
                        s_acc[OFF_B1 + j1] = aB1b;
                        *(float4*)&s_acc[OFF_W2 + j1 * 4] = make_float4(aW2b[0], aW2b[1], aW2b[2], aW2b[3]);
                        s_acc[OFF_W3 + j1] = aW3b;
                    }
                }
                __syncthreads();
                if (tid < 5) s_acc[OFF_B2 + tid] = s_bias[tid];   // B writes [0,640); biases here
                __syncthreads();
                {   // one striped red.add per parameter per block
                    float* dst = &gp[((parity * NSLOT) + (bid & (NSLOT - 1))) * NPAD];
                    for (int p = tid; p < NPARAM; p += TPB) atomicAdd(&dst[p], s_acc[p]);
                }
                GSYNC();
                // ---------------- Adam (redundant, block-local) -------------
                {
                    if (tid < 3 && mb == 0) stats[tid] = 0.0f;   // reset for next iter
                    // all threads: fold slots into the shared grad copy
                    float gpart = 0.0f;
                    {
                        const float4* srcv = (const float4*)&gp[parity * NSLOT * NPAD];
                        for (int p4 = tid * 4; p4 < NPARAM; p4 += TPB * 4) {
                            float4 g = srcv[p4 >> 2];
                            #pragma unroll
                            for (int sl = 1; sl < NSLOT; sl++) {
                                float4 t = srcv[(sl * NPAD + p4) >> 2];
                                g.x += t.x; g.y += t.y; g.z += t.z; g.w += t.w;
                            }
                            *(float4*)&s_acc[p4] = g;
                            gpart += g.x * g.x + g.y * g.y + g.z * g.z + g.w * g.w;
                        }
                    }
                    // grad-norm square: warp reduce in registers + one block sync
                    #pragma unroll
                    for (int off = 16; off; off >>= 1) gpart += __shfl_down_sync(~0u, gpart, off);
                    if ((tid & 31) == 0) s_red[tid >> 5] = gpart;
                    __syncthreads();
                    float gsum = 0.0f;
                    for (int w = 0; w < TPB / 32; w++) gsum += s_red[w];
                    float coef = fminf(1.0f, MAXGN / (sqrtf(gsum) + 1e-6f));
                    const int tstep = it * (EPOCHS * MINIBATCHES) + mb + 1;
                    const float bc1 = 1.0f - exp2f((float)tstep * (-0.1520030930450f));   // 0.9^t
                    const float bc2 = 1.0f - exp2f((float)tstep * (-0.00144418255002f));  // 0.999^t
                    const float lr_mc = LR / bc1, inv_bc2 = 1.0f / bc2;
                    float* zp = &gp[((parity * NSLOT) + (bid & (NSLOT - 1))) * NPAD];
                    for (int p = tid; p < NPARAM; p += TPB) {
                        float g = s_acc[p] * coef;
                        float m_ = 0.9f * sm[p] + 0.1f * g;
                        float v_ = 0.999f * sv[p] + 0.001f * g * g;
                        sp[p] -= lr_mc * m_ / (sqrtf(v_ * inv_bc2) + 1e-8f);
                        sm[p] = m_; sv[p] = v_;
                        zp[p] = 0.0f;
                    }
                    if (tid < 5) s_bias[tid] = 0.0f;   // consumed by next minibatch's A
                }
                for (int j = tid; j < 64; j += TPB) {
                    s_b1w3[j * 2] = sp[OFF_B1 + j];
                    s_b1w3[j * 2 + 1] = sp[OFF_W3 + j];
                }
                __syncthreads();
            }
        }
    }
}

// ------------------------------ host glue ----------------------------------
template <int NB, int TPB>
static void* kfn() { return (void*)ppo_fused<NB, TPB>; }

void run_fused(at::Tensor ws, at::Tensor curve, int64_t iters, int64_t seed) {
    TORCH_CHECK(ws.is_cuda() && curve.is_cuda(), "tensors must be CUDA");
    TORCH_CHECK(iters > 0, "iters must be > 0");
    char* base = (char*)ws.data_ptr();
    auto carve = [&](size_t bytes) { char* p = base; base += (bytes + 255) & ~255ULL; return p; };
    float* samp_buf  = (float*) carve(sizeof(float) * (size_t)BATCH * 8);
    float* gp       = (float*) carve(sizeof(float) * 2 * NSLOT * NPAD);
    float* stats    = (float*) carve(sizeof(float) * 16);

    // candidate configs in preference order
    struct Cfg { void* fn; int nb; int tpb; };
    static Cfg cfgs[4] = {
        {kfn<256, 128>(), 256, 128},
        {kfn<128, 256>(), 128, 256},
        {kfn<512, 128>(), 512, 128},
        {kfn<256, 256>(), 256, 256},
    };
    static int cfg = -2;
    if (cfg == -2) {
        const char* env = getenv("PPO_CFG");
        cfg = env ? atoi(env) : -1;
        if (cfg >= 0 && cfg <= 3) {
            printf("[ppo] config forced: %d (%d x %d)\n", cfg, cfgs[cfg].nb, cfgs[cfg].tpb);
        } else {
            int dev = 0; cudaGetDevice(&dev);
            cudaDeviceProp prop; cudaGetDeviceProperties(&prop, dev);
            cfg = -1;
            for (int i = 0; i < 4; i++) {
                int occ = 0;
                cudaOccupancyMaxActiveBlocksPerMultiprocessor(&occ, cfgs[i].fn, cfgs[i].tpb, 0);
                if (occ > 0 && occ * prop.multiProcessorCount >= cfgs[i].nb) { cfg = i; break; }
            }
            TORCH_CHECK(cfg >= 0, "no cooperative configuration fits this GPU");
        }
    }

    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    float* curve_ptr = (float*)curve.data_ptr();
    int it32 = (int)iters;
    unsigned long long s64 = (unsigned long long)seed;
    void* args[] = {&samp_buf,
                    &gp, &stats, &curve_ptr, &it32, (void*)&s64};
    cudaError_t err = cudaLaunchCooperativeKernel(cfgs[cfg].fn, dim3(cfgs[cfg].nb),
                                                  dim3(cfgs[cfg].tpb), args, 0, stream);
    TORCH_CHECK(err == cudaSuccess, "cooperative launch failed: ", cudaGetErrorString(err));
}
"""

_WS_BYTES = (
    4 * 131072 * 4      # obs
    + 4 * 131072        # act
    + 4 * 131072        # logp
    + 8 * 131072        # vr
    + 4 * 131072 * 2    # adv + ret
    + 4 * 2 * 4 * 645   # gp
    + 4 * 16            # stats
    + 4096              # alignment slack
)

_mod = None


def _get_mod():
    global _mod
    if _mod is None:
        _mod = load_inline(
            name="grid_ppo_fused_v6",
            cpp_sources="void run_fused(at::Tensor ws, at::Tensor curve, int64_t iters, int64_t seed);",
            cuda_sources=_CUDA_SRC,
            functions=["run_fused"],
            extra_cuda_cflags=["-O3", "-lineinfo"],
            verbose=False,
        )
    return _mod


def train(total_env_steps: int, seed: int) -> list[float]:
    mod = _get_mod()
    dev = torch.device("cuda:0")
    iters = max(1, int(total_env_steps) // (32 * 4096))
    ws = torch.empty(_WS_BYTES, dtype=torch.uint8, device=dev)
    curve = torch.empty(iters, dtype=torch.float32, device=dev)
    mod.run_fused(ws, curve, iters, int(seed))
    return curve.cpu().tolist()

20260716_233457_kinetic-claude_kinetic-0715_1m__01_rl_grid_ppo