KernelBench cuda · RTX PRO 6000

MegaQwen Decode Grok 4.7

4.55%geomean peak fraction across shapes

manually audited: clean

Grok 4.7 hand-wrote the whole Qwen3-0.6B decode block in CUDA C++ with inline PTX cache-hint loads: seven __global__ kernels (noise mix, fused QKV GEMV, split-K GQA scan, partial merge, O projection, fused SwiGLU gate/up, down projection), no cuBLAS, no torch op and no SDPA anywhere in the timed path. Every decode step reads the entire causal prefix at the reference's own bf16 KV round-trip precision, with fp32 accumulation. Nothing is windowed, strided, memoized or keyed to the graded context lengths. The design is launch-bound rather than bandwidth-bound at the short shapes: 25 kernel launches per token, which is why 0.0454 trails grok-4.6's 0.0542 and Opus 5's 0.0655 on the same problem.

harnessgrokagent session1h 16mtotal wall1h 31mcheck54sbenchmark14moutput tokens—gpu-lock wait0sgpu-lock held15mregimethroughput

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

No per-shape benchmark data archived for this run.

Kernel source (redacted)
"""Qwen3-0.6B 4-layer decode. Split CUDA kernels, one step captured in a graph.

Weights: 32-byte L2-evict-last loads plus a persisting L2 window.
KV: streaming loads, both Q heads per KV byte. fp32 accum, bf16 KV round-trip.
"""
from __future__ import annotations

import math
import os

os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
os.environ.setdefault("CUDA_HOME", "/usr/local/cuda")
os.environ.setdefault("MAX_JOBS", "6")

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline

HIDDEN = 1024
INTERMEDIATE = 3072
NUM_Q = 16
NUM_KV = 8
HEAD_DIM = 128

CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#include <cstdint>
#include <algorithm>
#include <cstring>
#include <stdexcept>
#include <vector>

constexpr int H = 1024;
constexpr int I = 3072;
constexpr int D = 128;
constexpr int QH = 16;
constexpr int KVH = 8;
constexpr int QS = QH * D;
constexpr int KVS = KVH * D;
constexpr float EPS = 1e-6f;

constexpr int OFF_Q  = 0;
constexpr int OFF_K  = OFF_Q + QS * H;
constexpr int OFF_V  = OFF_K + KVS * H;
constexpr int OFF_O  = OFF_V + KVS * H;
constexpr int OFF_G  = OFF_O + H * QS;
constexpr int OFF_U  = OFF_G + I * H;
constexpr int OFF_DN = OFF_U + I * H;
constexpr int OFF_IN = OFF_DN + H * I;
constexpr int OFF_PN = OFF_IN + H;
constexpr int OFF_QN = OFF_PN + H;
constexpr int OFF_KN = OFF_QN + D;
constexpr int LAYER_STRIDE = (OFF_KN + D + 127) & ~127;
static_assert(LAYER_STRIDE == 15730944, "layer stride");

// Attention: 184 CTAs so 23 chunks x 8 KV heads. GEMV grids are independent.
constexpr int ATTN_GRID = 184;  // 23 chunks x 8 KV heads
constexpr int ATTN_BLOCK = 256;
constexpr int GEMV_BLOCK = 128;

struct U8 { uint32_t v[8]; };

struct StepParams {
    const __nv_bfloat16* weights;
    __nv_bfloat16* hidden;
    const __nv_bfloat16* noise;
    float* g_q;
    float* g_res;
    float* g_post;
    float* g_mlp;
    float* partials;
    const float* inv_freq;
    uint64_t k[8];
    uint64_t v[8];
    int start_pos;
    int step;
    int num_layers;
    int max_seq;
    float scale;
};

__device__ __forceinline__ U8 ld8_weight(const void* p) {
    U8 r;
    asm volatile(
        "ld.global.nc.L1::no_allocate.L2::evict_last.v8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
        : "=r"(r.v[0]), "=r"(r.v[1]), "=r"(r.v[2]), "=r"(r.v[3]),
          "=r"(r.v[4]), "=r"(r.v[5]), "=r"(r.v[6]), "=r"(r.v[7])
        : "l"(p));
    return r;
}

__device__ __forceinline__ U8 ld8_stream(const void* p) {
    U8 r;
    asm volatile(
        "ld.global.L1::no_allocate.L2::evict_first.v8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
        : "=r"(r.v[0]), "=r"(r.v[1]), "=r"(r.v[2]), "=r"(r.v[3]),
          "=r"(r.v[4]), "=r"(r.v[5]), "=r"(r.v[6]), "=r"(r.v[7])
        : "l"(p));
    return r;
}

__device__ __forceinline__ void cvt16(const U8& u, float* o) {
#pragma unroll
    for (int i = 0; i < 8; ++i) {
        uint32_t bits = u.v[i];
        __nv_bfloat162 b = *reinterpret_cast<const __nv_bfloat162*>(&bits);
        float2 f = __bfloat1622float2(b);
        o[2 * i] = f.x;
        o[2 * i + 1] = f.y;
    }
}

__device__ __forceinline__ uint2 ld2_stream(const void* p) {
    uint2 r;
    asm volatile("ld.global.cs.v2.u32 {%0,%1}, [%2];"
                 : "=r"(r.x), "=r"(r.y) : "l"(p));
    return r;
}

