KernelBench cuda · B200

MegaQwen Decode Claude Opus 4.8

4.91%geomean peak fraction across shapes

manually audited: clean

harnessclaudeagent session1h 56mtotal wall2h 31mcheck10mbenchmark25moutput tokens357,034cost$40.31regimethroughput

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)
"""MegaQwen-style Qwen3-0.6B 4-layer decode, custom CUDA on RTX PRO 6000 (SM120).

Memory-bound autoregressive decode. Everything is a hand-written CUDA kernel:
  mix -> [ RMSNorm+QKV gemv -> Q/K RMSNorm+RoPE+KV store -> split-KV flash attn
           -> O proj + residual -> RMSNorm+gate/up gemv -> SwiGLU -> down + residual ] x4

Numerics mirror reference.py: fp32 math with bf16 weights upcast, K/V rounded to
bf16 in cache, block output rounded to bf16. A host loop launches the per-step
kernels with `pos` as a plain argument (no device counters).
"""
from __future__ import annotations

import glob
import os
import site

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


def _setup_cuda_toolkit() -> str | None:
    """Point CUDA_HOME at the venv's cu13 toolkit.

    torch here is a cu130 build, but the box only ships a system CUDA 12.8
    toolkit (major mismatch -> cpp_extension hard-fails). The matching cu13
    nvcc/cccl/nvvm/crt are pip-installed into nvidia/cu13; wire CUDA_HOME to it.
    """
    cu = None
    try:
        import nvidia
        for base in list(getattr(nvidia, "__path__", [])):
            cand = os.path.join(base, "cu13")
            if os.path.exists(os.path.join(cand, "bin", "nvcc")):
                cu = cand
                break
    except Exception:
        pass
    if cu is None:
        roots = list(site.getsitepackages()) if hasattr(site, "getsitepackages") else []
        roots += [p for p in os.sys.path if p.endswith("site-packages")]
        for sp in roots:
            cand = os.path.join(sp, "nvidia", "cu13")
            if os.path.exists(os.path.join(cand, "bin", "nvcc")):
                cu = cand
                break
    if cu is None:
        return None
    lib, lib64 = os.path.join(cu, "lib"), os.path.join(cu, "lib64")
    if os.path.isdir(lib) and not os.path.exists(lib64):
        try:
            os.symlink("lib", lib64)
        except OSError:
            pass
    # The pip runtime ships only versioned libs (libcudart.so.13); the linker's
    # `-lcudart` needs an unversioned .so symlink.
    for versioned in glob.glob(os.path.join(lib, "lib*.so.*")):
        stem = versioned
        while "." in os.path.basename(stem) and not stem.endswith(".so"):
            stem = stem.rsplit(".", 1)[0]
        if not stem.endswith(".so"):
            continue
        if not os.path.exists(stem):
            try:
                os.symlink(os.path.basename(versioned), stem)
            except OSError:
                pass
    os.environ["CUDA_HOME"] = cu
    os.environ["CUDA_PATH"] = cu
    os.environ["PATH"] = os.path.join(cu, "bin") + os.pathsep + os.environ.get("PATH", "")
    # torch.cpp_extension caches CUDA_HOME at import (resolved from `which nvcc`,
    # which is the GPU-lock wrapper on PATH -> wrong/12.8). Force it to cu13 so
    # load_inline uses cu13/bin/nvcc directly (matches torch's cu130 build).
    import torch.utils.cpp_extension as _cpp
    _cpp.CUDA_HOME = cu
    return cu

OP_TYPE = "megaqwen_decode"

HIDDEN = 1024
INTERMEDIATE = 3072
NUM_Q = 16
NUM_KV = 8
HEAD_DIM = 128
NUM_LAYERS = 4
EPS = 1e-6
MAXSPLIT = 2048

# ----------------------------------------------------------------------------
# CUDA
# ----------------------------------------------------------------------------
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>

typedef __nv_bfloat16 bf16;

#define H 1024
#define II 3072
#define NQ 16
#define NKV 8
#define DH 128
#define HALF 64
#define NL 4
#define EPSILON 1e-6f
#define SCALE 0.08838834764831845f   // 1/sqrt(128)
#define MAXSPLIT 2048
#define PW 132   // padded partial stride: [m, l, _, _, acc[128]]
#define PACC 4   // acc offset within a partial slot (16-byte aligned)

__device__ __forceinline__ float b2f(bf16 x){ return __bfloat162float(x); }
__device__ __forceinline__ bf16 f2b(float x){ return __float2bfloat16(x); }
__device__ __forceinline__ float to_f(float x){ return x; }
__device__ __forceinline__ float to_f(bf16 x){ return __bfloat162float(x); }
__device__ __forceinline__ void store_out(float* p, int i, float v){ p[i]=v; }
__device__ __forceinline__ void store_out(bf16*  p, int i, float v){ p[i]=f2b(v); }

