KernelBench cuda · RTX PRO 6000
MegaQwen Decode Claude Opus 4.8
manually audited: clean
Genuine hand-written CUDA multi-layer decode for the Qwen3-0.6B geometry: per-block RMSNorm, warp-per-row QKV GEMV, fused Q/K head-RMSNorm + RoPE + bf16 KV-cache store, split-KV online-softmax flash attention (W warps cooperating per split, shared-memory merge, parallel combine), K-split GEMV with fused residual for O/down proj, SwiGLU - a host loop launches ~11 kernels per layer per step, wrapped in a CUDA graph keyed on (start_pos, n_steps, split config) that copies the live hidden + randn into static buffers and replays over the persistent KV cache. Empirically verified to recompute on mutated KV/hidden, not replay stale outputs. cuda_language.json: framework=cuda_raw, no Triton/DSL, real __global__ CUDA - passes the CUDA-only gate. Numerics mirror reference exactly (fp32 math, bf16 rounding at cache and block boundaries; same seeded CPU-generator randn contract). 0.0491 geomean (4831 tok/s at ctx 2048 vs eager reference 206 tok/s) is an honest measurement.
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