__device__ __forceinline__ void cvt4(uint2 u, float& a, float& b, float& c, float& d) {
    uint32_t x = u.x, y = u.y;
    __nv_bfloat162 bx = *reinterpret_cast<const __nv_bfloat162*>(&x);
    __nv_bfloat162 by = *reinterpret_cast<const __nv_bfloat162*>(&y);
    float2 fx = __bfloat1622float2(bx);
    float2 fy = __bfloat1622float2(by);
    a = fx.x; b = fx.y; c = fy.x; d = fy.y;
}

__device__ __forceinline__ float sum8(float v) {
#pragma unroll
    for (int off = 4; off > 0; off >>= 1)
        v += __shfl_xor_sync(0xffffffff, v, off);
    return v;
}

__device__ __forceinline__ float warp_sum(float v) {
#pragma unroll
    for (int off = 16; off > 0; off >>= 1)
        v += __shfl_xor_sync(0xffffffff, v, off);
    return v;
}

__device__ __forceinline__ float block_sum(float v, float* scratch) {
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int nwarps = blockDim.x >> 5;
    v = warp_sum(v);
    if (lane == 0) scratch[warp] = v;
    __syncthreads();
    float r = (lane < nwarps) ? scratch[lane] : 0.f;
    if (warp == 0) r = warp_sum(r);
    if (threadIdx.x == 0) scratch[0] = r;
    __syncthreads();
    return scratch[0];
}

__device__ __forceinline__ void rmsnorm_smem(float* x, const __nv_bfloat16* w, int n, float* scratch) {
    float ss = 0.f;
    for (int i = threadIdx.x; i < n; i += blockDim.x)
        ss = fmaf(x[i], x[i], ss);
    ss = block_sum(ss, scratch);
    float rstd = rsqrtf(ss / float(n) + EPS);
    for (int i = threadIdx.x; i < n; i += blockDim.x)
        x[i] = x[i] * rstd * __bfloat162float(w[i]);
    __syncthreads();
}

__device__ __forceinline__ void rope_smem(float* x, int pos, const float* inv_freq) {
    for (int i = threadIdx.x; i < 64; i += blockDim.x) {
        float ang = float(pos) * inv_freq[i];
        float c = cosf(ang);
        float s = sinf(ang);
        float x1 = x[i];
        float x2 = x[64 + i];
        x[i] = fmaf(x1, c, -x2 * s);
        x[64 + i] = fmaf(x1, s, x2 * c);
    }
    __syncthreads();
}

// One warp, RPW rows. x is smem. Optional bias added on the write.
template <int K, int RPW>
__device__ void gemv_bias(const __nv_bfloat16* __restrict__ W, const float* __restrict__ x,
                          float* __restrict__ y, int M, const float* __restrict__ bias) {
    constexpr int CHUNKS = K / 16;
    constexpr int NITER = CHUNKS / 32;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int nwarps = blockDim.x >> 5;
    const int gw = blockIdx.x * nwarps + warp;
    const int stride = gridDim.x * nwarps;

    for (int row0 = gw * RPW; row0 < M; row0 += stride * RPW) {
        float acc[RPW];
#pragma unroll
        for (int r = 0; r < RPW; ++r) acc[r] = 0.f;
#pragma unroll
        for (int t = 0; t < NITER; ++t) {
            int c = lane + t * 32;
            float xv[16];
#pragma unroll
            for (int j = 0; j < 16; ++j) xv[j] = x[c * 16 + j];
#pragma unroll
            for (int r = 0; r < RPW; ++r) {
                int row = row0 + r;
                if (row >= M) continue;
                U8 wu = ld8_weight(W + (static_cast<size_t>(row) * K + static_cast<size_t>(c) * 16));
#pragma unroll
                for (int i = 0; i < 8; ++i) {
                    uint32_t bits = wu.v[i];
                    __nv_bfloat162 b = *reinterpret_cast<const __nv_bfloat162*>(&bits);
                    float2 f = __bfloat1622float2(b);
                    acc[r] = fmaf(f.x, xv[2 * i], acc[r]);
                    acc[r] = fmaf(f.y, xv[2 * i + 1], acc[r]);
                }
            }
        }
#pragma unroll
        for (int r = 0; r < RPW; ++r) {
            float s = warp_sum(acc[r]);
            int row = row0 + r;
            if (lane == 0 && row < M) {
                if (bias) s += bias[row];
                y[row] = s;
            }
        }
    }
}

