"""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 #include #include #include #include namespace cg = cooperative_groups; #include #include #include #include #include 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(&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(&x); __nv_bfloat162 by = *reinterpret_cast(&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 __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(row) * K + static_cast(c) * 16)); #pragma unroll for (int i = 0; i < 8; ++i) { uint32_t bits = wu.v[i]; __nv_bfloat162 b = *reinterpret_cast(&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 __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(row) * K + static_cast(c) * 16)); U8 uu = ld8_weight(Wu + (static_cast(row) * K + static_cast(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(&bg); __nv_bfloat162 hu = *reinterpret_cast(&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(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(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(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(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(kv) * max_seq + pos) * 128; __nv_bfloat16* vd = vcache + (static_cast(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(kv) * max_seq * 128; const __nv_bfloat16* vbase = vcache + static_cast(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(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(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(layer) * LAYER_STRIDE; gemv_bias(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(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(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(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(row) * K + static_cast(c) * 16)); #pragma unroll for (int i = 0; i < 8; ++i) { uint32_t bits = wu.v[i]; __nv_bfloat162 b = *reinterpret_cast(&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(weights.numel()) * sizeof(__nv_bfloat16); attr.accessPolicyWindow.base_ptr = weights.data_ptr(); attr.accessPolicyWindow.num_bytes = std::min(nbytes, static_cast(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<<>>(d_params, layer); attn_k<<>>(d_params, layer, pos); merge_k<<>>(d_params); o_k<<>>(d_params, layer); up_k<<>>(d_params, layer); down_k<<>>(d_params, layer); } } } void decode_launch( torch::Tensor weights, torch::Tensor hidden, torch::Tensor noise, torch::Tensor scratch, torch::Tensor inv_freq, std::vector k_ptrs, std::vector 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(); StepParams host{}; host.weights = reinterpret_cast(weights.data_ptr()); host.hidden = reinterpret_cast<__nv_bfloat16*>(hidden.data_ptr()); host.noise = reinterpret_cast(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(); 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& 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 #include #include void decode_launch( torch::Tensor weights, torch::Tensor hidden, torch::Tensor noise, torch::Tensor scratch, torch::Tensor inv_freq, std::vector k_ptrs, std::vector 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& 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}