// ---- mix: blk0 = round_bf16(0.5*randn + 0.5*hidden) ----
__global__ void mix_kernel(bf16* __restrict__ blk0, const bf16* __restrict__ rnd,
                           const bf16* __restrict__ hid){
    int i = blockIdx.x*blockDim.x + threadIdx.x;
    if(i<H) blk0[i] = f2b(0.5f*b2f(rnd[i]) + 0.5f*b2f(hid[i]));
}

// ---- RMSNorm over [H] -> fp32 h_normed.  1 block, 256 threads. ----
template<typename T>
__global__ void rmsnorm_kernel(float* __restrict__ out, const T* __restrict__ x_in,
                               const bf16* __restrict__ nw){
    __shared__ float red[8];
    int tid=threadIdx.x, lane=tid&31, warp=tid>>5;
    float v0=to_f(x_in[tid]), v1=to_f(x_in[tid+256]), v2=to_f(x_in[tid+512]), v3=to_f(x_in[tid+768]);
    float p=v0*v0+v1*v1+v2*v2+v3*v3;
    for(int o=16;o>0;o>>=1) p+=__shfl_down_sync(0xffffffffu,p,o);
    if(lane==0) red[warp]=p;
    __syncthreads();
    if(warp==0){ float t=(lane<8)?red[lane]:0.f; for(int o=4;o>0;o>>=1) t+=__shfl_down_sync(0xffffffffu,t,o); if(lane==0) red[0]=t; }
    __syncthreads();
    float inv=rsqrtf(red[0]/(float)H + EPSILON);
    out[tid]     = v0*inv*b2f(nw[tid]);
    out[tid+256] = v1*inv*b2f(nw[tid+256]);
    out[tid+512] = v2*inv*b2f(nw[tid+512]);
    out[tid+768] = v3*inv*b2f(nw[tid+768]);
}

// ---- pure QKV gemv over pre-normed h [H].  grid 512, block 256 (1 warp/row). ----
__global__ void qkv_kernel(
        float* __restrict__ q_raw, float* __restrict__ k_raw, float* __restrict__ v_raw,
        const bf16* __restrict__ Wq, const bf16* __restrict__ Wk, const bf16* __restrict__ Wv,
        const float* __restrict__ h){
    __shared__ float sh[H];
    int tid=threadIdx.x, lane=tid&31, warp=tid>>5;
    for(int i=tid;i<H;i+=256) sh[i]=h[i];
    __syncthreads();
    int grow=blockIdx.x*8+warp;
    if(grow>=4096) return;
    const bf16* wr; float* outp; int oi;
    if(grow<2048){ wr=Wq+(size_t)grow*H; outp=q_raw; oi=grow; }
    else if(grow<3072){ int r=grow-2048; wr=Wk+(size_t)r*H; outp=k_raw; oi=r; }
    else { int r=grow-3072; wr=Wv+(size_t)r*H; outp=v_raw; oi=r; }
    float acc=0.f;
    for(int i=lane*8;i<H;i+=256){
        float4 wv=*reinterpret_cast<const float4*>(wr+i);
        const bf16* wb=reinterpret_cast<const bf16*>(&wv);
        #pragma unroll
        for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t];
    }
    for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o);
    if(lane==0) outp[oi]=acc;
}

// ---- pure gate/up gemv over pre-normed h [H].  grid 768 (6144 rows / 8). ----
// gate_up layout: [0,3072)=gate, [3072,6144)=up
__global__ void gateup_kernel(
        float* __restrict__ gate_up,
        const bf16* __restrict__ Wg, const bf16* __restrict__ Wu,
        const float* __restrict__ h){
    __shared__ float sh[H];
    int tid=threadIdx.x, lane=tid&31, warp=tid>>5;
    for(int i=tid;i<H;i+=256) sh[i]=h[i];
    __syncthreads();
    int grow=blockIdx.x*8+warp;
    if(grow>=6144) return;
    const bf16* wr; int oi;
    if(grow<3072){ wr=Wg+(size_t)grow*H; oi=grow; }
    else { int r=grow-3072; wr=Wu+(size_t)r*H; oi=3072+r; }
    float acc=0.f;
    for(int i=lane*8;i<H;i+=256){
        float4 wv=*reinterpret_cast<const float4*>(wr+i);
        const bf16* wb=reinterpret_cast<const bf16*>(&wv);
        #pragma unroll
        for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t];
    }
    for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o);
    if(lane==0) gate_up[oi]=acc;
}