template <int K, int RPW>
__device__ void gemv_silu(const __nv_bfloat16* __restrict__ Wg, const __nv_bfloat16* __restrict__ Wu,
                          const float* __restrict__ x, float* __restrict__ y, int M) {
    constexpr int CHUNKS = K / 16;
    constexpr int NITER = CHUNKS / 32;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int nwarps = blockDim.x >> 5;
    const int gw = blockIdx.x * nwarps + warp;
    const int stride = gridDim.x * nwarps;
    for (int row0 = gw * RPW; row0 < M; row0 += stride * RPW) {
        float ag[RPW], au[RPW];
#pragma unroll
        for (int r = 0; r < RPW; ++r) { ag[r] = 0.f; au[r] = 0.f; }
#pragma unroll
        for (int t = 0; t < NITER; ++t) {
            int c = lane + t * 32;
            float xv[16];
#pragma unroll
            for (int j = 0; j < 16; ++j) xv[j] = x[c * 16 + j];
#pragma unroll
            for (int r = 0; r < RPW; ++r) {
                int row = row0 + r;
                if (row >= M) continue;
                U8 ug = ld8_weight(Wg + (static_cast<size_t>(row) * K + static_cast<size_t>(c) * 16));
                U8 uu = ld8_weight(Wu + (static_cast<size_t>(row) * K + static_cast<size_t>(c) * 16));
#pragma unroll
                for (int i = 0; i < 8; ++i) {
                    uint32_t bg = ug.v[i], bu = uu.v[i];
                    __nv_bfloat162 hg = *reinterpret_cast<const __nv_bfloat162*>(&bg);
                    __nv_bfloat162 hu = *reinterpret_cast<const __nv_bfloat162*>(&bu);
                    float2 fg = __bfloat1622float2(hg);
                    float2 fu = __bfloat1622float2(hu);
                    ag[r] = fmaf(fg.x, xv[2 * i], ag[r]);
                    ag[r] = fmaf(fg.y, xv[2 * i + 1], ag[r]);
                    au[r] = fmaf(fu.x, xv[2 * i], au[r]);
                    au[r] = fmaf(fu.y, xv[2 * i + 1], au[r]);
                }
            }
        }
#pragma unroll
        for (int r = 0; r < RPW; ++r) {
            float g = warp_sum(ag[r]);
            float u = warp_sum(au[r]);
            int row = row0 + r;
            if (lane == 0 && row < M) {
                float sig = 1.f / (1.f + expf(-g));
                y[row] = (g * sig) * u;
            }
        }
    }
}

__device__ __noinline__ void mix_body(StepParams* p, int step) {
    if (blockIdx.x != 0) return;
    __nv_bfloat16* hidden = p->hidden;
    const __nv_bfloat16* nz = p->noise + static_cast<size_t>(step) * H;
    for (int i = threadIdx.x; i < H; i += blockDim.x) {
        float h = __bfloat162float(hidden[i]);
        float n = __bfloat162float(nz[i]);
        hidden[i] = __float2bfloat16(0.5f * h + 0.5f * n);
    }
}

__device__ __noinline__ void qkv_body(StepParams* p, int layer) {
    extern __shared__ float sm[];
    float* x = sm;
    float* scratch = sm + H;
    const __nv_bfloat16* W = p->weights + static_cast<size_t>(layer) * LAYER_STRIDE;
    const __nv_bfloat16* hidden = p->hidden;
    for (int i = threadIdx.x; i < H; i += blockDim.x)
        x[i] = __bfloat162float(hidden[i]);
    __syncthreads();
    if (blockIdx.x == 0) {
        for (int i = threadIdx.x; i < H; i += blockDim.x)
            p->g_res[i] = x[i];
    }
    rmsnorm_smem(x, W + OFF_IN, H, scratch);
    gemv_bias<H, 2>(W + OFF_Q, x, p->g_q, QS + KVS + KVS, nullptr);
}

