KernelBench cuda · RTX PRO 6000
MegaQwen Decode Claude Fable 5.1
manually audited: clean
Claude Fable 5.1 through OpenRouter (or-fable harness), max effort, unlimited budget, RTX PRO 6000, 3h33m to a voluntary stop. One persistent cooperative megakernel: 188 blocks x 512 threads, every decode step and all four layers in a single cudaLaunchCooperativeKernel, five grid barriers per layer (QKV, attention, O-proj, gate-up, down), the 126 MB weight set held L2-resident with evict_last hints while the KV cache streams through a 6-stage cp.async.bulk/mbarrier ring tagged evict_first. Attention softmax uses a fixed per-layer reference max (p = exp(s - M) with M >= max|s| derived from the live q/k-norm weights plus a 2% margin), which is exact because the constant cancels in the division; an exact online-softmax fallback exists but is dead code under every grader-seeded model. 36 measured kernel builds: a cp.async L2-hint illegal instruction on sm_120 forced the bulk-copy redesign, one deadlocked build was killed, and the agent's own 8192-context spot check caught a NaN regression that was fixed before submission. Fresh h_out every call; the kernel_state cache keyed on weight data_ptr/_version stores pointers and two scalars, weights are re-read from GPU memory on every launch. Contended in-run grade 0.0638; isolated sequential regrade on the same box 2026-09-04: 0.0643 (+0.8%), 11.4k tok/s at ctx 2048 down to 655 tok/s at ctx 131072. Probe (probe.log): primed ctx 2048 matches the reference at cos 1.0000, a second call with a new seed on the same model produces a different output (cos(s1,s1b)=-0.1474) that matches at 1.0000, and overwriting the weights in place changes the output (0.0239) and matches at 1.0000. All template files byte-identical, no foreign reads, no network, no clock commands, no key strings, zero lock contention. Graded file byte-identical to the build the agent verified. OpenRouter cost $64.87 (672k output tokens, 56.3M cache reads). Second on the problem behind or-opus 0.0655.
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-geometry decode: persistent cooperative CUDA megakernel.
Design (RTX PRO 6000 Blackwell, SM120, 188 SMs, 128 MB L2, ~1.56 TB/s GDDR7):
* One cooperative kernel launch runs ALL decode steps x ALL layers: 188 blocks (one per SM) x 512 threads,
5 grid barriers per layer (QKV | attention | O-proj | gate-up | down).
* The 4-layer bf16 weight set (126 MB) stays L2-resident across steps: every weight fetch carries an
`evict_last` L2 policy, every KV-cache fetch an `evict_first` policy (measured: weights then stream at
~7 TB/s while 2 GB/step of KV pass through the L2).
* All bulk data movement is TMA-style `cp.async.bulk` into shared memory with mbarrier completion (plain
cp.async + L2::cache_hint raises an illegal-instruction fault on sm_120). Weight slices for the next phase
are issued as soon as their destination buffer is consumed; KV tiles for the next layer are prefetched into
a 6-stage ring right after a barrier arrive so at 2K context the attention loop never touches DRAM.
* Grid barrier = relaxed atomic counter + one polling thread. Phase outputs are published with *returning*
atomics so no gpu-scope fence is needed (a fence would wait for all in-flight TMA of the SM).
* Attention: split-K over 8 KV heads x ~23 position chunks (GQA: 2 Q heads share each K/V read), one position
per warp per tile. Because Q/K are RMSNormed, |score| <= scale*128*max|w_qn|*max|w_kn|; when that bound is
small (it is ~11.5 for unit norm weights) the softmax uses a fixed reference max so chunks merge with
fire-and-forget fp32 reductions. Otherwise the exact online-softmax + last-arriver combine path is used.
* RMSNorms, Q/K norm + RoPE and the new K/V cache entry are recomputed redundantly by the blocks that need
them; small activation vectors are replicated 4x to avoid 188-way hot L2 lines.
API (same as reference.py): Model, prefill, decode_steps, run.
"""
from __future__ import annotations
import math
import os
import sys
import torch
import torch.nn as nn
HIDDEN = 1024
INTERMEDIATE = 3072
NUM_Q = 16
NUM_KV = 8
HEAD_DIM = 128
NUM_LAYERS = 4
EPS = 1e-6
# ----------------------------------------------------------------------------
# CUDA source
# ----------------------------------------------------------------------------
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cstdint>
#include <vector>
typedef __nv_bfloat16 bf16;
constexpr int HID = 1024, INTER = 3072, NQ = 16, NKV = 8, HD = 128;
constexpr int QKV_ROWS = NQ * HD + 2 * NKV * HD; // 4096
constexpr int NTHR = 512, NWARP = 16;
constexpr int NSTAGE = 3; // dedicated KV ring stages (block-wide tiles)
constexpr int NSTAGE_EXT = 3; // extra stages living in W2[0, 24K) (free from B4 through P2)
constexpr int NPRE = NSTAGE + NSTAGE_EXT; // 6 stages usable for pre-filling the next layer's tiles
constexpr int NSTAGE_W1 = 0; // stages living in W1 (only while attention runs); >0 thrashes the ~2 MB of L2 left for in-flight KV
constexpr int NST = NPRE + NSTAGE_W1; // 10 stages total
constexpr int W2_QKV_ROW0 = 12; // QKV part-2 rows live at W2 rows 12.. (24 KB offset)
constexpr int NREP = 4; // replicas of the small activation vectors (hot-line relief)
constexpr int TILE_POS = 16; // positions per tile (one per warp)
constexpr int TILE_BYTES = 2 * TILE_POS * HD * 2; // 8192: [K rows | V rows]
constexpr int MAXCH = 32; // max chunks per kv head
constexpr int MAXL = 8;
constexpr int PART_STRIDE = 132; // acc[128], m, l, pad (16B aligned rows)
constexpr int POLL_T = 15 * 32; // barrier polling thread (warp 15, lane 0)
constexpr int W1_ROWS = 18; // K=1024 rows that fit W1 (36 KB)
constexpr float RMS_EPS = 1e-6f;
constexpr float ATTN_SCALE = 0.08838834764831845f;
// shared memory layout (bytes)
constexpr int SM_RING = NSTAGE * TILE_BYTES; // 24576 (also: P2 warp partials, P3 act)
constexpr int SM_W1 = 36 * 1024; // 36864: O / gate-up part 1 / down / QKV part 1
constexpr int SM_W2 = 32 * 1024; // 32768: gate-up part 2 / QKV part 2 / P5 act
constexpr int SM_ACT4 = 4096; // P1/P4 normalized input staging; P2 q/k/v scratch
constexpr int SM_RED = 128;
constexpr int SM_RES = 256;
constexpr int SM_MISC = 64;
constexpr int SM_TRIG = 512; // cos[64], sin[64]
constexpr int SM_MBAR = (2 * NST + 3) * 8; // full[NST], empty[NST], wm1, wm1b, wm2
constexpr int OFF_RING = 0;
constexpr int OFF_W1 = OFF_RING + SM_RING;
constexpr int OFF_W2 = OFF_W1 + SM_W1;
constexpr int OFF_ACT4 = OFF_W2 + SM_W2;
constexpr int OFF_RED = OFF_ACT4 + SM_ACT4;
constexpr int OFF_RES = OFF_RED + SM_RED;
constexpr int OFF_MISC = OFF_RES + SM_RES;
constexpr int OFF_TRIG = OFF_MISC + SM_MISC;
constexpr int OFF_MBAR = OFF_TRIG + SM_TRIG;
constexpr int SMEM_BYTES = OFF_MBAR + SM_MBAR;
static_assert(NWARP * 2 * PART_STRIDE * 4 <= SM_RING, "warp partials must fit in the ring region");
static_assert(NSTAGE_EXT * TILE_BYTES <= SM_W2, "extended ring stages must fit in W2");
static_assert(NSTAGE_W1 * TILE_BYTES <= SM_W1, "W1 ring stages must fit in W1");
static_assert(NQ * HD * 4 <= SM_W2, "P3 act must fit in W2");
static_assert(NSTAGE_EXT * TILE_BYTES + 4 * HID * 2 <= SM_W2, "W2: 3 ring stages + 4 QKV rows");
static_assert(SMEM_BYTES <= 99 * 1024, "smem budget");
struct Params {
const bf16* w_in_ln[MAXL];
const bf16* w_q[MAXL];
const bf16* w_k[MAXL];
const bf16* w_v[MAXL];
const bf16* w_qn[MAXL];
const bf16* w_kn[MAXL];
const bf16* w_o[MAXL];
const bf16* w_post_ln[MAXL];
const bf16* w_gate[MAXL];
const bf16* w_up[MAXL];
const bf16* w_down[MAXL];
bf16* kc[MAXL];
bf16* vc[MAXL];
long long kv_head_stride; // elements between kv heads (= max_seq * HD)
const bf16* randn; // [n_steps][HID]
bf16* h; // [NREP][HID] hidden (replicated; replica b%NREP is read by block b)
const bf16* h_in; // [HID] input hidden (read by step 0 / layer 0 instead of the replicas)
bf16* h_out; // [HID] final hidden (written at the last layer of the last step)
float* qkv_raw; // [QKV_ROWS]
float* partials; // [NKV][MAXCH][2][PART_STRIDE]
float* attn_out; // [NQ*HD]
float* h_prime; // [NREP][HID]
float* mlp; // [NREP][INTER]
unsigned* counters; // [NKV] attention last-arriver tickets
unsigned* gbar; // [1] grid barrier counter (monotonic, zeroed per launch)
float* attn_acc; // [2][NQ*HD] fixed-max softmax numerators (double-buffered by layer parity)
float* l_acc; // [2][NQ] fixed-max softmax denominators
float m_est[MAXL]; // per-layer upper bound of |score| (<=0: use exact online-softmax path)
const float* inv_freq; // [HD/2]
unsigned long long* timing; // optional per-phase globaltimer stamps (block 0), or nullptr
int start_pos;
int n_steps;
int num_layers;
};
// ---------------------------------------------------------------- helpers
__device__ __forceinline__ float bf_lo(uint32_t w) { return __uint_as_float(w << 16); }
__device__ __forceinline__ float bf_hi(uint32_t w) { return __uint_as_float(w & 0xffff0000u); }
__device__ __forceinline__ float bf2f(bf16 x) { return __bfloat162float(x); }
__device__ __forceinline__ float bfbits2f(unsigned short u) { return __uint_as_float(((uint32_t)u) << 16); }
// TMA-style bulk copies (cp.async with L2::cache_hint raises an illegal-instruction fault on sm_120;
// bulk copies with the hint work and reach ~1.5 TB/s).
__device__ __forceinline__ void mbar_init(uint32_t mb, uint32_t cnt) {
asm volatile("mbarrier.init.shared.b64 [%0], %1;" [REDACTED: IP]"r"(mb), "r"(cnt) : "memory");
}
__device__ __forceinline__ void mbar_expect_tx(uint32_t mb, uint32_t bytes) {
asm volatile("mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;" [REDACTED: IP]"r"(mb), "r"(bytes) : "memory");
}
__device__ __forceinline__ void mbar_arrive(uint32_t mb) {
asm volatile("mbarrier.arrive.shared.b64 _, [%0];" [REDACTED: IP]"r"(mb) : "memory");
}
__device__ __forceinline__ void bulk_g2s(uint32_t dst, const void* src, uint32_t bytes, uint32_t mb, uint64_t pol) {
asm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
[REDACTED: IP]"r"(dst), "l"(src), "r"(bytes), "r"(mb), "l"(pol) : "memory");
}
__device__ __forceinline__ void mbar_wait(uint32_t mb, uint32_t parity) {
uint32_t done = 0;
while (!done) {
asm volatile("{ .reg .pred p; mbarrier.try_wait.parity.shared.b64 p, [%1], %2; selp.u32 %0, 1, 0, p; }"
: "=r"(done) : "r"(mb), "r"(parity) : "memory");
}
}
__device__ __forceinline__ void fence_proxy_async() { asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); }
__device__ __forceinline__ unsigned long long gtimer() {
unsigned long long t; asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t) [REDACTED: IP] "memory"); return t;
}
// Publishing stores: returning atomics. Their completion (the returned value) proves the data is performed at L2,
// so the grid barrier can arrive with a relaxed add and no gpu-scope fence (a fence would wait for all in-flight TMA).
__device__ __forceinline__ unsigned pub_f32(float* p, float v) {
unsigned old;
asm volatile("atom.relaxed.gpu.global.exch.b32 %0, [%1], %2;" : "=r"(old) : "l"(p), "r"(__float_as_uint(v)) : "memory");
return old;
}
__device__ __forceinline__ unsigned pub_u32(unsigned* p, unsigned v) {
unsigned old;
asm volatile("atom.relaxed.gpu.global.exch.b32 %0, [%1], %2;" : "=r"(old) : "l"(p), "r"(v) : "memory");
return old;
}
__device__ __forceinline__ unsigned pub_add_f32(float* p, float v) {
float old;
asm volatile("atom.relaxed.gpu.global.add.f32 %0, [%1], %2;" : "=f"(old) : "l"(p), "f"(v) : "memory");
return __float_as_uint(old);
}
// Force the thread to wait for the returned values (otherwise nothing consumes them and the thread runs ahead).
__device__ __forceinline__ void pub_sink(unsigned v, int* misc) {
if (v == 0x7fc00001u) misc[7] = 1; // a signalling-NaN bit pattern that never occurs in practice
}
__device__ __forceinline__ unsigned ld_relaxed_gpu_u32(const unsigned* p) {
unsigned v;
asm volatile("ld.relaxed.gpu.global.u32 %0, [%1];" : "=r"(v) : "l"(p) : "memory");
return v;
}
// bf16 element `idx` of a bf16 array read through L2 as an aligned 32-bit word
__device__ __forceinline__ float ldcg_bf16_elem(const bf16* base, int idx) {
uint32_t w = __ldcg(reinterpret_cast<const unsigned int*>(base) + (idx >> 1));
return (idx & 1) ? bf_hi(w) : bf_lo(w);
}
__device__ __forceinline__ uint64_t make_policy_evict_last() {
uint64_t p; asm volatile("createpolicy.fractional.L2::evict_last.b64 %0, 1.0;" : "=l"(p)); return p;
}
__device__ __forceinline__ uint64_t make_policy_evict_first() {
uint64_t p; asm volatile("createpolicy.fractional.L2::evict_first.b64 %0, 1.0;" : "=l"(p)); return p;
}
__device__ __forceinline__ float warp_sum(float v) {
#pragma unroll
for (int o = 16; o > 0; o >>= 1) v += __shfl_xor_sync(0xffffffffu, v, o);
return v;
}
__device__ __forceinline__ void block_rows(int M, int b, int nblk, int& r0, int& n) {
int rpb = (M + nblk - 1) / nblk;
r0 = b * rpb;
int r1 = min(M, r0 + rpb);
n = r1 > r0 ? r1 - r0 : 0;
}
__device__ __forceinline__ void chunk_bounds(int P, int cidx, int nch, int& lo, int& hi) {
// P <= 2^17, cidx < 32: products fit in 32 bits
lo = (int)(((unsigned)P * (unsigned)cidx) / (unsigned)nch);
hi = (cidx == nch - 1) ? P : (int)(((unsigned)P * (unsigned)(cidx + 1)) / (unsigned)nch);
}
// dot of 8 bf16 (uint4) with 8 fp32 activations held in registers
__device__ __forceinline__ float dot8r(uint4 w, const float* a, float s) {
s = fmaf(bf_lo(w.x), a[0], s); s = fmaf(bf_hi(w.x), a[1], s);
s = fmaf(bf_lo(w.y), a[2], s); s = fmaf(bf_hi(w.y), a[3], s);
s = fmaf(bf_lo(w.z), a[4], s); s = fmaf(bf_hi(w.z), a[5], s);
s = fmaf(bf_lo(w.w), a[6], s); s = fmaf(bf_hi(w.w), a[7], s);
return s;
}
// K=1024 row from smem dotted with the lane's 32 register activations (act[c*256 + lane*8 + j])
__device__ __forceinline__ float row1024_dot(const bf16* row, const float (&a)[32], int lane) {
float s = 0.f;
#pragma unroll
for (int c = 0; c < 4; c++) {
uint4 w = reinterpret_cast<const uint4*>(row + c * 256)[lane];
s = dot8r(w, a + c * 8, s);
}
return warp_sum(s);
}
// row of K bf16 from smem dotted with fp32 activations in smem
template <int K>
__device__ __forceinline__ float smemrow_dot(const bf16* row, const float* act, int lane) {
float s = 0.f;
#pragma unroll
for (int c = 0; c < K / 256; c++) {
uint4 w = reinterpret_cast<const uint4*>(row + c * 256)[lane];
const float4* a = reinterpret_cast<const float4*>(act + c * 256 + lane * 8);
float4 a0 = a[0], a1 = a[1];
s = fmaf(bf_lo(w.x), a0.x, s); s = fmaf(bf_hi(w.x), a0.y, s);
s = fmaf(bf_lo(w.y), a0.z, s); s = fmaf(bf_hi(w.y), a0.w, s);
s = fmaf(bf_lo(w.z), a1.x, s); s = fmaf(bf_hi(w.z), a1.y, s);
s = fmaf(bf_lo(w.w), a1.z, s); s = fmaf(bf_hi(w.w), a1.w, s);
}
return warp_sum(s);
}
// ------------------------------------------------------------- attention
// One position per warp; lane owns dims 4*lane .. 4*lane+3 for both q heads of the kv head.
struct AttnState { float m[2]; float l[2]; float acc[2][4]; };
__device__ __forceinline__ void attn_init(AttnState& st) {
st.m[0] = st.m[1] = -INFINITY;
st.l[0] = st.l[1] = 0.f;
#pragma unroll
for (int j = 0; j < 4; j++) { st.acc[0][j] = 0.f; st.acc[1][j] = 0.f; }
}
__device__ __forceinline__ void attn_update(AttnState& st, const uint2* kp, const uint2* vp, const float (&q)[2][4]) {
uint2 kw = *kp;
float kf[4] = {bf_lo(kw.x), bf_hi(kw.x), bf_lo(kw.y), bf_hi(kw.y)};
float d0 = 0.f, d1 = 0.f;
#pragma unroll
for (int j = 0; j < 4; j++) { d0 = fmaf(q[0][j], kf[j], d0); d1 = fmaf(q[1][j], kf[j], d1); }
d0 = warp_sum(d0);
d1 = warp_sum(d1);
uint2 vw = *vp;
float vf[4] = {bf_lo(vw.x), bf_hi(vw.x), bf_lo(vw.y), bf_hi(vw.y)};
float s0 = d0 * ATTN_SCALE, s1 = d1 * ATTN_SCALE;
float m0 = fmaxf(st.m[0], s0), m1 = fmaxf(st.m[1], s1);
float c0 = __expf(st.m[0] - m0), c1 = __expf(st.m[1] - m1);
float p0 = __expf(s0 - m0), p1 = __expf(s1 - m1);
st.l[0] = fmaf(st.l[0], c0, p0);
st.l[1] = fmaf(st.l[1], c1, p1);
#pragma unroll
for (int j = 0; j < 4; j++) {
st.acc[0][j] = fmaf(p0, vf[j], st.acc[0][j] * c0);
st.acc[1][j] = fmaf(p1, vf[j], st.acc[1][j] * c1);
}
st.m[0] = m0;
st.m[1] = m1;
}
// Fixed-reference-max variant: p = exp(s - M) with a per-layer constant M >= max|s|; no rescaling needed.
__device__ __forceinline__ void attn_update_fixed(AttnState& st, const uint2* kp, const uint2* vp, const float (&q)[2][4],
float negM) {
uint2 kw = *kp;
float kf[4] = {bf_lo(kw.x), bf_hi(kw.x), bf_lo(kw.y), bf_hi(kw.y)};
float d0 = 0.f, d1 = 0.f;
#pragma unroll
for (int j = 0; j < 4; j++) { d0 = fmaf(q[0][j], kf[j], d0); d1 = fmaf(q[1][j], kf[j], d1); }
d0 = warp_sum(d0);
d1 = warp_sum(d1);
uint2 vw = *vp;
float vf[4] = {bf_lo(vw.x), bf_hi(vw.x), bf_lo(vw.y), bf_hi(vw.y)};
float p0 = __expf(fmaf(d0, ATTN_SCALE, negM)), p1 = __expf(fmaf(d1, ATTN_SCALE, negM));
st.l[0] += p0;
st.l[1] += p1;
#pragma unroll
for (int j = 0; j < 4; j++) {
st.acc[0][j] = fmaf(p0, vf[j], st.acc[0][j]);
st.acc[1][j] = fmaf(p1, vf[j], st.acc[1][j]);
}
}
// Two positions at once (independent dependency chains overlap).
__device__ __forceinline__ void attn_update_fixed2(AttnState& st, const uint2* kpA, const uint2* vpA,
const uint2* kpB, const uint2* vpB, const float (&q)[2][4], float negM) {
uint2 kwA = *kpA, kwB = *kpB;
float kA[4] = {bf_lo(kwA.x), bf_hi(kwA.x), bf_lo(kwA.y), bf_hi(kwA.y)};
float kB[4] = {bf_lo(kwB.x), bf_hi(kwB.x), bf_lo(kwB.y), bf_hi(kwB.y)};
float a0 = 0.f, a1 = 0.f, b0 = 0.f, b1 = 0.f;
#pragma unroll
for (int j = 0; j < 4; j++) {
a0 = fmaf(q[0][j], kA[j], a0); a1 = fmaf(q[1][j], kA[j], a1);
b0 = fmaf(q[0][j], kB[j], b0); b1 = fmaf(q[1][j], kB[j], b1);
}
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
a0 += __shfl_xor_sync(0xffffffffu, a0, o); a1 += __shfl_xor_sync(0xffffffffu, a1, o);
b0 += __shfl_xor_sync(0xffffffffu, b0, o); b1 += __shfl_xor_sync(0xffffffffu, b1, o);
}
uint2 vwA = *vpA, vwB = *vpB;
float vA[4] = {bf_lo(vwA.x), bf_hi(vwA.x), bf_lo(vwA.y), bf_hi(vwA.y)};
float vB[4] = {bf_lo(vwB.x), bf_hi(vwB.x), bf_lo(vwB.y), bf_hi(vwB.y)};
float pa0 = __expf(fmaf(a0, ATTN_SCALE, negM)), pa1 = __expf(fmaf(a1, ATTN_SCALE, negM));
float pb0 = __expf(fmaf(b0, ATTN_SCALE, negM)), pb1 = __expf(fmaf(b1, ATTN_SCALE, negM));
st.l[0] += pa0 + pb0;
st.l[1] += pa1 + pb1;
#pragma unroll
for (int j = 0; j < 4; j++) {
st.acc[0][j] = fmaf(pb0, vB[j], fmaf(pa0, vA[j], st.acc[0][j]));
st.acc[1][j] = fmaf(pb1, vB[j], fmaf(pa1, vA[j], st.acc[1][j]));
}
}
// Three positions at once.
__device__ __forceinline__ void attn_update_fixed3(AttnState& st, const uint2* const (&kp)[3], const uint2* const (&vp)[3],
const float (&q)[2][4], float negM) {
uint2 kw[3] = {*kp[0], *kp[1], *kp[2]};
float d0[3] = {0.f, 0.f, 0.f}, d1[3] = {0.f, 0.f, 0.f};
#pragma unroll
for (int i = 0; i < 3; i++) {
float kf[4] = {bf_lo(kw[i].x), bf_hi(kw[i].x), bf_lo(kw[i].y), bf_hi(kw[i].y)};
#pragma unroll
for (int j = 0; j < 4; j++) { d0[i] = fmaf(q[0][j], kf[j], d0[i]); d1[i] = fmaf(q[1][j], kf[j], d1[i]); }
}
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
#pragma unroll
for (int i = 0; i < 3; i++) { d0[i] += __shfl_xor_sync(0xffffffffu, d0[i], o); d1[i] += __shfl_xor_sync(0xffffffffu, d1[i], o); }
}
uint2 vw[3] = {*vp[0], *vp[1], *vp[2]};
#pragma unroll
for (int i = 0; i < 3; i++) {
float vf[4] = {bf_lo(vw[i].x), bf_hi(vw[i].x), bf_lo(vw[i].y), bf_hi(vw[i].y)};
float p0 = __expf(fmaf(d0[i], ATTN_SCALE, negM)), p1 = __expf(fmaf(d1[i], ATTN_SCALE, negM));
st.l[0] += p0;
st.l[1] += p1;
#pragma unroll
for (int j = 0; j < 4; j++) {
st.acc[0][j] = fmaf(p0, vf[j], st.acc[0][j]);
st.acc[1][j] = fmaf(p1, vf[j], st.acc[1][j]);
}
}
}
// Block-wide KV tile: positions [lo + 16t, min(hi, lo + 16t + 16)) of one kv head; K rows then V rows.
// Tile j of a layer lives in stage j % NST. Stages 0..NSTAGE-1 are dedicated; stages NSTAGE.. alias W2 and may only
// hold tiles while attention runs (W2 is otherwise a weight buffer). Per-stage phase parities are tracked in bits.
struct Ring {
uint32_t ring_addr, w2_addr, w1_addr, full_addr, empty_addr;
uint32_t full_par, empty_par; // bit k: parity to wait for on stage k's next use
};
__device__ __forceinline__ uint32_t ring_stage_addr(const Ring& r, int stage) {
if (stage < NSTAGE) return r.ring_addr + stage * TILE_BYTES;
if (stage < NPRE) return r.w2_addr + (stage - NSTAGE) * TILE_BYTES;
return r.w1_addr + (stage - NPRE) * TILE_BYTES;
}
__device__ __forceinline__ void ring_issue(Ring& r, int j, int lo, int hi, const bf16* kb, const bf16* vb, int tid, uint64_t pol,
int issuer = 0) {
int stage = j % NST;
if (tid == issuer) {
mbar_wait(r.empty_addr + stage * 8, (r.empty_par >> stage) & 1u); // previous tile in this stage consumed
int p0 = lo + j * TILE_POS;
int n = min(hi - p0, TILE_POS);
uint32_t bytes = (uint32_t)n * (HD * 2);
uint32_t mb = r.full_addr + stage * 8;
uint32_t dst = ring_stage_addr(r, stage);
mbar_expect_tx(mb, 2 * bytes);
bulk_g2s(dst, kb + (size_t)p0 * HD, bytes, mb, pol);
bulk_g2s(dst + TILE_BYTES / 2, vb + (size_t)p0 * HD, bytes, mb, pol);
}
r.empty_par ^= (1u << stage);
}
__device__ __forceinline__ void ring_wait_full(Ring& r, int j) {
int stage = j % NST;
mbar_wait(r.full_addr + stage * 8, (r.full_par >> stage) & 1u);
r.full_par ^= (1u << stage);
}
__device__ __forceinline__ void ring_release(const Ring& r, int j, int lane) {
int stage = j % NST;
if (lane == 0) mbar_arrive(r.empty_addr + stage * 8);
}
// ------------------------------------------------------------------ kernel
__global__ void __launch_bounds__(NTHR, 1) megaqwen_kernel(Params p) {
extern __shared__ __align__(128) unsigned char smem_raw[];
uint4* ring = reinterpret_cast<uint4*>(smem_raw + OFF_RING);
float* wpart = reinterpret_cast<float*>(smem_raw + OFF_RING); // P2 merge: [16][2][130]
float* act3 = reinterpret_cast<float*>(smem_raw + OFF_W2); // P3: attn_out [2048] (W2 stage 3 region, idle in P3)
bf16* w1 = reinterpret_cast<bf16*>(smem_raw + OFF_W1);
bf16* w2 = reinterpret_cast<bf16*>(smem_raw + OFF_W2);
float* act4 = reinterpret_cast<float*>(smem_raw + OFF_ACT4); // P1/P4: normalized input [1024]
float* q_s = act4; // P2: [2][128]
bf16* k_s = reinterpret_cast<bf16*>(act4 + 256); // P2: [128]
bf16* v_s = k_s + HD; // P2: [128]
float* red2 = act4 + 512; // P2: [128] combine scratch
float* red = reinterpret_cast<float*>(smem_raw + OFF_RED); // [32]
float* res = reinterpret_cast<float*>(smem_raw + OFF_RES); // [64]
int* misc = reinterpret_cast<int*>(smem_raw + OFF_MISC);
float* cs_tab = reinterpret_cast<float*>(smem_raw + OFF_TRIG); // [64]
float* sn_tab = cs_tab + 64; // [64]
const uint32_t mbar_base = (uint32_t)__cvta_generic_to_shared(smem_raw + OFF_MBAR);
const uint32_t w1_addr = (uint32_t)__cvta_generic_to_shared(w1);
const uint32_t w2_addr = (uint32_t)__cvta_generic_to_shared(w2);
const uint32_t wm1 = mbar_base + 2 * NST * 8; // W1 first-half slice barrier
const uint32_t wm1b = wm1 + 8; // W1 second-half slice barrier
const uint32_t wm2 = wm1 + 16; // W2 slice barrier
const int nblk = gridDim.x;
const int b = blockIdx.x, t = threadIdx.x, warp = t >> 5, lane = t & 31;
const int g = b & 7, cidx = b >> 3;
const int rep = b & (NREP - 1);
const bf16* h_rd = p.h + (size_t)rep * HID;
const float* hp_rd = p.h_prime + (size_t)rep * HID;
const float* mlp_rd = p.mlp + (size_t)rep * INTER;
const int nch = nblk / NKV + ((g < (nblk % NKV)) ? 1 : 0);
const bool designated = (cidx == nch - 1);
const uint64_t pol_w = make_policy_evict_last();
const uint64_t pol_kv = make_policy_evict_first();
Ring rg;
rg.ring_addr = (uint32_t)__cvta_generic_to_shared(ring);
rg.w2_addr = w2_addr;
rg.w1_addr = w1_addr;
rg.full_addr = mbar_base;
rg.empty_addr = mbar_base + NST * 8;
rg.full_par = 0u;
rg.empty_par = (1u << NST) - 1u; // fresh barriers: waiting on the "previous" phase (parity 1) passes immediately
uint32_t wseq1 = 0, wseq1b = 0, wseq2 = 0; // weight-slice barrier uses (uniform)
if (t == 0) {
for (int i = 0; i < NST; i++) mbar_init(rg.full_addr + i * 8, 1);
for (int i = 0; i < NST; i++) mbar_init(rg.empty_addr + i * 8, NWARP);
mbar_init(wm1, 1);
mbar_init(wm1b, 1);
mbar_init(wm2, 1);
}
fence_proxy_async();
__syncthreads();
unsigned long long* tm = (p.timing != nullptr && b == 0 && t == 0) ? p.timing : nullptr;
unsigned long long* tml = tm; // per-layer stamp base
unsigned long long* tdbg = (p.timing != nullptr && t == 0) ? p.timing + 2048 : nullptr; // per-block debug
#define TSTAMP(slot) do { if (tml) { int dep_ = *reinterpret_cast<volatile int*>(misc + 4); unsigned long long t_; \
asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t_) : "r"(dep_) : "memory"); tml[slot] = t_; } } while (0)
long long clk0 = 0; unsigned long long gt0 = 0;
if (tm) { asm volatile("mov.u64 %0, %%clock64;" : "=l"(clk0)); gt0 = gtimer(); }
// ---- per-block row slices
int qkv_r0, qkv_n, o_r0, o_n, gu_m0, gu_np;
block_rows(QKV_ROWS, b, nblk, qkv_r0, qkv_n);
block_rows(HID, b, nblk, o_r0, o_n);
block_rows(INTER, b, nblk, gu_m0, gu_np); // gate/up pairs
const int gu_nrows = 2 * gu_np; // interleaved rows: 2m = gate m, 2m+1 = up m
const int qkv_n1 = min(qkv_n, W1_ROWS), qkv_n2 = qkv_n - qkv_n1;
const int gu_n1 = min(gu_nrows, W1_ROWS), gu_n2 = gu_nrows - gu_n1;
auto qkv_row_ptr = [&](int L, int row) -> const bf16* {
if (row < NQ * HD) return p.w_q[L] + (size_t)row * HID;
if (row < NQ * HD + NKV * HD) return p.w_k[L] + (size_t)(row - NQ * HD) * HID;
return p.w_v[L] + (size_t)(row - NQ * HD - NKV * HD) * HID;
};
// rows [r, r+n) of the concatenated q/k/v matrix as contiguous segments -> smem dst (thread 0)
auto issue_qkv_rows = [&](int L, int r, int n, uint32_t dst, uint32_t mb) {
int rend = r + n, off = 0;
while (r < rend) {
int seg_end = r < NQ * HD ? NQ * HD : (r < NQ * HD + NKV * HD ? NQ * HD + NKV * HD : QKV_ROWS);
int k = min(rend, seg_end) - r;
bulk_g2s(dst + off, qkv_row_ptr(L, r), (uint32_t)k * HID * 2, mb, pol_w);
off += k * HID * 2;
r += k;
}
};
// QKV: W1 rows [0, qkv_n1) split into W1a = rows [0, 9) and W1b = rows [9, 18); W2 rows [0, qkv_n2).
const int qkv_n1a = min(qkv_n1, W1_ROWS / 2), qkv_n1b = qkv_n1 - qkv_n1a;
auto issue_qkv1a = [&](int L) {
if (t == POLL_T) { mbar_expect_tx(wm1, (uint32_t)qkv_n1a * HID * 2); if (qkv_n1a) issue_qkv_rows(L, qkv_r0, qkv_n1a, w1_addr, wm1); }
};
auto issue_qkv1b = [&](int L) {
if (t == POLL_T) {
mbar_expect_tx(wm1b, (uint32_t)qkv_n1b * HID * 2);
if (qkv_n1b) issue_qkv_rows(L, qkv_r0 + qkv_n1a, qkv_n1b, w1_addr + qkv_n1a * HID * 2, wm1b);
}
};
auto issue_qkv2 = [&](int L) {
if (t == POLL_T) {
mbar_expect_tx(wm2, (uint32_t)qkv_n2 * HID * 2);
if (qkv_n2) issue_qkv_rows(L, qkv_r0 + qkv_n1, qkv_n2, w2_addr + W2_QKV_ROW0 * HID * 2, wm2);
}
};
// gate/up part 1 (pairs [0, gu_n1/2)) in W1: gate rows -> W1a (wm1), up rows -> W1b (wm1b); part 2 in W2 (wm2).
const int gu_p1 = gu_n1 >> 1, gu_p2 = gu_n2 >> 1;
auto issue_gu1a = [&](int L) {
if (t == POLL_T) {
uint32_t bytes = (uint32_t)gu_p1 * HID * 2;
mbar_expect_tx(wm1, bytes);
if (bytes) bulk_g2s(w1_addr, p.w_gate[L] + (size_t)gu_m0 * HID, bytes, wm1, pol_w);
}
};
auto issue_gu1b = [&](int L) {
if (t == POLL_T) {
uint32_t bytes = (uint32_t)gu_p1 * HID * 2;
mbar_expect_tx(wm1b, bytes);
if (bytes) bulk_g2s(w1_addr + bytes, p.w_up[L] + (size_t)gu_m0 * HID, bytes, wm1b, pol_w);
}
};
auto issue_gu2 = [&](int L) {
if (t == POLL_T) {
uint32_t bytes = (uint32_t)gu_p2 * HID * 2;
mbar_expect_tx(wm2, 2 * bytes);
if (bytes) {
bulk_g2s(w2_addr, p.w_gate[L] + (size_t)(gu_m0 + gu_p1) * HID, bytes, wm2, pol_w);
bulk_g2s(w2_addr + bytes, p.w_up[L] + (size_t)(gu_m0 + gu_p1) * HID, bytes, wm2, pol_w);
}
}
};
// local interleaved row lr -> smem row pointer
auto gu_row_smem = [&](int lr) -> const bf16* {
if (lr < gu_n1) { int pr = lr >> 1; return w1 + (size_t)(((lr & 1) ? gu_p1 + pr : pr)) * HID; }
int l2 = lr - gu_n1, pr = l2 >> 1;
return w2 + (size_t)(((l2 & 1) ? gu_p2 + pr : pr)) * HID;
};
auto issue_o = [&](int L) {
if (t == POLL_T) {
mbar_expect_tx(wm1, (uint32_t)o_n * (NQ * HD) * 2);
if (o_n) bulk_g2s(w1_addr, p.w_o[L] + (size_t)o_r0 * (NQ * HD), (uint32_t)o_n * (NQ * HD) * 2, wm1, pol_w);
}
};
// down rows [0, o_n): rows [0, 3) -> W1a (wm1), rows [3, o_n) -> W1b (wm1b)
const int d_na = min(o_n, 3), d_nb = o_n - d_na;
auto issue_da = [&](int L) {
if (t == POLL_T) {
mbar_expect_tx(wm1, (uint32_t)d_na * INTER * 2);
if (d_na) bulk_g2s(w1_addr, p.w_down[L] + (size_t)o_r0 * INTER, (uint32_t)d_na * INTER * 2, wm1, pol_w);
}
};
auto issue_db = [&](int L) {
if (t == POLL_T) {
mbar_expect_tx(wm1b, (uint32_t)d_nb * INTER * 2);
if (d_nb) bulk_g2s(w1_addr + d_na * INTER * 2, p.w_down[L] + (size_t)(o_r0 + d_na) * INTER, (uint32_t)d_nb * INTER * 2, wm1b, pol_w);
}
};
auto wait_w1 = [&]() { mbar_wait(wm1, wseq1 & 1u); wseq1++; };
auto wait_w1b = [&]() { mbar_wait(wm1b, wseq1b & 1u); wseq1b++; };
auto wait_w2 = [&]() { mbar_wait(wm2, wseq2 & 1u); wseq2++; };
// ---- split grid barrier: arrive (release fence + relaxed add) ... prefetch ... wait (poll).
// A gpu-scope fence waits for ALL outstanding TMA/LDG traffic of the SM, so prefetches are
// issued strictly after the arrive, and the poller (warp 15) never has loads in flight.
unsigned bar_target = 0;
// All phase outputs are published with returning atomics (see pub_*), so no fence is needed here.
auto bar_arrive = [&]() {
__syncthreads();
bar_target += (unsigned)nblk;
if (t == 0) atomicAdd(p.gbar, 1u);
};
auto bar_wait = [&]() {
if (t == POLL_T) {
while (ld_relaxed_gpu_u32(p.gbar) < bar_target) { }
}
__syncthreads();
};
auto next_chunk = [&](int s, int L, int& nS, int& nL, int& nlo, int& nhi, int& nt) {
nL = L + 1; nS = s;
if (nL == p.num_layers) { nL = 0; nS = s + 1; }
nt = 0; nlo = 0; nhi = 0;
if (nS < p.n_steps) {
chunk_bounds(p.start_pos + nS, cidx, nch, nlo, nhi);
nt = (nhi > nlo) ? (nhi - nlo + TILE_POS - 1) / TILE_POS : 0;
}
};
auto ring_cur_tile = [&](int k, int cs, int cL) { // tile k of (step cs, layer cL)'s own chunk
int clo, chi;
chunk_bounds(p.start_pos + cs, cidx, nch, clo, chi);
int cnt = (chi > clo) ? (chi - clo + TILE_POS - 1) / TILE_POS : 0;
if (k < cnt && k < NPRE) {
const bf16* ckb = p.kc[cL] + (size_t)g * p.kv_head_stride;
const bf16* cvb = p.vc[cL] + (size_t)g * p.kv_head_stride;
ring_issue(rg, k, clo, chi, ckb, cvb, t, pol_kv, (6 + (k % 10)) * 32);
}
};
auto ring_prefill_tile = [&](int k, int nS, int nL, int nlo, int nhi, int nt) {
if (nS < p.n_steps && k < nt && k < NPRE) {
const bf16* nkb = p.kc[nL] + (size_t)g * p.kv_head_stride;
const bf16* nvb = p.vc[nL] + (size_t)g * p.kv_head_stride;
ring_issue(rg, k, nlo, nhi, nkb, nvb, t, pol_kv, (6 + (k % 10)) * 32); // tile k -> stage k; warps 6.. issue
}
};
if (b == 0) { // zero both fixed-max accumulator buffers (published; visible to all after the first barrier)
unsigned sink = 0;
for (int i = t; i < 2 * NQ * HD; i += NTHR) sink |= pub_f32(p.attn_acc + i, 0.f);
if (t < 2 * NQ) sink |= pub_f32(p.l_acc + t, 0.f);
pub_sink(sink, misc);
}
// ---- initial prefetch: layer-0 QKV weights, first KV tiles of (step 0, layer 0)
issue_qkv2(0);
issue_qkv1a(0);
issue_qkv1b(0);
{
int lo, hi;
chunk_bounds(p.start_pos, cidx, nch, lo, hi);
const bf16* kb = p.kc[0] + (size_t)g * p.kv_head_stride;
const bf16* vb = p.vc[0] + (size_t)g * p.kv_head_stride;
int nt = (hi > lo) ? (hi - lo + TILE_POS - 1) / TILE_POS : 0;
for (int tt = 0; tt < NPRE && tt < nt; tt++) ring_issue(rg, tt, lo, hi, kb, vb, t, pol_kv, (6 + (tt % 10)) * 32);
}
for (int s = 0; s < p.n_steps; s++) {
const int pos = p.start_pos + s;
for (int L = 0; L < p.num_layers; L++) {
if (tm) { int li = s * p.num_layers + L; tml = (li < 40) ? tm + li * 48 : nullptr; }
int nS, nL, nlo, nhi, nt;
next_chunk(s, L, nS, nL, nlo, nhi, nt);
// =============================== P1: RMSNorm + QKV projection
TSTAMP(0);
{
uint32_t hw = __ldcg(reinterpret_cast<const unsigned int*>((s + L == 0) ? p.h_in : h_rd) + t);
uint32_t ww = __ldg(reinterpret_cast<const unsigned int*>(p.w_in_ln[L]) + t);
uint32_t rw = 0;
if (L == 0) rw = __ldg(reinterpret_cast<const unsigned int*>(p.randn + (size_t)s * HID) + t);
// TMA issues right after the activation loads are in flight (their data returns first):
// this layer's KV tiles 3,4,5 (W2 stages; P5's act staging there is done) and the QKV W1b rows.
if (s + L > 0) { issue_qkv1a(L); issue_qkv1b(L); } // W1 rows (step 0, layer 0: issued at kernel start)
if (L == 0 && t >= NTHR - HD / 2) { // RoPE table for this position (used by P2 of every layer of this step)
int i = t - (NTHR - HD / 2);
float ang = __fmul_rn((float)pos, p.inv_freq[i]);
float sn, cs;
sincosf(ang, &sn, &cs);
cs_tab[i] = cs;
sn_tab[i] = sn;
}
float x0, x1;
if (L == 0) {
x0 = bf2f(__float2bfloat16_rn(0.5f * bf_lo(rw) + 0.5f * bf_lo(hw)));
x1 = bf2f(__float2bfloat16_rn(0.5f * bf_hi(rw) + 0.5f * bf_hi(hw)));
} else {
x0 = bf_lo(hw);
x1 = bf_hi(hw);
}
float ss = warp_sum(x0 * x0 + x1 * x1);
if (lane == 0) red[warp] = ss;
__syncthreads();
float tot = 0.f;
#pragma unroll
for (int i = 0; i < NWARP; i++) tot += red[i];
float r = rsqrtf(tot * (1.0f / HID) + RMS_EPS);
reinterpret_cast<float2*>(act4)[t] = make_float2((x0 * r) * bf_lo(ww), (x1 * r) * bf_hi(ww));
__syncthreads();
float a[32];
#pragma unroll
for (int c = 0; c < 4; c++) {
float4 v0 = reinterpret_cast<const float4*>(act4 + c * 256 + lane * 8)[0];
float4 v1 = reinterpret_cast<const float4*>(act4 + c * 256 + lane * 8)[1];
a[c * 8 + 0] = v0.x; a[c * 8 + 1] = v0.y; a[c * 8 + 2] = v0.z; a[c * 8 + 3] = v0.w;
a[c * 8 + 4] = v1.x; a[c * 8 + 5] = v1.y; a[c * 8 + 6] = v1.z; a[c * 8 + 7] = v1.w;
}
TSTAMP(1);
wait_w2();
unsigned sink = 0;
for (int lr = qkv_n1 + warp; lr < qkv_n; lr += NWARP) {
float sum = row1024_dot(w2 + (size_t)(W2_QKV_ROW0 + lr - qkv_n1) * HID, a, lane);
if (lane == 0) sink |= pub_f32(p.qkv_raw + qkv_r0 + lr, sum);
}
wait_w1();
for (int lr = warp; lr < qkv_n1a; lr += NWARP) {
float sum = row1024_dot(w1 + (size_t)lr * HID, a, lane);
if (lane == 0) sink |= pub_f32(p.qkv_raw + qkv_r0 + lr, sum);
}
wait_w1b();
for (int lr = qkv_n1a + warp; lr < qkv_n1; lr += NWARP) {
float sum = row1024_dot(w1 + (size_t)lr * HID, a, lane);
if (lane == 0) sink |= pub_f32(p.qkv_raw + qkv_r0 + lr, sum);
}
pub_sink(sink, misc);
fence_proxy_async();
}
TSTAMP(2);
bar_arrive();
bar_wait();
TSTAMP(3);
const bool dbg_layer = (tdbg != nullptr) && (s == 1) && (L == 1);
if (dbg_layer) tdbg[b] = gtimer();
// =============================== P2: attention (split over kv head x chunk)
{
const bf16* kb = p.kc[L] + (size_t)g * p.kv_head_stride;
const bf16* vb = p.vc[L] + (size_t)g * p.kv_head_stride;
// --- setup: threads 0..255: grp 0/1 = q heads 2g/2g+1, grp 2 = k head g, grp 3 = v head g; thread owns dims i, i+64
if (t < 256) {
int grp = t >> 6, i = t & 63;
const float* src = (grp < 2) ? p.qkv_raw + (2 * g + grp) * HD
: (grp == 2 ? p.qkv_raw + NQ * HD + g * HD : p.qkv_raw + NQ * HD + NKV * HD + g * HD);
float a = __ldcg(src + i), bq = __ldcg(src + i + 64);
const bf16* nw = (grp < 2) ? p.w_qn[L] : p.w_kn[L];
float nw0 = bf2f(nw[i]), nw1 = bf2f(nw[i + 64]); // norm weights fetched in the same round trip
TSTAMP(35);
float ss = warp_sum(a * a + bq * bq);
TSTAMP(36); // loads arrived
if (lane == 0) red[warp] = ss;
__syncthreads();
TSTAMP(37);
if (grp < 3) {
float gss = red[grp * 2] + red[grp * 2 + 1];
float r = rsqrtf(gss * (1.0f / HD) + RMS_EPS);
float an = (a * r) * nw0, bn = (bq * r) * nw1;
float cs = cs_tab[i], sn = sn_tab[i];
float o1 = __fsub_rn(__fmul_rn(an, cs), __fmul_rn(bn, sn));
float o2 = __fadd_rn(__fmul_rn(an, sn), __fmul_rn(bn, cs));
if (grp < 2) {
q_s[grp * HD + i] = o1;
q_s[grp * HD + i + 64] = o2;
} else {
bf16 k1 = __float2bfloat16_rn(o1), k2 = __float2bfloat16_rn(o2);
k_s[i] = k1;
k_s[i + 64] = k2;
}
} else {
bf16 v1 = __float2bfloat16_rn(a), v2 = __float2bfloat16_rn(bq);
v_s[i] = v1;
v_s[i + 64] = v2;
}
} else {
__syncthreads();
}
__syncthreads();
if (designated && t < 128) { // publish the new K/V cache rows (32-bit pairs) with returning atomics
const unsigned* ksrc = reinterpret_cast<const unsigned*>(k_s);
const unsigned* vsrc = reinterpret_cast<const unsigned*>(v_s);
unsigned* kr = reinterpret_cast<unsigned*>(const_cast<bf16*>(kb) + (size_t)pos * HD);
unsigned* vr = reinterpret_cast<unsigned*>(const_cast<bf16*>(vb) + (size_t)pos * HD);
unsigned sink = (t < 64) ? pub_u32(kr + t, ksrc[t]) : pub_u32(vr + (t - 64), vsrc[t - 64]);
pub_sink(sink, misc);
}
float q[2][4];
#pragma unroll
for (int j = 0; j < 4; j++) { q[0][j] = q_s[lane * 4 + j]; q[1][j] = q_s[HD + lane * 4 + j]; }
TSTAMP(4);
// --- stream this block's chunk of cached positions through the ring (one position per warp per tile)
int lo, hi;
chunk_bounds(pos, cidx, nch, lo, hi);
AttnState st;
attn_init(st);
const float mest = p.m_est[L];
const bool fixedmax = (mest > 0.f);
const float negM = -mest;
int ntiles = (hi > lo) ? (hi - lo + TILE_POS - 1) / TILE_POS : 0;
// tiles 0..NPRE-1 were pre-filled during the previous phases; W1 is free now -> issue tiles NPRE..NST-1
for (int j = NPRE; j < NST && j < ntiles; j++) ring_issue(rg, j, lo, hi, kb, vb, t, pol_kv, (6 + (j % 10)) * 32);
auto tile_ptrs = [&](int j, const uint2*& kp, const uint2*& vp) {
const unsigned char* tile = smem_raw + (ring_stage_addr(rg, j % NST) - rg.ring_addr) + OFF_RING;
kp = reinterpret_cast<const uint2*>(tile + warp * (HD * 2)) + lane;
vp = reinterpret_cast<const uint2*>(tile + TILE_BYTES / 2 + warp * (HD * 2)) + lane;
};
if (ntiles <= NPRE) {
// fast path (short contexts): everything is already in shared memory
for (int j = 0; j < ntiles; j++) ring_wait_full(rg, j);
if (tm) TSTAMP(25);
for (int tt = 0; tt < ntiles; tt += 3) {
const int nb = min(3, ntiles - tt);
const uint2* kp[3]; const uint2* vp[3]; bool ok[3];
#pragma unroll
for (int k = 0; k < 3; k++) {
tile_ptrs(tt + (k < nb ? k : 0), kp[k], vp[k]);
ok[k] = (k < nb) && (lo + (tt + k) * TILE_POS + warp < hi);
}
if (fixedmax) {
if (ok[0] && ok[1] && ok[2]) attn_update_fixed3(st, kp, vp, q, negM);
else if (ok[0] && ok[1]) attn_update_fixed2(st, kp[0], vp[0], kp[1], vp[1], q, negM);
else if (ok[0]) attn_update_fixed(st, kp[0], vp[0], q, negM);
} else {
#pragma unroll
for (int k = 0; k < 3; k++) if (ok[k]) attn_update(st, kp[k], vp[k], q);
}
}
if (tm) TSTAMP(26);
__syncwarp();
for (int j = 0; j < ntiles; j++) ring_release(rg, j, lane);
} else {
for (int tt = 0; tt < ntiles; tt += 3) {
const int nb = min(3, ntiles - tt);
for (int k = 0; k < nb; k++) ring_wait_full(rg, tt + k);
if (tt < 3) TSTAMP(25 + 2 * tt);
const uint2* kp[3]; const uint2* vp[3]; bool ok[3];
#pragma unroll
for (int k = 0; k < 3; k++) {
tile_ptrs(tt + (k < nb ? k : 0), kp[k], vp[k]);
ok[k] = (k < nb) && (lo + (tt + k) * TILE_POS + warp < hi);
}
if (fixedmax) {
if (ok[0] && ok[1] && ok[2]) attn_update_fixed3(st, kp, vp, q, negM);
else if (ok[0] && ok[1]) attn_update_fixed2(st, kp[0], vp[0], kp[1], vp[1], q, negM);
else if (ok[0]) attn_update_fixed(st, kp[0], vp[0], q, negM);
} else {
#pragma unroll
for (int k = 0; k < 3; k++) if (ok[k]) attn_update(st, kp[k], vp[k], q);
}
if (tt < 3) TSTAMP(26 + 2 * tt);
__syncwarp();
for (int k = 0; k < nb; k++) ring_release(rg, tt + k, lane);
for (int k = 0; k < nb; k++)
if (tt + k + NST < ntiles) ring_issue(rg, tt + k + NST, lo, hi, kb, vb, t, pol_kv, (6 + ((tt + k) % 10)) * 32);
}
}
TSTAMP(5);
if (dbg_layer) tdbg[256 + b] = gtimer();
// --- the new position (designated block only, warp 0)
if (designated && warp == 0) {
const uint2* kp = reinterpret_cast<const uint2*>(k_s) + lane;
const uint2* vp = reinterpret_cast<const uint2*>(v_s) + lane;
if (fixedmax) attn_update_fixed(st, kp, vp, q, negM);
else attn_update(st, kp, vp, q);
}
// --- merge the 16 warps (wpart aliases the drained ring)
fence_proxy_async();
__syncthreads(); // all warps done reading their ring slots before wpart overwrites the ring region
issue_o(L); // O-proj slice -> W1 (W1 ring stages drained; lands during merge + barrier + P3 act load)
#pragma unroll
for (int hh = 0; hh < 2; hh++) {
float* wp = wpart + (warp * 2 + hh) * PART_STRIDE;
reinterpret_cast<float4*>(wp)[lane] = make_float4(st.acc[hh][0], st.acc[hh][1], st.acc[hh][2], st.acc[hh][3]);
if (lane == 0) { wp[128] = st.m[hh]; wp[129] = st.l[hh]; }
}
__syncthreads();
if (fixedmax) {
// fixed-max path: block sum -> fire-and-forget fp32 reductions into the per-head accumulators
const int buf = (s * p.num_layers + L) & 1;
if (t < 256) {
int hh = t >> 7, d = t & 127;
float A = 0.f, Lsum = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; w++) {
const float* wp = wpart + (w * 2 + hh) * PART_STRIDE;
A += wp[d];
Lsum += wp[129];
}
unsigned sink = pub_add_f32(p.attn_acc + buf * (NQ * HD) + (2 * g + hh) * HD + d, A);
if (d == 0) sink |= pub_add_f32(p.l_acc + buf * NQ + 2 * g + hh, Lsum);
pub_sink(sink, misc);
}
TSTAMP(6); TSTAMP(7); TSTAMP(8);
} else if (t < 256) {
int hh = t >> 7, d = t & 127;
float M = -INFINITY;
#pragma unroll
for (int w = 0; w < NWARP; w++) M = fmaxf(M, wpart[(w * 2 + hh) * PART_STRIDE + 128]);
float A = 0.f, Lsum = 0.f;
if (M > -INFINITY) {
#pragma unroll
for (int w = 0; w < NWARP; w++) {
const float* wp = wpart + (w * 2 + hh) * PART_STRIDE;
float mw = wp[128];
float sc = (mw == -INFINITY) ? 0.f : __expf(mw - M);
A += wp[d] * sc;
Lsum += wp[129] * sc;
}
}
float* gp = p.partials + ((size_t)(g * MAXCH + cidx) * 2 + hh) * PART_STRIDE;
gp[d] = A;
if (d == 0) { gp[128] = M; gp[129] = Lsum; }
}
if (!fixedmax) {
// --- last-arriver combine for this kv head (thread 0 fences only; data read via L2)
__syncthreads();
TSTAMP(6);
if (t == 0) {
__threadfence();
unsigned old = atomicAdd(p.counters + g, 1u);
int last = (old == (unsigned)(nch - 1)) ? 1 : 0;
if (last) { p.counters[g] = 0u; __threadfence(); }
misc[0] = last;
}
__syncthreads();
TSTAMP(7);
TSTAMP(8);
if (misc[0]) {
unsigned long long* cdbg = (dbg_layer && t == 0) ? tdbg + 1216 + 8 * g : nullptr;
if (cdbg) cdbg[0] = gtimer();
const float* base = p.partials + (size_t)(g * MAXCH) * 2 * PART_STRIDE;
{
// thread -> (head hh, dim d, chunk parity); one L2 round trip: acc[d], m, l for its 12 chunks
int hh = t >> 8, d = (t >> 1) & 127, par = t & 1;
float av[MAXCH / 2], mv[MAXCH / 2], lv[MAXCH / 2];
#pragma unroll
for (int k = 0; k < MAXCH / 2; k++) {
int c = 2 * k + par;
const float* cp = base + ((size_t)c * 2 + hh) * PART_STRIDE;
bool ok = (c < nch);
av[k] = ok ? __ldcg(cp + d) : 0.f;
mv[k] = ok ? __ldcg(cp + 128) : -INFINITY;
lv[k] = ok ? __ldcg(cp + 129) : 0.f;
}
if (cdbg) cdbg[1] = gtimer();
float M = -INFINITY;
#pragma unroll
for (int k = 0; k < MAXCH / 2; k++) M = fmaxf(M, mv[k]);
M = fmaxf(M, __shfl_xor_sync(0xffffffffu, M, 1)); // both parities -> head max
float A = 0.f, Lsum = 0.f;
#pragma unroll
for (int k = 0; k < MAXCH / 2; k++) {
float sc = (mv[k] == -INFINITY) ? 0.f : __expf(mv[k] - M);
A += av[k] * sc;
Lsum += lv[k] * sc;
}
A += __shfl_xor_sync(0xffffffffu, A, 1);
Lsum += __shfl_xor_sync(0xffffffffu, Lsum, 1);
if (par == 0) p.attn_out[(2 * g + hh) * HD + d] = A / Lsum;
}
if (cdbg) { cdbg[2] = gtimer(); cdbg[3] = cdbg[2]; }
}
} // !fixedmax
}
fence_proxy_async(); // ring / W2 generic reads done before later TMA writes
TSTAMP(9);
if (dbg_layer) { tdbg[512 + b] = gtimer(); tdbg[768 + b] = misc[0]; }
bar_arrive();
ring_prefill_tile(0, nS, nL, nlo, nhi, nt); // next attention's tiles -> dedicated stages (drained; merge done)
ring_prefill_tile(1, nS, nL, nlo, nhi, nt);
ring_prefill_tile(2, nS, nL, nlo, nhi, nt);
bar_wait();
TSTAMP(10);
if (dbg_layer) tdbg[1024 + b] = gtimer();
// =============================== P3: O-proj + residual (act in the ring region; W1 = O slice)
{
float resid = 0.f;
if (t < o_n) {
int row = o_r0 + t;
float hv = ldcg_bf16_elem((s + L == 0) ? p.h_in : h_rd, row);
if (L == 0) {
float rv = bf2f(p.randn[(size_t)s * HID + row]);
resid = bf2f(__float2bfloat16_rn(0.5f * rv + 0.5f * hv));
} else {
resid = hv;
}
}
if (t < o_n) res[40 + t] = resid;
TSTAMP(32);
float4 av4;
if (p.m_est[L] > 0.f) {
const int buf = (s * p.num_layers + L) & 1;
av4 = __ldcg(reinterpret_cast<const float4*>(p.attn_acc + buf * (NQ * HD)) + t);
float linv = 1.0f / __ldcg(p.l_acc + buf * NQ + (t >> 5)); // head of dims 4t..4t+3
av4.x *= linv; av4.y *= linv; av4.z *= linv; av4.w *= linv;
} else {
av4 = __ldcg(reinterpret_cast<const float4*>(p.attn_out) + t);
}
TSTAMP(33);
reinterpret_cast<float4*>(act3)[t] = av4;
TSTAMP(34);
__syncthreads();
TSTAMP(11);
if (dbg_layer) tdbg[1280 + b] = gtimer();
wait_w1();
TSTAMP(12);
if (warp < 2 * o_n) { // row = warp/2, K-half = warp&1
const int row = warp >> 1, half = warp & 1;
float sum = smemrow_dot<NQ * HD / 2>(w1 + (size_t)row * (NQ * HD) + half * (NQ * HD / 2), act3 + half * (NQ * HD / 2), lane);
if (lane == 0) res[16 + warp] = sum;
}
TSTAMP(13);
fence_proxy_async(); // W1 / W2 (act3) generic accesses done before the TMA writes after the arrive
__syncthreads();
unsigned sink = 0;
if (t < NREP * o_n) { // publish h' replicas: thread -> (replica t / o_n, row t % o_n)
const int rr = t / o_n, row = t - rr * o_n;
float v = res[16 + 2 * row] + res[16 + 2 * row + 1] + res[40 + row];
sink = pub_f32(p.h_prime + (size_t)rr * HID + o_r0 + row, v);
}
pub_sink(sink, misc);
}
TSTAMP(14);
if (dbg_layer) tdbg[1536 + b] = gtimer();
bar_arrive();
bar_wait();
TSTAMP(15);
// =============================== P4: post-norm + gate/up + SiLU (act in registers)
unsigned zero_sink = 0;
{
{ // zero the fixed-max accumulators for the next layer (last read in P3 two layers ago); spread over blocks
const int nbuf = (s * p.num_layers + L + 1) & 1;
const int per = (NQ * HD + NQ + nblk - 1) / nblk; // elements per block (acc then l)
const int i0 = b * per;
if (t < per && i0 + t < NQ * HD + NQ) {
const int i = i0 + t;
unsigned sv = (i < NQ * HD) ? pub_f32(p.attn_acc + nbuf * (NQ * HD) + i, 0.f) : pub_f32(p.l_acc + nbuf * NQ + (i - NQ * HD), 0.f);
zero_sink = sv;
}
}
float2 xv = __ldcg(reinterpret_cast<const float2*>(hp_rd) + t);
uint32_t ww = __ldg(reinterpret_cast<const unsigned int*>(p.w_post_ln[L]) + t);
// TMA issues while the activation loads are in flight (their data returns first): gate/up part 1 (gate rows ->
// W1a, up rows -> W1b; O rows consumed in P3) and part 2 -> W2 (act3 consumed in P3).
issue_gu1a(L);
issue_gu1b(L);
issue_gu2(L);
float ss = warp_sum(xv.x * xv.x + xv.y * xv.y);
if (lane == 0) red[warp] = ss;
__syncthreads();
float tot = 0.f;
#pragma unroll
for (int i = 0; i < NWARP; i++) tot += red[i];
float r = rsqrtf(tot * (1.0f / HID) + RMS_EPS);
reinterpret_cast<float2*>(act4)[t] = make_float2((xv.x * r) * bf_lo(ww), (xv.y * r) * bf_hi(ww));
__syncthreads();
float a[32];
#pragma unroll
for (int c = 0; c < 4; c++) {
float4 v0 = reinterpret_cast<const float4*>(act4 + c * 256 + lane * 8)[0];
float4 v1 = reinterpret_cast<const float4*>(act4 + c * 256 + lane * 8)[1];
a[c * 8 + 0] = v0.x; a[c * 8 + 1] = v0.y; a[c * 8 + 2] = v0.z; a[c * 8 + 3] = v0.w;
a[c * 8 + 4] = v1.x; a[c * 8 + 5] = v1.y; a[c * 8 + 6] = v1.z; a[c * 8 + 7] = v1.w;
}
TSTAMP(16);
wait_w1();
for (int lr = 2 * warp; lr < gu_n1; lr += 2 * NWARP) { // part 1 gate rows (even lr)
float sum = row1024_dot(gu_row_smem(lr), a, lane);
if (lane == 0) res[lr] = sum;
}
wait_w1b();
for (int lr = 2 * warp + 1; lr < gu_n1; lr += 2 * NWARP) { // part 1 up rows (odd lr)
float sum = row1024_dot(gu_row_smem(lr), a, lane);
if (lane == 0) res[lr] = sum;
}
wait_w2();
for (int lr = warp + gu_n1; lr < gu_nrows; lr += NWARP) { // part 2 (W2; issued at the end of P3)
float sum = row1024_dot(gu_row_smem(lr), a, lane);
if (lane == 0) res[lr] = sum;
}
fence_proxy_async(); // W1 / W2 row reads done before the TMA writes after the arrive
__syncthreads();
TSTAMP(17);
unsigned sink = zero_sink;
if (t < NREP * gu_np) { // publish mlp replicas: thread -> (replica t / gu_np, pair t % gu_np)
const int rr = t / gu_np, m = t - rr * gu_np;
float gv = res[2 * m], uv = res[2 * m + 1];
sink |= pub_f32(p.mlp + (size_t)rr * INTER + gu_m0 + m, gv / (1.0f + expf(-gv)) * uv);
}
pub_sink(sink, misc);
}
TSTAMP(18);
bar_arrive();
ring_prefill_tile(3, nS, nL, nlo, nhi, nt); // W2 stage 3 (gate/up part 2 consumed); stages 4,5 host P5's act
bar_wait();
TSTAMP(19);
// =============================== P5: down + residual -> h (bf16); activations staged in smem:
// mlp[0..2048) at W2+24K (QKV part-2 slot, free now), mlp[2048..3072) in act4
{
float* act5a = reinterpret_cast<float*>(reinterpret_cast<unsigned char*>(w2) + TILE_BYTES); // W2 stages 4,5 (idle now)
float* act5b = act4; // 1024 floats
float hp = 0.f;
if (t < o_n) hp = __ldcg(hp_rd + o_r0 + t);
float4 m0 = __ldcg(reinterpret_cast<const float4*>(mlp_rd) + t);
float4 m1 = make_float4(0.f, 0.f, 0.f, 0.f);
if (t < 256) m1 = __ldcg(reinterpret_cast<const float4*>(mlp_rd) + 512 + t);
// TMA issues while the activation loads are in flight: down rows -> W1 (gate/up part 1 consumed in P4) and the
// next layer's QKV W2 rows (W2 tail).
issue_da(L);
issue_db(L);
issue_qkv2(nL);
reinterpret_cast<float4*>(act5a)[t] = m0;
if (t < 256) reinterpret_cast<float4*>(act5b)[t] = m1;
if (t < o_n) res[40 + t] = hp;
__syncthreads();
TSTAMP(20);
wait_w1();
TSTAMP(21);
// down rows: K split in halves across warps (warp -> row warp/2, half warp&1); rows 0..2 from W1a, 3..5 from W1b
auto down_half = [&](int row, int half) {
const bf16* wrow = w1 + (size_t)row * INTER + half * (INTER / 2);
float sum = 0.f;
#pragma unroll
for (int cc = 0; cc < INTER / 512; cc++) {
const int c = half * (INTER / 512) + cc;
uint4 w = reinterpret_cast<const uint4*>(wrow + cc * 256)[lane];
const float* ap = (c < 8) ? act5a + c * 256 + lane * 8 : act5b + (c - 8) * 256 + lane * 8;
float4 a0 = reinterpret_cast<const float4*>(ap)[0], a1 = reinterpret_cast<const float4*>(ap)[1];
sum = fmaf(bf_lo(w.x), a0.x, sum); sum = fmaf(bf_hi(w.x), a0.y, sum);
sum = fmaf(bf_lo(w.y), a0.z, sum); sum = fmaf(bf_hi(w.y), a0.w, sum);
sum = fmaf(bf_lo(w.z), a1.x, sum); sum = fmaf(bf_hi(w.z), a1.y, sum);
sum = fmaf(bf_lo(w.w), a1.z, sum); sum = fmaf(bf_hi(w.w), a1.w, sum);
}
return warp_sum(sum);
};
const int drow = warp >> 1, dhalf = warp & 1;
if (warp < 2 * d_na) {
float sum = down_half(drow, dhalf);
if (lane == 0) res[16 + warp] = sum;
}
wait_w1b();
if (warp >= 2 * d_na && warp < 2 * o_n) {
float sum = down_half(drow, dhalf);
if (lane == 0) res[16 + warp] = sum;
}
fence_proxy_async(); // W1 / act5 generic accesses done before the TMA writes after the arrive
__syncthreads();
const int npair = o_n >> 1; // bf16 pairs (o_r0 is even)
if (t < NREP * npair + npair) { // replicas 0..NREP-1, then (t >= NREP*npair) the final output when last
const int rr = t / npair, pi = t - rr * npair;
const bool final_out = (rr == NREP);
if (!final_out || (s == p.n_steps - 1 && L == p.num_layers - 1)) {
float v0 = res[16 + 4 * pi] + res[16 + 4 * pi + 1] + res[40 + 2 * pi];
float v1 = res[16 + 4 * pi + 2] + res[16 + 4 * pi + 3] + res[40 + 2 * pi + 1];
bf16 lo_b = __float2bfloat16_rn(v0), hi_b = __float2bfloat16_rn(v1);
unsigned word = (unsigned)__bfloat16_as_ushort(lo_b) | ((unsigned)__bfloat16_as_ushort(hi_b) << 16);
unsigned* dst = final_out ? reinterpret_cast<unsigned*>(p.h_out) : reinterpret_cast<unsigned*>(p.h + (size_t)rr * HID);
pub_sink(pub_u32(dst + ((o_r0 >> 1) + pi), word), misc);
}
}
TSTAMP(22);
}
TSTAMP(23);
bar_arrive();
ring_prefill_tile(4, nS, nL, nlo, nhi, nt); // W2 stages 4,5 (P5's act staging consumed)
ring_prefill_tile(5, nS, nL, nlo, nhi, nt);
bar_wait();
TSTAMP(24);
}
}
if (tm) { long long clk1; asm volatile("mov.u64 %0, %%clock64;" : "=l"(clk1)); tm[2000] = (unsigned long long)(clk1 - clk0); tm[2001] = gtimer() - gt0; }
#undef TSTAMP
}
// ------------------------------------------------------------------- host
static bool g_attr_set = false;
void megaqwen_run(std::vector<torch::Tensor> w, std::vector<torch::Tensor> kc, std::vector<torch::Tensor> vc,
torch::Tensor randn, torch::Tensor h, torch::Tensor qkv_raw, torch::Tensor partials,
torch::Tensor attn_out, torch::Tensor h_prime, torch::Tensor mlp, torch::Tensor counters,
torch::Tensor inv_freq, torch::Tensor timing, torch::Tensor attn_acc, torch::Tensor l_acc,
std::vector<double> m_est, torch::Tensor h_rep, torch::Tensor h_in, int64_t start_pos, int64_t n_steps) {
int L = (int)kc.size();
TORCH_CHECK(L >= 1 && L <= MAXL, "num_layers out of range");
TORCH_CHECK((int)w.size() == 11 * L, "weights list size mismatch");
TORCH_CHECK((int)vc.size() == L, "v cache list size mismatch");
Params p;
memset(&p, 0, sizeof(p));
for (int i = 0; i < L; i++) {
p.w_in_ln[i] = reinterpret_cast<const bf16*>(w[i * 11 + 0].data_ptr());
p.w_q[i] = reinterpret_cast<const bf16*>(w[i * 11 + 1].data_ptr());
p.w_k[i] = reinterpret_cast<const bf16*>(w[i * 11 + 2].data_ptr());
p.w_v[i] = reinterpret_cast<const bf16*>(w[i * 11 + 3].data_ptr());
p.w_qn[i] = reinterpret_cast<const bf16*>(w[i * 11 + 4].data_ptr());
p.w_kn[i] = reinterpret_cast<const bf16*>(w[i * 11 + 5].data_ptr());
p.w_o[i] = reinterpret_cast<const bf16*>(w[i * 11 + 6].data_ptr());
p.w_post_ln[i] = reinterpret_cast<const bf16*>(w[i * 11 + 7].data_ptr());
p.w_gate[i] = reinterpret_cast<const bf16*>(w[i * 11 + 8].data_ptr());
p.w_up[i] = reinterpret_cast<const bf16*>(w[i * 11 + 9].data_ptr());
p.w_down[i] = reinterpret_cast<const bf16*>(w[i * 11 + 10].data_ptr());
TORCH_CHECK(kc[i].is_contiguous() && vc[i].is_contiguous(), "kv cache must be contiguous");
TORCH_CHECK(kc[i].size(0) == NKV && kc[i].size(2) == HD, "kv cache shape");
TORCH_CHECK(kc[i].stride(0) == kc[0].stride(0) && vc[i].stride(0) == kc[0].stride(0), "kv cache strides differ");
p.kc[i] = reinterpret_cast<bf16*>(kc[i].data_ptr());
p.vc[i] = reinterpret_cast<bf16*>(vc[i].data_ptr());
}
p.kv_head_stride = kc[0].stride(0);
TORCH_CHECK(start_pos + n_steps <= kc[0].size(1), "kv cache too small");
p.randn = reinterpret_cast<const bf16*>(randn.data_ptr());
p.h = reinterpret_cast<bf16*>(h_rep.data_ptr());
p.h_in = reinterpret_cast<const bf16*>(h_in.data_ptr());
p.h_out = reinterpret_cast<bf16*>(h.data_ptr());
TORCH_CHECK(h_rep.numel() == NREP * HID && h_prime.numel() == NREP * HID && mlp.numel() == NREP * INTER, "replica buffer sizes");
p.qkv_raw = qkv_raw.data_ptr<float>();
p.partials = partials.data_ptr<float>();
p.attn_out = attn_out.data_ptr<float>();
p.h_prime = h_prime.data_ptr<float>();
p.mlp = mlp.data_ptr<float>();
p.counters = reinterpret_cast<unsigned*>(counters.data_ptr<int32_t>());
p.gbar = reinterpret_cast<unsigned*>(counters.data_ptr<int32_t>()) + NKV;
p.inv_freq = inv_freq.data_ptr<float>();
p.attn_acc = attn_acc.data_ptr<float>();
p.l_acc = l_acc.data_ptr<float>();
TORCH_CHECK((int)m_est.size() == L, "m_est size mismatch");
for (int i = 0; i < L; i++) p.m_est[i] = (float)m_est[i];
p.timing = timing.numel() > 0 ? reinterpret_cast<unsigned long long*>(timing.data_ptr<int64_t>()) : nullptr;
p.start_pos = (int)start_pos;
p.n_steps = (int)n_steps;
p.num_layers = L;
int dev = 0;
cudaGetDevice(&dev);
int nsm = 0;
cudaDeviceGetAttribute(&nsm, cudaDevAttrMultiProcessorCount, dev);
TORCH_CHECK(nsm >= 8 && nsm <= 8 * MAXCH, "unexpected SM count");
TORCH_CHECK((QKV_ROWS + nsm - 1) / nsm <= W1_ROWS + SM_W2 / (HID * 2), "QKV rows per block exceed W1+W2");
TORCH_CHECK(2 * ((INTER + nsm - 1) / nsm) <= W1_ROWS + SM_W2 / (HID * 2), "gate/up rows per block exceed W1+W2");
TORCH_CHECK(((HID + nsm - 1) / nsm) * INTER * 2 <= SM_W1, "down rows per block exceed W1");
if (!g_attr_set) {
cudaError_t e = cudaFuncSetAttribute((const void*)megaqwen_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES);
TORCH_CHECK(e == cudaSuccess, "set smem attr: ", cudaGetErrorString(e));
g_attr_set = true;
}
int max_blocks = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&max_blocks, megaqwen_kernel, NTHR, SMEM_BYTES);
TORCH_CHECK(max_blocks >= 1, "kernel does not fit on an SM");
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
void* args[] = {&p};
cudaError_t e = cudaLaunchCooperativeKernel((const void*)megaqwen_kernel, dim3(nsm), dim3(NTHR), args, SMEM_BYTES, stream);
TORCH_CHECK(e == cudaSuccess, "cooperative launch failed: ", cudaGetErrorString(e));
}
int64_t megaqwen_smem_bytes() { return SMEM_BYTES; }
"""
CPP_SRC = r"""
#include <torch/extension.h>
#include <vector>
void megaqwen_run(std::vector<torch::Tensor> w, std::vector<torch::Tensor> kc, std::vector<torch::Tensor> vc,
torch::Tensor randn, torch::Tensor h, torch::Tensor qkv_raw, torch::Tensor partials,
torch::Tensor attn_out, torch::Tensor h_prime, torch::Tensor mlp, torch::Tensor counters,
torch::Tensor inv_freq, torch::Tensor timing, torch::Tensor attn_acc, torch::Tensor l_acc,
std::vector<double> m_est, torch::Tensor h_rep, torch::Tensor h_in, int64_t start_pos, int64_t n_steps);
int64_t megaqwen_smem_bytes();
"""
_EXT = None
def _ensure_ninja_on_path() -> None:
cand = os.path.dirname(os.path.abspath(sys.executable))
if os.path.exists(os.path.join(cand, "ninja")) and cand not in os.environ.get("PATH", "").split(os.pathsep):
os.environ["PATH"] = cand + os.pathsep + os.environ.get("PATH", "")
try:
import ninja # noqa: F401
bin_dir = getattr(ninja, "BIN_DIR", None)
if bin_dir and bin_dir not in os.environ.get("PATH", "").split(os.pathsep):
os.environ["PATH"] = bin_dir + os.pathsep + os.environ.get("PATH", "")
except Exception:
pass
def _get_ext():
global _EXT
if _EXT is None:
_ensure_ninja_on_path()
from torch.utils.cpp_extension import load_inline
major, minor = torch.cuda.get_device_capability(0)
arch = f"-gencode=arch=compute_{major}{minor},code=sm_{major}{minor}"
verbose = os.environ.get("MEGAQWEN_VERBOSE", "0") == "1"
flags = ["-O3", arch, "-lineinfo", "--expt-relaxed-constexpr"]
if verbose:
flags.append("-Xptxas=-v")
_EXT = load_inline(
name="megaqwen_decode_mk_v36",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["megaqwen_run", "megaqwen_smem_bytes"],
extra_cuda_cflags=flags,
verbose=verbose,
)
return _EXT
# ----------------------------------------------------------------------------
# Model (same parameterization / 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 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 = 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)])
# Upper bound of |attention score| per layer: |q.k|*scale <= scale*HD*max|w_qn|*max|w_kn| (RMSNorm outputs have
# norm sqrt(HD)). If it is small enough (<= 40) the kernel uses a fixed reference max in the softmax and merges
# chunks with fp32 reductions; otherwise it falls back to the exact online-softmax + last-arriver combine.
_M_EST_LIMIT = 40.0
def kernel_state(self):
"""(weight tensor list, per-layer score bounds), cached until any parameter is re-pointed or modified."""
key = tuple((p.data_ptr(), p._version) for p in self.parameters())
cached = getattr(self, "_kstate", None)
if cached is not None and cached[0] == key:
return cached[1], cached[2]
ws = []
for blk in self.blocks:
for name in ("input_ln", "q_proj", "k_proj", "v_proj", "q_norm", "k_norm", "o_proj", "post_ln",
"gate_proj", "up_proj", "down_proj"):
w = getattr(blk, name).data
if w.dtype != torch.bfloat16 or not w.is_cuda or not w.is_contiguous():
raise RuntimeError("model weights must be contiguous bf16 CUDA tensors")
ws.append(w)
bounds = []
for blk in self.blocks:
mq = blk.q_norm.detach().float().abs().max()
mk = blk.k_norm.detach().float().abs().max()
m = float((mq * mk).item()) * (HEAD_DIM / math.sqrt(HEAD_DIM)) * 1.02
bounds.append(m if m <= self._M_EST_LIMIT else 0.0)
self._kstate = (key, ws, bounds)
return ws, bounds
def kernel_weights(self):
return self.kernel_state()[0]
def score_bounds(self):
return self.kernel_state()[1]
# ----------------------------------------------------------------------------
# Scratch buffers
# ----------------------------------------------------------------------------
_SCRATCH: dict = {}
def _scratch(device: torch.device) -> dict:
key = (device.type, device.index)
sc = _SCRATCH.get(key)
if sc is None:
half = HEAD_DIM // 2
inv = 1.0 / (10000 ** (torch.arange(0, half, device=device, dtype=torch.float32) / half))
sc = {
"qkv_raw": torch.zeros(NUM_Q * HEAD_DIM + 2 * NUM_KV * HEAD_DIM, device=device, dtype=torch.float32),
"partials": torch.zeros(NUM_KV * 32 * 2 * 132, device=device, dtype=torch.float32),
"attn_out": torch.zeros(NUM_Q * HEAD_DIM, device=device, dtype=torch.float32),
"h_prime": torch.zeros(4 * HIDDEN, device=device, dtype=torch.float32),
"mlp": torch.zeros(4 * INTERMEDIATE, device=device, dtype=torch.float32),
"h_rep": torch.zeros(4 * HIDDEN, device=device, dtype=torch.bfloat16),
"counters": torch.zeros(NUM_KV + 64, device=device, dtype=torch.int32),
"inv_freq": inv.contiguous(),
"timing": torch.zeros(0, device=device, dtype=torch.int64),
"attn_acc": torch.zeros(2 * NUM_Q * HEAD_DIM, device=device, dtype=torch.float32),
"l_acc": torch.zeros(2 * NUM_Q, device=device, dtype=torch.float32),
}
_SCRATCH[key] = sc
return sc
def _launch(model: Model, h_in: torch.Tensor, h_out: torch.Tensor, k_caches, v_caches, randn_dev: torch.Tensor,
start_pos: int, n_steps: int):
"""Run n_steps decode steps: reads h_in, updates the caches in place, writes h_out."""
ext = _get_ext()
sc = _scratch(h_in.device)
ws, bounds = model.kernel_state()
sc["counters"].zero_()
ext.megaqwen_run(
ws, list(k_caches), list(v_caches), randn_dev, h_out,
sc["qkv_raw"], sc["partials"], sc["attn_out"], sc["h_prime"], sc["mlp"], sc["counters"], sc["inv_freq"],
sc["timing"], sc["attn_acc"], sc["l_acc"], bounds, sc["h_rep"], h_in, int(start_pos), int(n_steps),
)
def empty_caches(num_layers: int, max_seq: int, device, dtype=torch.bfloat16):
k = [torch.zeros(NUM_KV, max_seq, HEAD_DIM, device=device, dtype=dtype) for _ in range(num_layers)]
v = [torch.zeros(NUM_KV, max_seq, HEAD_DIM, device=device, dtype=dtype) for _ in range(num_layers)]
return k, v
def _seeded_hidden(seed: int, device) -> torch.Tensor:
g = torch.Generator(device="cpu")
g.manual_seed(seed)
return torch.randn(HIDDEN, generator=g, dtype=torch.bfloat16).to(device)
def _randn_rows(g: torch.Generator, n: int, device) -> torch.Tensor:
"""n rows of the reference's per-step `torch.randn(HIDDEN, generator=g, dtype=bfloat16)` stream.
Row-batched generation is bit-identical to the per-step calls (the CPU normal kernel fills 16 values at a time
and HIDDEN is a multiple of 16), and drawing float32 then rounding to bf16 is bit-identical to drawing bf16
directly (that is how the bf16 path is implemented) but several times faster on the CPU."""
x = torch.randn(n, HIDDEN, generator=g, dtype=torch.float32).to(torch.bfloat16)
return x.to(device)
# ----------------------------------------------------------------------------
# Protocol
# ----------------------------------------------------------------------------
_PREFILL_CHUNK = 4096
@torch.no_grad()
def prefill(model: Model, ctx_len: int, seed: int, device: torch.device | None = None):
"""Build KV of length ctx_len by sequential decode steps at positions 0..ctx_len-1 (same numerics as the
reference). Untimed."""
device = device or next(model.parameters()).device
model = model.to(device).eval()
assert ctx_len <= model.max_seq
h = _seeded_hidden(seed, device)
k_caches, v_caches = empty_caches(model.num_layers, model.max_seq, device)
g = torch.Generator(device="cpu")
g.manual_seed(seed + 1)
for s0 in range(0, ctx_len, _PREFILL_CHUNK):
n = min(_PREFILL_CHUNK, ctx_len - s0)
rn = _randn_rows(g, n, device)
h_next = torch.empty_like(h)
_launch(model, h, h_next, k_caches, v_caches, rn, s0, n)
h = h_next
return h, k_caches, v_caches
@torch.no_grad()
def decode_steps(model: Model, hidden: torch.Tensor, k_caches, v_caches, start_pos: int, n_steps: int, seed: int):
"""n_steps decode steps starting at start_pos (timed). Caches are updated in place."""
device = hidden.device
g = torch.Generator(device="cpu")
g.manual_seed(seed + 2)
h_in = hidden.to(torch.bfloat16).contiguous()
rn = _randn_rows(g, n_steps, device)
h_out = torch.empty_like(h_in)
_launch(model, h_in, h_out, k_caches, v_caches, rn, start_pos, n_steps)
return h_out, k_caches, v_caches
@torch.no_grad()
def run(ctx_len: int, n_decode: int, seed: int, model: Model | None = None, max_seq: int | None = None) -> dict:
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(NUM_LAYERS, max_seq)
else:
if getattr(model, "max_seq", 0) < ctx_len + n_decode:
raise ValueError(
f"model.max_seq={getattr(model, 'max_seq', None)} too small for 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}
def get_init_inputs():
return [NUM_LAYERS, 131072]
def get_inputs():
return []
20260904_072208_or-fable_anthropic_claude-fable-5-1_03_megaqwen_decode