// ---- plain gemv with residual add.  out[row] = resid[row] + sum_j W[row,j]*x[j] ----
// K-split gemv with residual add.  block=256 (8 warps); SK warps share a row so
// M small (=H) still yields 8/SK rows/block * M/(8/SK) blocks = M*SK warps.
template<typename OUT, typename RES, int SK>
__global__ void gemv_resid_kernel(OUT* __restrict__ out, const bf16* __restrict__ W,
        const float* __restrict__ x_in, const RES* __restrict__ resid, int M, int K){
    extern __shared__ float sh[];        // sh[0..K-1]=x; sh[K..]=partials[ROWS*SK]
    const int ROWS = 8/SK;
    int tid=threadIdx.x, lane=tid&31, warp=tid>>5;
    for(int i=tid;i<K;i+=blockDim.x) sh[i]=x_in[i];
    __syncthreads();
    int r = warp / SK, ks = warp % SK;
    int grow = blockIdx.x*ROWS + r;
    if(grow>=M) return;
    const bf16* wr=W+(size_t)grow*K;
    int kslice=K/SK, kstart=ks*kslice, kend=kstart+kslice;
    float acc=0.f;
    for(int i=kstart+lane*8;i<kend;i+=256){
        float4 wv=*reinterpret_cast<const float4*>(wr+i);
        const bf16* wb=reinterpret_cast<const bf16*>(&wv);
        #pragma unroll
        for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t];
    }
    for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o);
    float* part=sh+K;
    if(lane==0) part[r*SK+ks]=acc;
    __syncthreads();
    if(ks==0 && lane==0){
        float s=0.f;
        #pragma unroll
        for(int j=0;j<SK;j++) s+=part[r*SK+j];
        store_out(out, grow, to_f(resid[grow]) + s);
    }
}

// ---- Q/K per-head RMSNorm + RoPE ; store K,V to cache at pos ----
// grid = 24 (0..15 q heads, 16..23 kv heads), block 128.
__global__ void qknr_kernel(
        float* __restrict__ q_out, bf16* __restrict__ kc_layer, bf16* __restrict__ vc_layer,
        const float* __restrict__ q_raw, const float* __restrict__ k_raw, const float* __restrict__ v_raw,
        const bf16* __restrict__ q_norm, const bf16* __restrict__ k_norm,
        int pos, int max_seq){
    int h=blockIdx.x, d=threadIdx.x;
    __shared__ float sh[DH];
    __shared__ float cs[HALF], sn[HALF];
    __shared__ float red[4];
    if(d<HALF){ float invf=powf(10000.f, -(float)d/(float)HALF); float ang=(float)pos*invf; cs[d]=cosf(ang); sn[d]=sinf(ang); }
    bool isq = (h<16);
    const float* raw = isq ? q_raw : k_raw;
    const bf16* nw = isq ? q_norm : k_norm;
    int hh = isq ? h : (h-16);
    float val, ss;
    if(isq || h<16+NKV){
        val = raw[hh*DH + d];
        ss = val*val;
        int lane=d&31, warp=d>>5;
        for(int o=16;o>0;o>>=1) ss+=__shfl_down_sync(0xffffffffu,ss,o);
        if(lane==0) red[warp]=ss;
        __syncthreads();
        if(warp==0){ float t=(lane<4)?red[lane]:0.f; for(int o=2;o>0;o>>=1) t+=__shfl_down_sync(0xffffffffu,t,o); if(lane==0) red[0]=t; }
        __syncthreads();
        float inv=rsqrtf(red[0]/(float)DH + EPSILON);
        sh[d]=val*inv*b2f(nw[d]);
        __syncthreads();
        float outv;
        if(d<HALF) outv = sh[d]*cs[d] - sh[d+HALF]*sn[d];
        else       outv = sh[d-HALF]*sn[d-HALF] + sh[d]*cs[d-HALF];
        if(isq){
            q_out[hh*DH+d]=outv;
        } else {
            kc_layer[(size_t)hh*max_seq*DH + (size_t)pos*DH + d] = f2b(outv);
            // v: no norm, no rope
            vc_layer[(size_t)hh*max_seq*DH + (size_t)pos*DH + d] = f2b(v_raw[hh*DH+d]);
        }
    }
}