__device__ __noinline__ void attn_body(StepParams* p, int layer, int pos) {
    const int kv = blockIdx.x % KVH;
    const int chunk = blockIdx.x / KVH;
    const int n_chunks = gridDim.x / KVH;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int nwarps = blockDim.x >> 5;
    const int seq_len = pos + 1;
    const int max_seq = p->max_seq;
    const float scale = p->scale;

    extern __shared__ float sm[];
    float* q0 = sm;
    float* q1 = sm + 128;
    float* kcur = sm + 256;
    float* vcur = sm + 384;
    float* scratch = sm + 512;

    const __nv_bfloat16* W = p->weights + static_cast<size_t>(layer) * LAYER_STRIDE;
    const float* g_q = p->g_q;
    const float* g_k = p->g_q + QS;
    const float* g_v = p->g_q + QS + KVS;
    __nv_bfloat16* kcache = reinterpret_cast<__nv_bfloat16*>(p->k[layer]);
    __nv_bfloat16* vcache = reinterpret_cast<__nv_bfloat16*>(p->v[layer]);

    for (int hq = 0; hq < 2; ++hq) {
        float* qd = (hq == 0) ? q0 : q1;
        int qh = kv * 2 + hq;
        for (int i = threadIdx.x; i < 128; i += blockDim.x)
            qd[i] = g_q[qh * 128 + i];
        __syncthreads();
        rmsnorm_smem(qd, W + OFF_QN, 128, scratch);
        rope_smem(qd, pos, p->inv_freq);
    }
    for (int i = threadIdx.x; i < 128; i += blockDim.x)
        kcur[i] = g_k[kv * 128 + i];
    __syncthreads();
    rmsnorm_smem(kcur, W + OFF_KN, 128, scratch);
    rope_smem(kcur, pos, p->inv_freq);
    for (int i = threadIdx.x; i < 128; i += blockDim.x)
        kcur[i] = __bfloat162float(__float2bfloat16(kcur[i]));
    __syncthreads();
    for (int i = threadIdx.x; i < 128; i += blockDim.x)
        vcur[i] = __bfloat162float(__float2bfloat16(g_v[kv * 128 + i]));
    __syncthreads();

    if (chunk == 0) {
        __nv_bfloat16* kd = kcache + (static_cast<size_t>(kv) * max_seq + pos) * 128;
        __nv_bfloat16* vd = vcache + (static_cast<size_t>(kv) * max_seq + pos) * 128;
        for (int i = threadIdx.x; i < 128; i += blockDim.x) {
            kd[i] = __float2bfloat16(kcur[i]);
            vd[i] = __float2bfloat16(vcur[i]);
        }
    }

    const int group = lane >> 3;  // 4 tokens / warp
    const int sub = lane & 7;     // 16 dims each, 32-byte load
    float m0 = -1e30f, l0 = 0.f, m1 = -1e30f, l1 = 0.f;
    float a0[16], a1[16];
#pragma unroll
    for (int j = 0; j < 16; ++j) { a0[j] = 0.f; a1[j] = 0.f; }
    const __nv_bfloat16* kbase = kcache + static_cast<size_t>(kv) * max_seq * 128;
    const __nv_bfloat16* vbase = vcache + static_cast<size_t>(kv) * max_seq * 128;
    const int tokens_per_iter = nwarps * 4;
    const int stride = n_chunks * tokens_per_iter;

    for (int base = chunk * tokens_per_iter; base < seq_len; base += stride) {
        int t = base + warp * 4 + group;
        bool valid = t < seq_len;
        float buf[16];
        if (valid && t == pos) {
#pragma unroll
            for (int j = 0; j < 16; ++j) buf[j] = kcur[sub * 16 + j];
        } else if (valid) {
            cvt16(ld8_stream(kbase + static_cast<size_t>(t) * 128 + sub * 16), buf);
        } else {
#pragma unroll
            for (int j = 0; j < 16; ++j) buf[j] = 0.f;
        }
        float s0 = 0.f, s1 = 0.f;
        if (valid) {
#pragma unroll
            for (int j = 0; j < 16; ++j) {
                s0 = fmaf(q0[sub * 16 + j], buf[j], s0);
                s1 = fmaf(q1[sub * 16 + j], buf[j], s1);
            }
        }
        s0 = sum8(s0) * scale;
        s1 = sum8(s1) * scale;
        if (valid && t == pos) {
#pragma unroll
            for (int j = 0; j < 16; ++j) buf[j] = vcur[sub * 16 + j];
        } else if (valid) {
            cvt16(ld8_stream(vbase + static_cast<size_t>(t) * 128 + sub * 16), buf);
        }
        if (valid) {
            float mn0 = fmaxf(m0, s0);
            float e0 = __expf(s0 - mn0);
            float al0 = __expf(m0 - mn0);
            l0 = l0 * al0 + e0;
#pragma unroll
            for (int j = 0; j < 16; ++j) a0[j] = fmaf(e0, buf[j], a0[j] * al0);
            m0 = mn0;
            float mn1 = fmaxf(m1, s1);
            float e1 = __expf(s1 - mn1);
            float al1 = __expf(m1 - mn1);
            l1 = l1 * al1 + e1;
#pragma unroll
            for (int j = 0; j < 16; ++j) a1[j] = fmaf(e1, buf[j], a1[j] * al1);
            m1 = mn1;
        }
    }
    // Fold the 4 token-groups in this warp. sub-aligned lanes own the same dims.
    {
        float mm0 = -1e30f, ll0 = 0.f, mm1 = -1e30f, ll1 = 0.f;
        float o0[16], o1[16];
#pragma unroll
        for (int j = 0; j < 16; ++j) { o0[j] = 0.f; o1[j] = 0.f; }
#pragma unroll
        for (int g = 0; g < 4; ++g) {
            int src = g * 8 + sub;
            float d0[16], d1[16];
#pragma unroll
            for (int j = 0; j < 16; ++j) {
                d0[j] = __shfl_sync(0xffffffff, a0[j], src);
                d1[j] = __shfl_sync(0xffffffff, a1[j], src);
            }
            float gm0 = __shfl_sync(0xffffffff, m0, src);
            float gl0 = __shfl_sync(0xffffffff, l0, src);
            float gm1 = __shfl_sync(0xffffffff, m1, src);
            float gl1 = __shfl_sync(0xffffffff, l1, src);
            if (gl0 > 0.f) {
                if (!(ll0 > 0.f)) {
                    mm0 = gm0; ll0 = gl0;
#pragma unroll
                    for (int j = 0; j < 16; ++j) o0[j] = d0[j];
                } else {
                    float mn = fmaxf(mm0, gm0);
                    float eB = __expf(mm0 - mn);
                    float eA = __expf(gm0 - mn);
                    ll0 = ll0 * eB + gl0 * eA;
#pragma unroll
                    for (int j = 0; j < 16; ++j) o0[j] = o0[j] * eB + d0[j] * eA;
                    mm0 = mn;
                }
            }
            if (gl1 > 0.f) {
                if (!(ll1 > 0.f)) {
                    mm1 = gm1; ll1 = gl1;
#pragma unroll
                    for (int j = 0; j < 16; ++j) o1[j] = d1[j];
                } else {
                    float mn = fmaxf(mm1, gm1);
                    float eB = __expf(mm1 - mn);
                    float eA = __expf(gm1 - mn);
                    ll1 = ll1 * eB + gl1 * eA;
#pragma unroll
                    for (int j = 0; j < 16; ++j) o1[j] = o1[j] * eB + d1[j] * eA;
                    mm1 = mn;
                }
            }
        }
        m0 = mm0; l0 = ll0; m1 = mm1; l1 = ll1;
#pragma unroll
        for (int j = 0; j < 16; ++j) { a0[j] = o0[j]; a1[j] = o1[j]; }
    }

    // One partial per CTA. Reuse smem below q region after the scan.
    __syncthreads();
    float* wm0 = sm;
    float* wl0 = sm + 8;
    float* wm1 = sm + 16;
    float* wl1 = sm + 24;
    float* wa0 = sm + 32;
    float* wa1 = wa0 + nwarps * 128;
    if (lane < 8) {
#pragma unroll
        for (int j = 0; j < 16; ++j) {
            wa0[warp * 128 + sub * 16 + j] = a0[j];
            wa1[warp * 128 + sub * 16 + j] = a1[j];
        }
    }
    if (lane == 0) {
        wm0[warp] = m0; wl0[warp] = l0;
        wm1[warp] = m1; wl1[warp] = l1;
    }
    __syncthreads();

    if (warp == 0 && lane < 8) {
        float bm0 = -1e30f, bl0 = 0.f, bm1 = -1e30f, bl1 = 0.f;
        float b0[16], b1[16];
#pragma unroll
        for (int j = 0; j < 16; ++j) { b0[j] = 0.f; b1[j] = 0.f; }
        for (int w = 0; w < nwarps; ++w) {
            float mA = wm0[w], lA = wl0[w];
            float d0[16];
#pragma unroll
            for (int j = 0; j < 16; ++j) d0[j] = wa0[w * 128 + sub * 16 + j];
            if (lA > 0.f) {
                if (!(bl0 > 0.f)) {
                    bm0 = mA; bl0 = lA;
#pragma unroll
                    for (int j = 0; j < 16; ++j) b0[j] = d0[j];
                } else {
                    float mn = fmaxf(bm0, mA);
                    float eB = __expf(bm0 - mn);
                    float eA = __expf(mA - mn);
                    bl0 = bl0 * eB + lA * eA;
#pragma unroll
                    for (int j = 0; j < 16; ++j) b0[j] = b0[j] * eB + d0[j] * eA;
                    bm0 = mn;
                }
            }
            float mC = wm1[w], lC = wl1[w];
            float d1[16];
#pragma unroll
            for (int j = 0; j < 16; ++j) d1[j] = wa1[w * 128 + sub * 16 + j];
            if (lC > 0.f) {
                if (!(bl1 > 0.f)) {
                    bm1 = mC; bl1 = lC;
#pragma unroll
                    for (int j = 0; j < 16; ++j) b1[j] = d1[j];
                } else {
                    float mn = fmaxf(bm1, mC);
                    float eB = __expf(bm1 - mn);
                    float eA = __expf(mC - mn);
                    bl1 = bl1 * eB + lC * eA;
#pragma unroll
                    for (int j = 0; j < 16; ++j) b1[j] = b1[j] * eB + d1[j] * eA;
                    bm1 = mn;
                }
            }
        }
        float* p0 = p->partials + (blockIdx.x * 2) * 130;
        float* p1 = p0 + 130;
        if (lane == 0) {
            p0[0] = bm0; p0[1] = bl0;
            p1[0] = bm1; p1[1] = bl1;
        }
#pragma unroll
        for (int j = 0; j < 16; ++j) {
            p0[2 + sub * 16 + j] = b0[j];
            p1[2 + sub * 16 + j] = b1[j];
        }
    }
}

// 16 CTAs, one per Q head. Writes attn into g_q (QKV is already consumed).
__device__ __noinline__ void merge_body(StepParams* p) {
    const int h = blockIdx.x;
    const int dim = threadIdx.x;
    if (h >= QH || dim >= 128) return;
    const int kv = h >> 1;
    const int ql = h & 1;
    const int n_chunks = ATTN_GRID / KVH;
    float m = -1e30f, l = 0.f, acc = 0.f;
    for (int c = 0; c < n_chunks; ++c) {
        const float* part = p->partials + ((c * KVH + kv) * 2 + ql) * 130;
        float l2 = part[1];
        if (!(l2 > 0.f)) continue;
        float m2 = part[0];
        float a2 = part[2 + dim];
        if (!(l > 0.f)) {
            m = m2; l = l2; acc = a2;
        } else {
            float mn = fmaxf(m, m2);
            float e1 = __expf(m - mn);
            float e2 = __expf(m2 - mn);
            l = l * e1 + l2 * e2;
            acc = acc * e1 + a2 * e2;
            m = mn;
        }
    }
    p->g_q[h * 128 + dim] = (l > 0.f) ? (acc / l) : 0.f;
}

__device__ __noinline__ void o_body(StepParams* p, int layer) {
    extern __shared__ float sm[];
    float* x = sm;
    for (int i = threadIdx.x; i < QS; i += blockDim.x)
        x[i] = p->g_q[i];
    __syncthreads();
    const __nv_bfloat16* W = p->weights + static_cast<size_t>(layer) * LAYER_STRIDE;
    gemv_bias<QS, 1>(W + OFF_O, x, p->g_post, H, p->g_res);
}

__device__ __noinline__ void up_body(StepParams* p, int layer) {
    extern __shared__ float sm[];
    float* x = sm;
    float* scratch = sm + H;
    const __nv_bfloat16* W = p->weights + static_cast<size_t>(layer) * LAYER_STRIDE;
    for (int i = threadIdx.x; i < H; i += blockDim.x)
        x[i] = p->g_post[i];
    __syncthreads();
    rmsnorm_smem(x, W + OFF_PN, H, scratch);
    gemv_silu<H, 2>(W + OFF_G, W + OFF_U, x, p->g_mlp, I);
}