// ---- split-KV flash attention partials.  block per (kvhead, split), W warps.
// grid = NKV*nsplit, block = W*32.  The W warps each scan a slice of the split's
// positions, then merge (shared) into one partial -> high occupancy, small nsplit.
__global__ void attn_partial_kernel(
        float* __restrict__ partials, const float* __restrict__ q_out,
        const bf16* __restrict__ kc_layer, const bf16* __restrict__ vc_layer,
        int seqlen, int nsplit, int max_seq, int W){
    int blk = blockIdx.x;
    int kvhead = blk % NKV;
    int split  = blk / NKV;
    int warp = threadIdx.x >> 5, lane = threadIdx.x & 31, dd0 = lane*4;
    int qh0=2*kvhead, qh1=qh0+1;
    int psplit = (seqlen + nsplit - 1)/nsplit;
    int p0 = split*psplit, p1 = min(p0+psplit, seqlen);
    // this warp's contiguous slice of the split
    int np = p1 - p0; if(np<0) np=0;
    int wc = (np + W - 1)/W;
    int wp0 = p0 + warp*wc, wp1 = min(wp0+wc, p1);

    float4 q0 = *reinterpret_cast<const float4*>(q_out + qh0*DH + dd0);
    float4 q1 = *reinterpret_cast<const float4*>(q_out + qh1*DH + dd0);
    float m0=-1e30f,l0=0.f,a0x=0,a0y=0,a0z=0,a0w=0;
    float m1=-1e30f,l1=0.f,a1x=0,a1y=0,a1z=0,a1w=0;
    const bf16* kh=kc_layer + (size_t)kvhead*max_seq*DH + dd0;
    const bf16* vh=vc_layer + (size_t)kvhead*max_seq*DH + dd0;
    #pragma unroll 4
    for(int p=wp0;p<wp1;p++){
        float2 kf = *reinterpret_cast<const float2*>(kh + (size_t)p*DH);
        const bf16* kb=reinterpret_cast<const bf16*>(&kf);
        float k0=b2f(kb[0]),k1=b2f(kb[1]),k2=b2f(kb[2]),k3=b2f(kb[3]);
        float d0d=q0.x*k0+q0.y*k1+q0.z*k2+q0.w*k3;
        float d1d=q1.x*k0+q1.y*k1+q1.z*k2+q1.w*k3;
        for(int o=16;o>0;o>>=1){ d0d+=__shfl_xor_sync(0xffffffffu,d0d,o); d1d+=__shfl_xor_sync(0xffffffffu,d1d,o); }
        float s0=d0d*SCALE, s1=d1d*SCALE;
        float2 vf = *reinterpret_cast<const float2*>(vh + (size_t)p*DH);
        const bf16* vb=reinterpret_cast<const bf16*>(&vf);
        float v0=b2f(vb[0]),v1=b2f(vb[1]),v2=b2f(vb[2]),v3=b2f(vb[3]);
        float m0n=fmaxf(m0,s0), c0=__expf(m0-m0n), pe0=__expf(s0-m0n);
        l0=l0*c0+pe0; a0x=a0x*c0+pe0*v0; a0y=a0y*c0+pe0*v1; a0z=a0z*c0+pe0*v2; a0w=a0w*c0+pe0*v3; m0=m0n;
        float m1n=fmaxf(m1,s1), c1=__expf(m1-m1n), pe1=__expf(s1-m1n);
        l1=l1*c1+pe1; a1x=a1x*c1+pe1*v0; a1y=a1y*c1+pe1*v1; a1z=a1z*c1+pe1*v2; a1w=a1w*c1+pe1*v3; m1=m1n;
    }
    // shared: m0s[W], l0s[W], m1s[W], l1s[W], acc0[W*DH], acc1[W*DH]
    extern __shared__ float sh[];
    float* m0s=sh; float* l0s=m0s+W; float* m1s=l0s+W; float* l1s=m1s+W;
    float* acc0=l1s+W; float* acc1=acc0+W*DH;
    if(lane==0){ m0s[warp]=m0; l0s[warp]=l0; m1s[warp]=m1; l1s[warp]=l1; }
    *reinterpret_cast<float4*>(acc0 + warp*DH + dd0) = make_float4(a0x,a0y,a0z,a0w);
    *reinterpret_cast<float4*>(acc1 + warp*DH + dd0) = make_float4(a1x,a1y,a1z,a1w);
    __syncthreads();
    // merge W warps -> one partial (threads 0..127 handle dims)
    if(threadIdx.x < DH){
        int d=threadIdx.x;
        float M0=-1e30f, M1=-1e30f;
        for(int w=0;w<W;w++){ M0=fmaxf(M0,m0s[w]); M1=fmaxf(M1,m1s[w]); }
        float acc0d=0.f, acc1d=0.f, L0=0.f, L1=0.f;
        for(int w=0;w<W;w++){
            float f0=__expf(m0s[w]-M0); acc0d += acc0[w*DH+d]*f0;
            float f1=__expf(m1s[w]-M1); acc1d += acc1[w*DH+d]*f1;
            if(d==0){ L0 += l0s[w]*f0; L1 += l1s[w]*f1; }
        }
        int b0=(qh0*nsplit+split)*PW, b1=(qh1*nsplit+split)*PW;
        partials[b0+PACC+d]=acc0d; partials[b1+PACC+d]=acc1d;
        if(d==0){ partials[b0]=M0; partials[b0+1]=L0; partials[b1]=M1; partials[b1+1]=L1; }
    }
}