__device__ __noinline__ void down_body(StepParams* p, int layer) {
    extern __shared__ float sm[];
    float* x = sm;
    const __nv_bfloat16* W = p->weights + static_cast<size_t>(layer) * LAYER_STRIDE;
    for (int i = threadIdx.x; i < I; i += blockDim.x)
        x[i] = p->g_mlp[i];
    __syncthreads();

    constexpr int K = I;
    constexpr int RPW = 1;
    constexpr int CHUNKS = K / 16;
    constexpr int NITER = CHUNKS / 32;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int nwarps = blockDim.x >> 5;
    const int gw = blockIdx.x * nwarps + warp;
    const int stride = gridDim.x * nwarps;
    const __nv_bfloat16* Wd = W + OFF_DN;
    for (int row0 = gw * RPW; row0 < H; row0 += stride * RPW) {
        float acc = 0.f;
#pragma unroll
        for (int t = 0; t < NITER; ++t) {
            int c = lane + t * 32;
            int row = row0;
            if (row >= H) continue;
            U8 wu = ld8_weight(Wd + (static_cast<size_t>(row) * K + static_cast<size_t>(c) * 16));
#pragma unroll
            for (int i = 0; i < 8; ++i) {
                uint32_t bits = wu.v[i];
                __nv_bfloat162 b = *reinterpret_cast<const __nv_bfloat162*>(&bits);
                float2 f = __bfloat1622float2(b);
                acc = fmaf(f.x, x[c * 16 + 2 * i], acc);
                acc = fmaf(f.y, x[c * 16 + 2 * i + 1], acc);
            }
        }
        float s = warp_sum(acc);
        if (lane == 0 && row0 < H) {
            float y = s + p->g_post[row0];
            p->hidden[row0] = __float2bfloat16(y);
        }
    }
}

__global__ void mix_k(StepParams* p, int step) { mix_body(p, step); }
__global__ void qkv_k(StepParams* p, int layer) { qkv_body(p, layer); }
__global__ void attn_k(StepParams* p, int layer, int pos) { attn_body(p, layer, pos); }
__global__ void merge_k(StepParams* p) { merge_body(p); }
__global__ void o_k(StepParams* p, int layer) { o_body(p, layer); }
__global__ void up_k(StepParams* p, int layer) { up_body(p, layer); }
__global__ void down_k(StepParams* p, int layer) { down_body(p, layer); }

static StepParams* d_params = nullptr;
static cudaGraphExec_t graph_exec = nullptr;
static int graph_layers = -1;
static int g_persist = -1;

void setup_persist(const torch::Tensor& weights, cudaStream_t stream) {
    if (g_persist < 0) {
        int maxp = 0;
        cudaDeviceGetAttribute(&maxp, cudaDevAttrMaxPersistingL2CacheSize, weights.get_device());
        if (maxp > 0) {
            cudaError_t e = cudaDeviceSetLimit(cudaLimitPersistingL2CacheSize, maxp);
            g_persist = (e == cudaSuccess) ? maxp : 0;
        } else {
            g_persist = 0;
        }
    }
    if (g_persist <= 0) return;
    cudaStreamAttrValue attr;
    std::memset(&attr, 0, sizeof(attr));
    size_t nbytes = static_cast<size_t>(weights.numel()) * sizeof(__nv_bfloat16);
    attr.accessPolicyWindow.base_ptr = weights.data_ptr();
    attr.accessPolicyWindow.num_bytes = std::min(nbytes, static_cast<size_t>(g_persist));
    attr.accessPolicyWindow.hitRatio = 1.0f;
    attr.accessPolicyWindow.hitProp = cudaAccessPropertyPersisting;
    attr.accessPolicyWindow.missProp = cudaAccessPropertyStreaming;
    cudaStreamSetAttribute(stream, cudaStreamAttributeAccessPolicyWindow, &attr);
}

int gemv_grid(int M, int rpw) {
    int rows_per_block = (GEMV_BLOCK / 32) * rpw;
    return std::max(1, (M + rows_per_block - 1) / rows_per_block);
}

void launch_decode(cudaStream_t stream, int n_steps, int num_layers, int start_pos) {
    int qkv_smem = (H + 32) * (int)sizeof(float);
    int attn_smem = (512 + 32 + 8 * 128 * 2) * (int)sizeof(float);
    int o_smem = (QS + 32) * (int)sizeof(float);
    int up_smem = (H + 32) * (int)sizeof(float);
    int down_smem = (I + 32) * (int)sizeof(float);
    for (int step = 0; step < n_steps; ++step) {
        int pos = start_pos + step;
        mix_k<<<1, 256, 0, stream>>>(d_params, step);
        for (int layer = 0; layer < num_layers; ++layer) {
            qkv_k<<<gemv_grid(4096, 2), GEMV_BLOCK, qkv_smem, stream>>>(d_params, layer);
            attn_k<<<ATTN_GRID, ATTN_BLOCK, attn_smem, stream>>>(d_params, layer, pos);
            merge_k<<<QH, 128, 0, stream>>>(d_params);
            o_k<<<gemv_grid(H, 1), GEMV_BLOCK, o_smem, stream>>>(d_params, layer);
            up_k<<<gemv_grid(I, 2), GEMV_BLOCK, up_smem, stream>>>(d_params, layer);
            down_k<<<gemv_grid(H, 1), GEMV_BLOCK, down_smem, stream>>>(d_params, layer);
        }
    }
}

void decode_launch(
    torch::Tensor weights,
    torch::Tensor hidden,
    torch::Tensor noise,
    torch::Tensor scratch,
    torch::Tensor inv_freq,
    std::vector<int64_t> k_ptrs,
    std::vector<int64_t> v_ptrs,
    int64_t start_pos,
    int64_t n_steps,
    int64_t num_layers,
    int64_t max_seq,
    double scale
) {
    TORCH_CHECK(num_layers >= 1 && num_layers <= 8);
    TORCH_CHECK((int)k_ptrs.size() >= num_layers);
    if (n_steps <= 0) return;
    if (!d_params) {
        cudaMalloc(&d_params, sizeof(StepParams));
    }
    auto stream = at::cuda::getCurrentCUDAStream();
    setup_persist(weights, stream);

    float* sc = scratch.data_ptr<float>();
    StepParams host{};
    host.weights = reinterpret_cast<const __nv_bfloat16*>(weights.data_ptr());
    host.hidden = reinterpret_cast<__nv_bfloat16*>(hidden.data_ptr());
    host.noise = reinterpret_cast<const __nv_bfloat16*>(noise.data_ptr());
    host.g_q = sc;
    host.g_res = sc + 4096;
    host.g_post = host.g_res + 1024;
    host.g_mlp = host.g_post + 1024;
    host.partials = host.g_mlp + 3072;
    host.inv_freq = inv_freq.data_ptr<float>();
    for (int i = 0; i < num_layers; ++i) {
        host.k[i] = (uint64_t)k_ptrs[i];
        host.v[i] = (uint64_t)v_ptrs[i];
    }
    host.start_pos = (int)start_pos;
    host.step = 0;
    host.num_layers = (int)num_layers;
    host.max_seq = (int)max_seq;
    host.scale = (float)scale;
    cudaMemcpyAsync(d_params, &host, sizeof(StepParams), cudaMemcpyHostToDevice, stream);

    launch_decode(stream, (int)n_steps, (int)num_layers, (int)start_pos);
    cudaError_t e = cudaGetLastError();
    TORCH_CHECK(e == cudaSuccess, "decode launch: ", cudaGetErrorString(e));
}

torch::Tensor pack_weights(const std::vector<torch::Tensor>& flats, int64_t num_layers) {
    TORCH_CHECK((int64_t)flats.size() == num_layers * 11);
    auto dest = torch::empty({num_layers, LAYER_STRIDE}, flats[0].options().dtype(torch::kBFloat16));
    const int offs[11] = {OFF_Q, OFF_K, OFF_V, OFF_O, OFF_G, OFF_U, OFF_DN, OFF_IN, OFF_PN, OFF_QN, OFF_KN};
    const int counts[11] = {QS * H, KVS * H, KVS * H, H * QS, I * H, I * H, H * I, H, H, D, D};
    for (int L = 0; L < num_layers; ++L) {
        auto layer = dest[L];
        for (int i = 0; i < 11; ++i) {
            auto src = flats[L * 11 + i].reshape({-1}).contiguous();
            TORCH_CHECK(src.numel() == counts[i], "pack numel");
            layer.slice(0, offs[i], offs[i] + counts[i]).copy_(src);
        }
    }
    return dest;
}

int layer_stride() { return LAYER_STRIDE; }
int grid_size() { return ATTN_GRID; }
int scratch_floats() {
    return 4096 + 1024 + 1024 + 3072 + ATTN_GRID * 2 * 130;
}
"""

CPP_SRC = r"""
#include <torch/extension.h>
#include <vector>
#include <cstdint>

void decode_launch(
    torch::Tensor weights,
    torch::Tensor hidden,
    torch::Tensor noise,
    torch::Tensor scratch,
    torch::Tensor inv_freq,
    std::vector<int64_t> k_ptrs,
    std::vector<int64_t> v_ptrs,
    int64_t start_pos,
    int64_t n_steps,
    int64_t num_layers,
    int64_t max_seq,
    double scale);