// ---- combine split partials -> attn_out (fully parallel) ----
// grid = 16 (one per q head), block 128 (one per dim).
__global__ void attn_combine_kernel(float* __restrict__ attn_out,
        const float* __restrict__ partials, int nsplit){
    int qh=blockIdx.x, tid=threadIdx.x;   // 128 threads
    extern __shared__ float sh[];         // sm[0..nsplit-1], sl[nsplit..2nsplit-1]
    float* sm = sh;
    float* sl = sh + nsplit;
    __shared__ float red[128];
    const float* base = partials + (size_t)qh*nsplit*PW;
    // cooperative coalesced load of m,l
    for(int s=tid;s<nsplit;s+=128){ sm[s]=base[s*PW]; sl[s]=base[s*PW+1]; }
    __syncthreads();
    // parallel max
    float mx=-1e30f;
    for(int s=tid;s<nsplit;s+=128) mx=fmaxf(mx,sm[s]);
    red[tid]=mx; __syncthreads();
    for(int o=64;o>0;o>>=1){ if(tid<o) red[tid]=fmaxf(red[tid],red[tid+o]); __syncthreads(); }
    float M=red[0]; __syncthreads();
    // factors (overwrite sm) + partial l-sum
    float lp=0.f;
    for(int s=tid;s<nsplit;s+=128){ float f=__expf(sm[s]-M); sm[s]=f; lp+=sl[s]*f; }
    red[tid]=lp; __syncthreads();
    for(int o=64;o>0;o>>=1){ if(tid<o) red[tid]+=red[tid+o]; __syncthreads(); }
    float L=red[0]; __syncthreads();
    // weighted acc for dim=tid (coalesced across threads per split)
    int d=tid;
    float accd=0.f;
    for(int s=0;s<nsplit;s++) accd += base[s*PW+PACC+d]*sm[s];
    attn_out[qh*DH+d] = accd / L;
}

// ---- SwiGLU: act = silu(gate)*up ----
__global__ void silu_kernel(float* __restrict__ act, const float* __restrict__ gate_up){
    int j=blockIdx.x*blockDim.x+threadIdx.x;
    if(j<II){ float g=gate_up[j], u=gate_up[II+j]; act[j]=(g/(1.f+__expf(-g)))*u; }
}

// ---- host driver: n_steps decode steps ----
void run_steps(
        torch::Tensor input_ln, torch::Tensor post_ln, torch::Tensor q_norm, torch::Tensor k_norm,
        torch::Tensor q_proj, torch::Tensor k_proj, torch::Tensor v_proj, torch::Tensor o_proj,
        torch::Tensor gate_proj, torch::Tensor up_proj, torch::Tensor down_proj,
        torch::Tensor k_cache_all, torch::Tensor v_cache_all,
        torch::Tensor hidden, torch::Tensor randn,
        torch::Tensor blk, torch::Tensor q_raw, torch::Tensor k_raw, torch::Tensor v_raw,
        torch::Tensor q_out, torch::Tensor partials, torch::Tensor attn_out,
        torch::Tensor h_attn, torch::Tensor gate_up, torch::Tensor act,
        torch::Tensor h_normed,
        int64_t start_pos, int64_t n_steps, int64_t max_seq, int64_t pos_per_split,
        int64_t warps_per_split){
    cudaStream_t st = at::cuda::getCurrentCUDAStream().stream();
    int W = (int)warps_per_split;
    size_t apsmem = (size_t)(4*W + 2*W*DH)*sizeof(float);
    const bf16* IN_LN=(const bf16*)input_ln.data_ptr();
    const bf16* PO_LN=(const bf16*)post_ln.data_ptr();
    const bf16* QN=(const bf16*)q_norm.data_ptr();
    const bf16* KN=(const bf16*)k_norm.data_ptr();
    const bf16* WQ=(const bf16*)q_proj.data_ptr();
    const bf16* WK=(const bf16*)k_proj.data_ptr();
    const bf16* WV=(const bf16*)v_proj.data_ptr();
    const bf16* WO=(const bf16*)o_proj.data_ptr();
    const bf16* WG=(const bf16*)gate_proj.data_ptr();
    const bf16* WU=(const bf16*)up_proj.data_ptr();
    const bf16* WD=(const bf16*)down_proj.data_ptr();
    bf16* KC=(bf16*)k_cache_all.data_ptr();
    bf16* VC=(bf16*)v_cache_all.data_ptr();
    bf16* HID=(bf16*)hidden.data_ptr();
    const bf16* RND=(const bf16*)randn.data_ptr();
    bf16* BLK=(bf16*)blk.data_ptr();
    float* QRAW=(float*)q_raw.data_ptr();
    float* KRAW=(float*)k_raw.data_ptr();
    float* VRAW=(float*)v_raw.data_ptr();
    float* QOUT=(float*)q_out.data_ptr();
    float* PART=(float*)partials.data_ptr();
    float* AOUT=(float*)attn_out.data_ptr();
    float* HATT=(float*)h_attn.data_ptr();
    float* GU=(float*)gate_up.data_ptr();
    float* ACT=(float*)act.data_ptr();
    float* HN=(float*)h_normed.data_ptr();

    const size_t WQL=(size_t)NQ*DH*H, WKVL=(size_t)NKV*DH*H, WOL=(size_t)H*NQ*DH;
    const size_t WGL=(size_t)II*H, WDL=(size_t)H*II, LNL=H, QKNL=DH;
    const size_t KCL=(size_t)NKV*max_seq*DH;

    for(int s=0;s<n_steps;s++){
        int pos=start_pos+s;
        int seqlen=pos+1;
        int nsplit=(seqlen + (int)pos_per_split - 1)/(int)pos_per_split;
        if(nsplit<1) nsplit=1; if(nsplit>MAXSPLIT) nsplit=MAXSPLIT;
        mix_kernel<<<(H+255)/256,256,0,st>>>(BLK, RND+(size_t)s*H, HID);
        for(int L=0;L<NL;L++){
            const bf16* blkL = BLK + (size_t)L*H;
            bf16* blkOut = (L<NL-1) ? (BLK+(size_t)(L+1)*H) : HID;
            bf16* kcL = KC + (size_t)L*KCL;
            bf16* vcL = VC + (size_t)L*KCL;
            rmsnorm_kernel<bf16><<<1,256,0,st>>>(HN, blkL, IN_LN+L*LNL);
            qkv_kernel<<<512,256,0,st>>>(QRAW,KRAW,VRAW, WQ+L*WQL,WK+L*WKVL,WV+L*WKVL, HN);
            qknr_kernel<<<24,128,0,st>>>(QOUT, kcL, vcL, QRAW,KRAW,VRAW, QN+L*QKNL, KN+L*QKNL, pos, (int)max_seq);
            attn_partial_kernel<<<NKV*nsplit, W*32, apsmem, st>>>(PART, QOUT, kcL, vcL, seqlen, nsplit, (int)max_seq, W);
            attn_combine_kernel<<<NQ,128,2*nsplit*sizeof(float),st>>>(AOUT, PART, nsplit);
            gemv_resid_kernel<float,bf16,4><<<H/2,256,(NQ*DH+8)*sizeof(float),st>>>(HATT, WO+L*WOL, AOUT, blkL, H, NQ*DH);
            rmsnorm_kernel<float><<<1,256,0,st>>>(HN, HATT, PO_LN+L*LNL);
            gateup_kernel<<<768,256,0,st>>>(GU, WG+L*WGL, WU+L*WGL, HN);
            silu_kernel<<<(II+255)/256,256,0,st>>>(ACT, GU);
            gemv_resid_kernel<bf16,float,4><<<H/2,256,(II+8)*sizeof(float),st>>>(blkOut, WD+L*WDL, ACT, HATT, H, II);
        }
    }
}
"""

CPP_SRC = r"""
#include <torch/extension.h>
void run_steps(
    torch::Tensor input_ln, torch::Tensor post_ln, torch::Tensor q_norm, torch::Tensor k_norm,
    torch::Tensor q_proj, torch::Tensor k_proj, torch::Tensor v_proj, torch::Tensor o_proj,
    torch::Tensor gate_proj, torch::Tensor up_proj, torch::Tensor down_proj,
    torch::Tensor k_cache_all, torch::Tensor v_cache_all,
    torch::Tensor hidden, torch::Tensor randn,
    torch::Tensor blk, torch::Tensor q_raw, torch::Tensor k_raw, torch::Tensor v_raw,
    torch::Tensor q_out, torch::Tensor partials, torch::Tensor attn_out,
    torch::Tensor h_attn, torch::Tensor gate_up, torch::Tensor act,
    torch::Tensor h_normed,
    int64_t start_pos, int64_t n_steps, int64_t max_seq, int64_t pos_per_split,
    int64_t warps_per_split);