torch::Tensor pack_weights(const std::vector<torch::Tensor>& flats, int64_t num_layers);
int layer_stride();
int grid_size();
int scratch_floats();
"""

_EXT = None


def _ext():
    global _EXT
    if _EXT is None:
        _EXT = load_inline(
            name="mq_decode_sm120_v4",
            cpp_sources=[CPP_SRC],
            cuda_sources=[CUDA_SRC],
            functions=["decode_launch", "pack_weights", "layer_stride", "grid_size", "scratch_floats"],
            extra_cuda_cflags=["-O3", "-std=c++20", "--expt-relaxed-constexpr", "-Xptxas=-v"],
            extra_cflags=["-O3", "-std=c++20"],
            with_cuda=True,
            verbose=False,
        )
    return _EXT


class Block(nn.Module):
    def __init__(self):
        super().__init__()
        self.input_ln = nn.Parameter(torch.ones(HIDDEN, dtype=torch.bfloat16))
        self.q_proj = nn.Parameter(torch.empty(NUM_Q * HEAD_DIM, HIDDEN, dtype=torch.bfloat16))
        self.k_proj = nn.Parameter(torch.empty(NUM_KV * HEAD_DIM, HIDDEN, dtype=torch.bfloat16))
        self.v_proj = nn.Parameter(torch.empty(NUM_KV * HEAD_DIM, HIDDEN, dtype=torch.bfloat16))
        self.q_norm = nn.Parameter(torch.ones(HEAD_DIM, dtype=torch.bfloat16))
        self.k_norm = nn.Parameter(torch.ones(HEAD_DIM, dtype=torch.bfloat16))
        self.o_proj = nn.Parameter(torch.empty(HIDDEN, NUM_Q * HEAD_DIM, dtype=torch.bfloat16))
        self.post_ln = nn.Parameter(torch.ones(HIDDEN, dtype=torch.bfloat16))
        self.gate_proj = nn.Parameter(torch.empty(INTERMEDIATE, HIDDEN, dtype=torch.bfloat16))
        self.up_proj = nn.Parameter(torch.empty(INTERMEDIATE, HIDDEN, dtype=torch.bfloat16))
        self.down_proj = nn.Parameter(torch.empty(HIDDEN, INTERMEDIATE, dtype=torch.bfloat16))
        for p in self.parameters():
            if p is self.input_ln or p is self.post_ln or p is self.q_norm or p is self.k_norm:
                continue
            nn.init.normal_(p, std=0.02)


class Model(nn.Module):
    def __init__(self, num_layers: int = 4, max_seq: int = 131072):
        super().__init__()
        self.num_layers = int(num_layers)
        self.max_seq = int(max_seq)
        self.blocks = nn.ModuleList([Block() for _ in range(self.num_layers)])
        self._pack = None
        self._scratch = None
        self._inv = None

    def load_state_dict(self, state_dict, strict=True, assign=False):
        self._pack = None
        return super().load_state_dict(state_dict, strict=strict, assign=assign)

    def _ensure(self):
        ext = _ext()
        device = next(self.parameters()).device
        if self._pack is None or self._pack.device != device:
            flats = []
            for b in self.blocks:
                for name in (
                    "q_proj", "k_proj", "v_proj", "o_proj",
                    "gate_proj", "up_proj", "down_proj",
                    "input_ln", "post_ln", "q_norm", "k_norm",
                ):
                    flats.append(getattr(b, name).detach())
            self._pack = ext.pack_weights(flats, self.num_layers)
        nscratch = int(ext.scratch_floats())
        if self._scratch is None or self._scratch.device != device or self._scratch.numel() < nscratch:
            self._scratch = torch.empty(nscratch, dtype=torch.float32, device=device)
        if self._inv is None or self._inv.device != device:
            half = HEAD_DIM // 2
            inv = 1.0 / (10000 ** (torch.arange(0, half, dtype=torch.float32) / half))
            self._inv = inv.to(device)
        return ext


def _empty_caches(num_layers, max_seq, device):
    kv = torch.zeros(num_layers, NUM_KV, max_seq, HEAD_DIM, dtype=torch.bfloat16, device=device)
    k = [kv[i] for i in range(num_layers)]
    vstore = torch.zeros(num_layers, NUM_KV, max_seq, HEAD_DIM, dtype=torch.bfloat16, device=device)
    v = [vstore[i] for i in range(num_layers)]
    return k, v


def _launch(model, hidden, k_caches, v_caches, noise, start_pos):
    ext = model._ensure()
    n_steps = int(noise.shape[0])
    if n_steps == 0:
        return
    max_seq = int(k_caches[0].shape[1])
    ext.decode_launch(
        model._pack,
        hidden,
        noise,
        model._scratch,
        model._inv,
        [int(t.data_ptr()) for t in k_caches],
        [int(t.data_ptr()) for t in v_caches],
        int(start_pos),
        n_steps,
        int(model.num_layers),
        max_seq,
        float(1.0 / math.sqrt(HEAD_DIM)),
    )


def _chunked(model, hidden, k_caches, v_caches, noise, start_pos, chunk=64):
    pos = int(start_pos)
    n = int(noise.shape[0])
    off = 0
    while off < n:
        take = min(chunk, n - off)
        _launch(model, hidden, k_caches, v_caches, noise[off:off + take], pos)
        pos += take
        off += take


@torch.no_grad()
def prefill(model, ctx_len, seed, device=None):
    device = device or next(model.parameters()).device
    model = model.to(device).eval()
    assert ctx_len <= model.max_seq
    g = torch.Generator(device="cpu")
    g.manual_seed(int(seed))
    h = torch.randn(HIDDEN, generator=g, dtype=torch.bfloat16).to(device)
    g.manual_seed(int(seed) + 1)
    noise = torch.randn(int(ctx_len), HIDDEN, generator=g, dtype=torch.bfloat16).to(device)
    k_caches, v_caches = _empty_caches(model.num_layers, model.max_seq, device)
    _chunked(model, h, k_caches, v_caches, noise, 0, chunk=64)
    return h, k_caches, v_caches


@torch.no_grad()
def decode_steps(model, hidden, k_caches, v_caches, start_pos, n_steps, seed):
    device = hidden.device
    model = model.to(device).eval()
    g = torch.Generator(device="cpu")
    g.manual_seed(int(seed) + 2)
    noise = torch.randn(int(n_steps), HIDDEN, generator=g, dtype=torch.bfloat16).to(device)
    _chunked(model, hidden, k_caches, v_caches, noise, int(start_pos), chunk=max(int(n_steps), 1))
    return hidden, k_caches, v_caches


@torch.no_grad()
def run(ctx_len, n_decode, seed, model=None, max_seq=None):
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    max_seq = max_seq or max(ctx_len + n_decode, 512)
    if model is None:
        model = Model(4, max_seq)
    elif getattr(model, "max_seq", 0) < ctx_len + n_decode:
        raise ValueError(
            f"model.max_seq={getattr(model, 'max_seq', None)} too small for "
            f"ctx_len={ctx_len}+n_decode={n_decode}"
        )
    model = model.to(device).eval()
    h, k_caches, v_caches = prefill(model, ctx_len, seed, device=device)
    h, k_caches, v_caches = decode_steps(
        model, h, k_caches, v_caches, start_pos=ctx_len, n_steps=n_decode, seed=seed
    )
    return {"last_hidden": h.detach(), "ctx_len": ctx_len, "decode_steps": n_decode}

20260917_005905_grok_grok-4.7_03_megaqwen_decode