"""

_EXT = None


def _ext():
    global _EXT
    if _EXT is None:
        cu = _setup_cuda_toolkit()
        inc = [f"-I{os.path.join(cu, 'include')}"] if cu else []
        cap = torch.cuda.get_device_capability()
        arch = f"{cap[0]}{cap[1]}"  # e.g. "100" (B200/sm_100), "120" (RTX PRO 6000)
        gencode = [f"-gencode=arch=compute_{arch},code=sm_{arch}",
                   f"-gencode=arch=compute_{arch},code=compute_{arch}"]
        _EXT = load_inline(
            name=f"megaqwen_decode_ext_sm{arch}",
            cpp_sources=CPP_SRC,
            cuda_sources=CUDA_SRC,
            functions=["run_steps"],
            extra_cflags=["-O3"] + inc,
            extra_cuda_cflags=["-O3", "--expt-relaxed-constexpr"] + gencode + inc,
            verbose=False,
        )
    return _EXT


# ----------------------------------------------------------------------------
# Model (same state_dict as reference.Model)
# ----------------------------------------------------------------------------
class Block(nn.Module):
    def __init__(self):
        super().__init__()
        H, I, D = HIDDEN, INTERMEDIATE, HEAD_DIM
        self.input_ln = nn.Parameter(torch.ones(H, dtype=torch.bfloat16))
        self.q_proj = nn.Parameter(torch.empty(NUM_Q * D, H, dtype=torch.bfloat16))
        self.k_proj = nn.Parameter(torch.empty(NUM_KV * D, H, dtype=torch.bfloat16))
        self.v_proj = nn.Parameter(torch.empty(NUM_KV * D, H, dtype=torch.bfloat16))
        self.q_norm = nn.Parameter(torch.ones(D, dtype=torch.bfloat16))
        self.k_norm = nn.Parameter(torch.ones(D, dtype=torch.bfloat16))
        self.o_proj = nn.Parameter(torch.empty(H, NUM_Q * D, dtype=torch.bfloat16))
        self.post_ln = nn.Parameter(torch.ones(H, dtype=torch.bfloat16))
        self.gate_proj = nn.Parameter(torch.empty(I, H, dtype=torch.bfloat16))
        self.up_proj = nn.Parameter(torch.empty(I, H, dtype=torch.bfloat16))
        self.down_proj = nn.Parameter(torch.empty(H, I, dtype=torch.bfloat16))
        for p in self.parameters():
            if p.dim() >= 2:
                nn.init.normal_(p, std=0.02)


class Model(nn.Module):
    def __init__(self, num_layers: int = NUM_LAYERS, max_seq: int = 131072):
        super().__init__()
        self.num_layers = num_layers
        self.max_seq = max_seq
        self.blocks = nn.ModuleList([Block() for _ in range(num_layers)])
        self._prepped = False

    # ---- one-time: stack weights + allocate scratch on device ----
    def _prep(self):
        if self._prepped:
            return
        dev = next(self.parameters()).device
        bl = self.blocks

        def stack(name):
            return torch.stack([getattr(b, name) for b in bl]).contiguous()

        self.w = {n: stack(n) for n in
                  ["input_ln", "post_ln", "q_norm", "k_norm", "q_proj", "k_proj",
                   "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]}
        f32 = dict(dtype=torch.float32, device=dev)
        bf = dict(dtype=torch.bfloat16, device=dev)
        self._hidden = torch.zeros(HIDDEN, **bf)
        self._blk = torch.zeros(NUM_LAYERS, HIDDEN, **bf)
        self._q_raw = torch.zeros(NUM_Q * HEAD_DIM, **f32)
        self._k_raw = torch.zeros(NUM_KV * HEAD_DIM, **f32)
        self._v_raw = torch.zeros(NUM_KV * HEAD_DIM, **f32)
        self._q_out = torch.zeros(NUM_Q * HEAD_DIM, **f32)
        self._partials = torch.zeros(NUM_Q * MAXSPLIT * 132, **f32)
        self._attn_out = torch.zeros(NUM_Q * HEAD_DIM, **f32)
        self._h_attn = torch.zeros(HIDDEN, **f32)
        self._gate_up = torch.zeros(2 * INTERMEDIATE, **f32)
        self._act = torch.zeros(INTERMEDIATE, **f32)
        self._hn = torch.zeros(HIDDEN, **f32)
        # Persistent KV cache + randn buffer (stable pointers for CUDA graphs).
        self._maxdec = 64
        self._randn_buf = torch.zeros(self._maxdec, HIDDEN, **bf)
        self._kc = torch.zeros(NUM_LAYERS, NUM_KV, self.max_seq, HEAD_DIM, **bf)
        self._vc = torch.zeros(NUM_LAYERS, NUM_KV, self.max_seq, HEAD_DIM, **bf)
        self._graphs = {}
        self._prepped = True

    def _run_eager(self, k_cache_all, v_cache_all, randn, start_pos, n_steps, pos_per_split, warps):
        w = self.w
        _ext().run_steps(
            w["input_ln"], w["post_ln"], w["q_norm"], w["k_norm"],
            w["q_proj"], w["k_proj"], w["v_proj"], w["o_proj"],
            w["gate_proj"], w["up_proj"], w["down_proj"],
            k_cache_all, v_cache_all, self._hidden, randn,
            self._blk, self._q_raw, self._k_raw, self._v_raw,
            self._q_out, self._partials, self._attn_out,
            self._h_attn, self._gate_up, self._act, self._hn,
            int(start_pos), int(n_steps), int(self.max_seq), int(pos_per_split), int(warps))

    def _decode(self, hidden, randn, start_pos, n_steps, pos_per_split, warps):
        """Graph-replayed decode over persistent _kc/_vc. randn is [n_steps, H]."""
        n = int(n_steps)
        # `hidden` may alias self._hidden (prefill returns it); snapshot it so the
        # eager warmup below cannot clobber the input we reset to before replay.
        hidden_in = hidden.clone()
        if not _USE_GRAPH or n > self._maxdec:
            self._hidden.copy_(hidden_in)
            self._run_eager(self._kc, self._vc, randn, start_pos, n, pos_per_split, warps)
            return
        key = (int(start_pos), n, int(pos_per_split), int(warps))
        g = self._graphs.get(key)
        if g is None:
            self._randn_buf[:n].copy_(randn)
            self._hidden.copy_(hidden_in)
            # Warmup (eager) loads kernels & does lazy init before capture.
            self._run_eager(self._kc, self._vc, self._randn_buf, start_pos, n, pos_per_split, warps)
            torch.cuda.synchronize()
            g = torch.cuda.CUDAGraph()
            with torch.cuda.graph(g):
                self._run_eager(self._kc, self._vc, self._randn_buf, start_pos, n, pos_per_split, warps)
            self._graphs[key] = g
        self._randn_buf[:n].copy_(randn)
        self._hidden.copy_(hidden_in)
        g.replay()


_USE_GRAPH = True


def _split_cfg(ctx_len: int) -> tuple[int, int]:
    """(pos_per_split, warps_per_split). W warps/block cooperate on each split, so
    nsplit stays small (cheap combine) while occupancy = 8*nsplit*W stays high."""
    if ctx_len <= 2048:
        return max(1, ctx_len // 128), 4      # nsplit~128, W=4
    if ctx_len <= 8192:
        return max(1, ctx_len // 128), 8      # nsplit~64,  W=8
    if ctx_len <= 32768:
        return max(1, ctx_len // 256), 8      # nsplit~128, W=8
    return max(1, ctx_len // 512), 4          # 131072: nsplit~256, W=4


def _pos_per_split(ctx_len: int) -> int:
    return _split_cfg(ctx_len)[0]


def _warps_per_split(ctx_len: int) -> int:
    return _split_cfg(ctx_len)[1]


def _randn(seed: int, n: int, device) -> torch.Tensor:
    g = torch.Generator(device="cpu")
    g.manual_seed(seed)
    return torch.randn((n, HIDDEN), generator=g, dtype=torch.bfloat16).to(device)


@torch.no_grad()
def prefill(model, ctx_len: int, seed: int):
    model = model.cuda().eval()
    model._prep()
    dev = model._hidden.device
    model._kc.zero_()
    model._vc.zero_()
    g = torch.Generator(device="cpu")
    g.manual_seed(seed)
    h0 = torch.randn(HIDDEN, generator=g, dtype=torch.bfloat16).to(dev)
    model._hidden.copy_(h0)
    rnd = _randn(seed + 1, ctx_len, dev)
    model._run_eager(model._kc, model._vc, rnd, 0, ctx_len,
                     _pos_per_split(ctx_len), _warps_per_split(ctx_len))
    k_list = [model._kc[i] for i in range(NUM_LAYERS)]
    v_list = [model._vc[i] for i in range(NUM_LAYERS)]
    return model._hidden, k_list, v_list


@torch.no_grad()
def decode_steps(model, hidden, k_caches, v_caches, start_pos, n_steps, seed):
    model._prep()
    dev = model._hidden.device
    rnd = _randn(seed + 2, n_steps, dev)
    model._decode(hidden, rnd, start_pos, n_steps,
                  _pos_per_split(start_pos), _warps_per_split(start_pos))
    return model._hidden, k_caches, v_caches


_decode_impl = decode_steps


def run(ctx_len, decode_steps, seed, model=None, max_seq=None):
    n_decode = int(decode_steps)
    if model is None:
        ms = max_seq or max(ctx_len + n_decode, 512)
        model = Model(NUM_LAYERS, ms)
    model = model.cuda().eval()
    h, k, v = prefill(model, ctx_len, seed)
    h, k, v = _decode_impl(model, h, k, v, ctx_len, n_decode, seed)
    return {"last_hidden": h.detach(), "ctx_len": ctx_len, "decode_steps": n_decode}


def get_init_inputs():
    return [NUM_LAYERS, 131072]


def get_inputs():
    return []

20260719_103301_claude_claude-opus-4-8_03_megaqwen_decode