kernelbench.com

KernelBench mega · H100

Kimi-Linear Decode Claude Fable 5

19.12×geomean speedup across shapes

manually audited: clean

harnessor-fable
Kernel source (redacted)
"""Kimi-Linear W4A16 hybrid decode unit -- single-launch Triton megakernel.

One decode step == ONE persistent-grid @triton.jit kernel launch. Inside that
launch, a fixed sequence of stages runs the entire per-token forward -- fused
int4 dequant-GEMVs for every projection and MoE expert, the short causal conv,
the KDA gated delta-rule state update, absorbed-MLA latent-cache attention
(including the cache append), the MoE router/top-8 and expert accumulation,
both RMSNorms and all residual adds -- synchronized by software global
barriers (an atomic counter with a volatile spin + nanosleep backoff) between
stages. Weights are streamed exactly once per step in int4; no bf16 weight
matrix is ever materialized.

Fused dequant-GEMV: per 128-group, the activation chunk is quantized to two
int8 levels (value + quantization residual, so the fused product is exact to
~1e-5) once per stage into a private per-program scratch slice; each GEMV
task then multiplies the packed int4 nibbles against those int8 rows on
tensor cores (imma), and rescales with the group's (scale, zero) on the fly:
    y[n] += s*(xs1*d1 + xs2*d2) - s*z*(xs1*sum(q1) + xs2*sum(q2)).

MLA runs in the absorbed form: q_eff = W_kb_k^T q_nope is computed once
(32x512), scores are taken directly against the compressed latent cache with
a two-pass softmax (scores+max, then exp+weighted latent sum), and the value
projection is folded back through W_kb_v -- so attention touches only the
512-d latents plus 64-d rope keys, never a materialized per-token K/V.

Numerics mirror the reference closely: the residual stream is bf16-rounded at
every residual add exactly like the reference (h = bf16(x + bf16(attn));
x' = bf16(h + bf16(moe))) via committed buffers XA/XB plus f32 accumulators
AACC/MACC (consumers between accumulation and commit derive the rounded value
on the fly), and the router logits are bf16-rounded before softmax/top-8 so
expert selection matches the reference.

The MLA latent cache lives in a preallocated internal buffer; the incoming
state's cache is ingested by the kernel itself (stage 0) on the first step of
a run, and the returned state is a view of that buffer, so step() never runs
any extra kernel or copy. KDA state (S, conv windows) is updated in place.

Grid co-residency (required for the in-kernel barriers) is verified at load
time by a bounded-spin probe launch; the grid configuration falls back
automatically if the target GPU cannot co-schedule the preferred config.
"""
from __future__ import annotations

import torch
import torch.nn as nn
import triton
import triton.language as tl

OP_TYPE = "kimi_linear_w4a16_decode"
HARDWARE_REQUIRED = ["RTX_PRO_6000"]
GROUP_SIZE = 128
SCR_SIZE = 83360
PRIV_F = 2592           # per-program f32 scratch stride
PRIV_Q = 18432          # per-program int8 scratch stride

# ------------------------------------------------------------------------- #
# kernel constants (scratch layout offsets in f32 elements, model dims)
# ------------------------------------------------------------------------- #
EPS = tl.constexpr(1.0e-6)
OFF_XA = tl.constexpr(0)            # 2304   committed residual (layer input)
OFF_QRAW = tl.constexpr(2304)       # 16384  raw GEMV accumulators
OFF_QKVG = tl.constexpr(18688)      # 16384  conv'd q,k,v + decay gate
OFF_BETA = tl.constexpr(35072)      # 32     beta logits
OFF_OATT = tl.constexpr(35104)      # 4096   attention output vector
OFF_GU = tl.constexpr(39200)        # 9*2048 gate/up expert outputs
OFF_CTX = tl.constexpr(57632)       # 32*512 MLA context accumulator
OFF_GMAX = tl.constexpr(74016)      # 32     per-head score max
OFF_DEN = tl.constexpr(74048)       # 32     per-head softmax denominator
OFF_XB = tl.constexpr(74080)        # 2304   committed residual (post-attn h)
OFF_AACC = tl.constexpr(76384)      # 2304   attention o_proj accumulator
OFF_MACC = tl.constexpr(78688)      # 2304   MoE output accumulator
OFF_ZERO = tl.constexpr(80992)      # 2304   always-zero pad
OFF_LOGITS = tl.constexpr(83296)    # 64     router logits accumulator

# private per-program scratch layout
PQ2 = tl.constexpr(9216)            # offset of second-level quant (int8)
PXS1 = tl.constexpr(0)              # [72] group scales, level 1
PXS2 = tl.constexpr(72)             # [72] group scales, level 2
PSQ1 = tl.constexpr(144)            # [72] group int sums, level 1
PSQ2 = tl.constexpr(216)            # [72] group int sums, level 2
PXN = tl.constexpr(288)             # [2304] normed activation (f32 of bf16)
PSF = tl.constexpr(2592)
PSQ = tl.constexpr(18432)

D = tl.constexpr(2304)
GD = tl.constexpr(18)               # D // 128
NQKVG = tl.constexpr(16384)
NMQ = tl.constexpr(6720)            # MLA q (6144) + kv_a (576)
KDA_SCALE = tl.constexpr(0.08838834764831845)   # 128 ** -0.5
MLA_SCALE = tl.constexpr(0.07216878364870323)   # 192 ** -0.5
SPIN_CAP = tl.constexpr(2000000)


# ------------------------------------------------------------------------- #
# device helpers
# ------------------------------------------------------------------------- #
@triton.jit
def _bar(BAR, ABORT, TIMES, phase, sid, NPROG: tl.constexpr):
    tl.debug_barrier()
    tl.atomic_add(BAR, 1, sem="acq_rel", scope="gpu")
    phase += NPROG
    n = 0
    done = tl.load(BAR, volatile=True) >= phase
    while (done == 0) & (n < SPIN_CAP):
        n += 1
        if n > 32:
            tl.inline_asm_elementwise(
                "nanosleep.u32 $1; mov.u32 $0, 0;", "=r,r", [tl.full((), 1024, tl.int32)],
                dtype=tl.int32, is_pure=False, pack=1)
        done = tl.load(BAR, volatile=True) >= phase
    if done == 0:
        tl.atomic_xchg(ABORT, 1)
    tl.atomic_add(BAR, 0, sem="acquire", scope="gpu")
    if tl.program_id(0) == 0:
        tl.store(TIMES + sid, tl.extra.cuda.globaltimer())
    tl.debug_barrier()
    return phase, sid + 1


@triton.jit
def _rstd(SCR, offa, offb):
    ss = 0.0
    for part in tl.static_range(2):
        offs = part * 2048 + tl.arange(0, 2048)
        m = offs < D
        a = tl.load(SCR + offa + offs, mask=m, other=0.0, cache_modifier=".cg")
        b = tl.load(SCR + offb + offs, mask=m, other=0.0, cache_modifier=".cg")
        v = (a + b.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16).to(tl.float32)
        ss += tl.sum(v * v)
    return tl.rsqrt(ss / D + EPS)


@triton.jit
def _pick3(a, b, c, i):
    if i == 0:
        r = a
    elif i == 1:
        r = b
    else:
        r = c
    return r


@triton.jit
def _prequant(SCR, XQP, XSP, offa, offb, rstd, NRM, MODE: tl.constexpr):
    """Quantize this stage's activation vector once per program (two-level
    int8 per 128-group) into private scratch. MODE 0: rmsnorm-ed derived
    residual (18 groups, also stores the normed f32 vector); MODE 1: bf16 of
    OATT (32 groups); MODE 2: silu(gate)*up from GU (72 groups)."""
    pid = tl.program_id(0)
    qpb = XQP + pid.to(tl.int64) * PSQ
    psb = XSP + pid.to(tl.int64) * PSF
    r2048 = tl.arange(0, 2048)
    r64 = tl.arange(0, 64)
    if MODE == 0:
        PARTS: tl.constexpr = 2
        NG: tl.constexpr = 18
    elif MODE == 1:
        PARTS: tl.constexpr = 2
        NG: tl.constexpr = 32
    else:
        PARTS: tl.constexpr = 5
        NG: tl.constexpr = 72
    for part in tl.static_range(PARTS):
        idx = part * 2048 + r2048
        if MODE == 0:
            m = idx < D
            a = tl.load(SCR + offa + idx, mask=m, other=0.0, cache_modifier=".cg")
            b = tl.load(SCR + offb + idx, mask=m, other=0.0, cache_modifier=".cg")
            x2 = (a + b.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16).to(tl.float32)
            nw = tl.load(NRM + idx, mask=m, other=0.0)
            xn = (x2 * rstd * nw).to(tl.bfloat16).to(tl.float32)
            tl.store(psb + PXN + idx, xn, mask=m)
        elif MODE == 1:
            xn = tl.load(SCR + offa + idx,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
        else:
            m = idx < 9216
            slot = idx // 1024
            mm = idx % 1024
            ga = tl.load(SCR + OFF_GU + slot * 2048 + mm, mask=m, other=0.0,
                         cache_modifier=".cg")
            gb = tl.load(SCR + OFF_GU + slot * 2048 + 1024 + mm, mask=m, other=0.0,
                         cache_modifier=".cg")
            xn = ga * tl.sigmoid(ga) * gb
        X = tl.reshape(xn, (16, 128))
        xs1 = tl.max(tl.abs(X), 1) / 127.0 + 1e-30
        q1 = tl.floor(X / xs1[:, None] + 0.5)
        rr = X - xs1[:, None] * q1
        xs2 = tl.max(tl.abs(rr), 1) / 127.0 + 1e-30
        q2 = tl.floor(rr / xs2[:, None] + 0.5)
        sq1 = tl.sum(q1, 1)
        sq2 = tl.sum(q2, 1)
        g16 = part * 16 + tl.arange(0, 16)
        gm = g16 < NG
        fm = (part * 2048 + tl.arange(0, 2048)) < NG * 128
        # natural order; the consumer separates even/odd k via int16 tricks
        tl.store(qpb + part * 2048 + tl.arange(0, 2048),
                 tl.reshape(q1.to(tl.int8), (2048,)), mask=fm)
        tl.store(qpb + PQ2 + part * 2048 + tl.arange(0, 2048),
                 tl.reshape(q2.to(tl.int8), (2048,)), mask=fm)
        tl.store(psb + PXS1 + g16, xs1, mask=gm)
        tl.store(psb + PXS2 + g16, xs2, mask=gm)
        tl.store(psb + PSQ1 + g16, sq1, mask=gm)
        tl.store(psb + PSQ2 + g16, sq2, mask=gm)
    tl.debug_barrier()


@triton.jit(noinline=True)
def _gemv_task(WB, SB, ZB, NW, qg0, gc, nb, QP16, psb, OUTP, oncols, wscale,
               GPT: tl.constexpr):
    """Fused int4-dequant GEMV partial task: gpt contiguous k-groups x 256
    cols, accumulated atomically into OUTP. Activations come pre-quantized
    (two int8 levels) from private per-program scratch (read as int16 pairs
    to split even/odd k); the weight tile for the next group is prefetched
    while the current one is consumed. Rescale:
      y += s*(xs1*d1 + xs2*d2) - s*z*(xs1*sum(q1) + xs2*sum(q2))
    Shared (noinline) across every projection/expert stage of the kernel."""
    r64 = tl.arange(0, 64)
    r16 = tl.arange(0, 16)
    cols = nb * 256 + tl.arange(0, 256)
    cmask = cols < oncols
    acc = tl.zeros([256], tl.float32)
    g0 = gc * GPT
    qpb16 = QP16 + (qg0 + g0) * 64
    for gi in tl.range(0, GPT, num_stages=1):
        g = g0 + gi
        qg = qg0 + g
        p1 = tl.load(qpb16 + gi * 64 + r64)
        p2 = tl.load(qpb16 + PQ2 // 2 + gi * 64 + r64)
        q1e = ((p1 << 8) >> 8).to(tl.int8)
        q1o = (p1 >> 8).to(tl.int8)
        q2e = ((p2 << 8) >> 8).to(tl.int8)
        q2o = (p2 >> 8).to(tl.int8)
        xs1 = tl.load(psb + PXS1 + qg)
        xs2 = tl.load(psb + PXS2 + qg)
        sq1 = tl.load(psb + PSQ1 + qg)
        sq2 = tl.load(psb + PSQ2 + qg)
        Xe = tl.where(r16[:, None] == 0, q1e[None, :], 0) \
             + tl.where(r16[:, None] == 1, q2e[None, :], 0)
        Xo = tl.where(r16[:, None] == 0, q1o[None, :], 0) \
             + tl.where(r16[:, None] == 1, q2o[None, :], 0)
        w = tl.load(WB + (g * 64 + r64)[:, None] * NW + cols[None, :],
                    mask=cmask[None, :], other=0)
        lo = (w & 15).to(tl.int8)
        hi = (w >> 4).to(tl.int8)
        dd = tl.dot(Xe, lo, out_dtype=tl.int32) + tl.dot(Xo, hi, out_dtype=tl.int32)
        d1 = tl.sum(tl.where(r16[:, None] == 0, dd, 0), 0).to(tl.float32)
        d2 = tl.sum(tl.where(r16[:, None] == 1, dd, 0), 0).to(tl.float32)
        sv = tl.load(SB + g * NW + cols, mask=cmask, other=0).to(tl.float32)
        zv = tl.load(ZB + g * NW + cols, mask=cmask, other=0).to(tl.float32)
        acc += sv * (xs1 * d1 + xs2 * d2) - sv * zv * (xs1 * sq1 + xs2 * sq2)
    tl.atomic_add(OUTP + cols, acc * wscale, mask=cmask)


@triton.jit
def _topk(SCR):
    """Top-8 experts from the LOGITS accumulator (bf16-rounded like ref)."""
    r64o = tl.arange(0, 64)
    lg = tl.load(SCR + OFF_LOGITS + r64o,
                 cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
    mx = tl.max(lg)
    p = tl.exp(lg - mx)
    probs = p / tl.sum(p)
    rank = tl.full([64], 99, tl.int32)
    for it in range(8):
        cur = tl.where(rank < 99, -1.0, probs)
        mxc = tl.max(cur)
        cand = tl.where(cur == mxc, r64o, 99999)
        sel = tl.min(cand)
        rank = tl.where(r64o == sel, it, rank)
    wsum = tl.sum(tl.where(rank < 8, probs, 0.0))
    return rank, probs, wsum


@triton.jit
def _oproj(pid, SCR, XQP, XQP16, XSP, OW, OS, OZ, WQL, WKL, WVL, WNT,
           GPTO: tl.constexpr, NPROG: tl.constexpr):
    """o_proj GEMV (K=4096) into AACC; zeros GU, MACC, LOGITS; for KDA layers
    (WNT=12) also shifts the short-conv window state."""
    _prequant(SCR, XQP, XSP, OFF_OATT, OFF_ZERO, 0.0, SCR, 1)
    qp16 = XQP16 + pid.to(tl.int64) * (PSQ // 2)
    psb = XSP + pid.to(tl.int64) * PSF
    NGC: tl.constexpr = 32 // GPTO
    NTO = 9 * NGC
    NT = NTO + 9 + 3 + WNT
    for t in range(pid, NT, NPROG):
        if t < NTO:
            nb = t // NGC
            gc = t % NGC
            _gemv_task(OW, OS, OZ, 2304, 0, gc, nb, qp16, psb,
                       SCR + OFF_AACC, 2304, 1.0, GPTO)
        elif t < NTO + 9:
            zt = t - NTO
            zo = zt * 2048 + tl.arange(0, 2048)
            tl.store(SCR + OFF_GU + zo, tl.zeros([2048], tl.float32))
        elif t < NTO + 11:
            zt = t - NTO - 9
            zo = zt * 2048 + tl.arange(0, 2048)
            zm = zo < D
            tl.store(SCR + OFF_MACC + zo, tl.zeros([2048], tl.float32), mask=zm)
        elif t < NTO + 12:
            tl.store(SCR + OFF_LOGITS + tl.arange(0, 64), tl.zeros([64], tl.float32))
        else:
            wt = t - NTO - 12
            p = wt // 4
            c4 = (wt % 4) * 1024 + tl.arange(0, 1024)
            winp = _pick3(WQL, WKL, WVL, p)
            raw = tl.load(SCR + OFF_QRAW + p * 4096 + c4,
                          cache_modifier=".cg").to(tl.bfloat16)
            o1 = tl.load(winp + 4096 + c4, cache_modifier=".cg")
            o2 = tl.load(winp + 8192 + c4, cache_modifier=".cg")
            tl.store(winp + c4, o1)
            tl.store(winp + 4096 + c4, o2)
            tl.store(winp + 8192 + c4, raw)


@triton.jit
def _moe_logits(pid, SCR, XQP, XQP16, XSP, NRM, RW, GUW, GUS, GUZ, layer,
                GPT4A: tl.constexpr, NPROG: tl.constexpr):
    """Router logits (from private normed x) + shared-expert gate/up + zero
    QRAW. Runs on h1 derived from (XA, AACC)."""
    rstd = _rstd(SCR, OFF_XA, OFF_AACC)
    _prequant(SCR, XQP, XSP, OFF_XA, OFF_AACC, rstd, NRM, 0)
    qp16 = XQP16 + pid.to(tl.int64) * (PSQ // 2)
    psb = XSP + pid.to(tl.int64) * PSF
    r128 = tl.arange(0, 128)
    NGC: tl.constexpr = GD // GPT4A
    NTL = 24
    NTS = NTL + 8 * NGC
    NT = NTS + 8
    cols0 = tl.arange(0, 256)
    for t in range(pid, NT, NPROG):
        if t < NTL:
            ob = t // 3
            kc = t % 3
            rows = ob * 8 + tl.arange(0, 8)
            acc8 = tl.zeros([8], tl.float32)
            for ch in range(6):
                cc = kc * 6 + ch
                xn = tl.load(psb + PXN + cc * 128 + r128)
                w = tl.load(RW + rows[:, None] * D + cc * 128 + r128[None, :]).to(tl.float32)
                acc8 += tl.sum(w * xn[None, :], 1)
            tl.atomic_add(SCR + OFF_LOGITS + rows, acc8)
        elif t < NTS:
            ts = t - NTL
            nb = ts // NGC
            gc = ts % NGC
            eb = (layer * 65 + 64) + tl.zeros((), dtype=tl.int64)
            _gemv_task(GUW + eb * (1152 * 2048), GUS + eb * (GD * 2048),
                       GUZ + eb * (GD * 2048), 2048, 0, gc, nb, qp16, psb,
                       SCR + OFF_GU + 8 * 2048, 2048, 1.0, GPT4A)
        else:
            zt = t - NTS
            zo = zt * 2048 + tl.arange(0, 2048)
            tl.store(SCR + OFF_QRAW + zo, tl.zeros([2048], tl.float32))


@triton.jit
def _moe_gu(pid, SCR, XQP16, XSP, GUW, GUS, GUZ, layer,
            GPT4: tl.constexpr, NPROG: tl.constexpr):
    """Routed-expert gate/up (private quant from the logits stage is reused);
    commits h1 into XB."""
    rank, probs, wsum = _topk(SCR)
    qp16 = XQP16 + pid.to(tl.int64) * (PSQ // 2)
    psb = XSP + pid.to(tl.int64) * PSF
    r64o = tl.arange(0, 64)
    NGC: tl.constexpr = GD // GPT4
    NTG = 8 * 8 * NGC
    NT = NTG + 2
    cols0 = tl.arange(0, 256)
    for t in range(pid, NT, NPROG):
        if t < NTG:
            slot = t // (8 * NGC)
            rem = t % (8 * NGC)
            nb = rem // NGC
            gc = rem % NGC
            e = tl.sum(tl.where(rank == slot, r64o, 0))
            eb = (layer * 65 + e).to(tl.int64)
            _gemv_task(GUW + eb * (1152 * 2048), GUS + eb * (GD * 2048),
                       GUZ + eb * (GD * 2048), 2048, 0, gc, nb, qp16, psb,
                       SCR + OFF_GU + slot * 2048, 2048, 1.0, GPT4)
        else:
            ct = t - NTG
            co = ct * 2048 + tl.arange(0, 2048)
            cm = co < D
            av = tl.load(SCR + OFF_XA + co, mask=cm, other=0.0, cache_modifier=".cg")
            bv = tl.load(SCR + OFF_AACC + co, mask=cm, other=0.0, cache_modifier=".cg")
            h1 = (av + bv.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16).to(tl.float32)
            tl.store(SCR + OFF_XB + co, h1, mask=cm)


@triton.jit
def _moe_down(pid, SCR, XQP, XQP16, XSP, DWW, DWS, DWZ, layer,
              GPT5: tl.constexpr, NPROG: tl.constexpr):
    """MoE down projections into MACC; zeros AACC."""
    rank, probs, wsum = _topk(SCR)
    _prequant(SCR, XQP, XSP, 0, 0, 0.0, SCR, 2)
    qp16 = XQP16 + pid.to(tl.int64) * (PSQ // 2)
    psb = XSP + pid.to(tl.int64) * PSF
    r64o = tl.arange(0, 64)
    NGC: tl.constexpr = 8 // GPT5
    NTD = 9 * 9 * NGC
    NT = NTD + 2
    cols0 = tl.arange(0, 256)
    for t in range(pid, NT, NPROG):
        if t < NTD:
            slot = t // (9 * NGC)
            rem = t % (9 * NGC)
            nb = rem // NGC
            gc = rem % NGC
            e = tl.where(slot < 8, tl.sum(tl.where(rank == slot, r64o, 0)), 64)
            ws = tl.where(slot < 8,
                          tl.sum(tl.where(rank == slot, probs, 0.0)) / (wsum + 1e-9) * 2.446,
                          1.0)
            eb = (layer * 65 + e).to(tl.int64)
            _gemv_task(DWW + eb * (512 * 2304), DWS + eb * (8 * 2304),
                       DWZ + eb * (8 * 2304), 2304, slot * 8, gc, nb,
                       qp16, psb, SCR + OFF_MACC, 2304, ws, GPT5)
        else:
            zt = t - NTD
            zo = zt * 2048 + tl.arange(0, 2048)
            zm = zo < D
            tl.store(SCR + OFF_AACC + zo, tl.zeros([2048], tl.float32), mask=zm)


# ------------------------------------------------------------------------- #
# the megakernel
# ------------------------------------------------------------------------- #
@triton.jit(do_not_specialize=["pos", "copy_flag", "bar_base"])
def _mega(
    pos, copy_flag, bar_base,
    BAR, ABORT, TIMES,
    XIN, XOUT, SCR, XQP, XQP16, XSP, QEFFB, QROPEB, SCORE,
    CKV, KR, CKV32, KR32, CKVS32, KRS32,
    S0, S1, S2,
    CQ0, CK0, CV0, CQ1, CK1, CV1, CQ2, CK2, CV2,
    KW, KS, KZ, KBETA, KCONV, KOW, KOS, KOZ,
    ANORM, MNORM, ROUTER,
    GUW, GUS, GUZ, DWW, DWS, DWZ,
    MQW, MQS, MQZ, KVBW, KVBS, KVBZ, MOW, MOS, MOZ,
    RINV,
    NPROG: tl.constexpr, GPT1: tl.constexpr, GPT4: tl.constexpr,
    GPT4A: tl.constexpr, GPT5: tl.constexpr, GPTO: tl.constexpr,
    GPTM1: tl.constexpr,
):
    pid = tl.program_id(0)
    phase = bar_base
    sid = 1
    if pid == 0:
        tl.store(TIMES + 0, tl.extra.cuda.globaltimer())
    r128 = tl.arange(0, 128)
    r64 = tl.arange(0, 64)
    r32 = tl.arange(0, 32)
    hh = tl.arange(0, 32)
    qp16 = XQP16 + pid.to(tl.int64) * (PSQ // 2)
    psb = XSP + pid.to(tl.int64) * PSF

    # ---------------- stage 0: copy-in ----------------
    if copy_flag == 1:
        NT0 = 2 + (pos * 288 + 2047) // 2048
    else:
        NT0 = 2
    for t in range(pid, NT0, NPROG):
        if t < 2:
            o2 = t * 2048 + tl.arange(0, 2048)
            xm = o2 < D
            xv = tl.load(XIN + o2, mask=xm, other=0).to(tl.float32)
            tl.store(SCR + OFF_XB + o2, xv, mask=xm)
            tl.store(SCR + OFF_MACC + o2, tl.zeros([2048], tl.float32), mask=xm)
        else:
            idx = (t - 2) * 2048 + tl.arange(0, 2048)
            ce = pos * 256
            m1 = idx < ce
            v1 = tl.load(CKVS32 + idx, mask=m1, other=0)
            tl.store(CKV32 + idx, v1, mask=m1)
            idx2 = idx - ce
            m2 = (idx >= ce) & (idx2 < pos * 32)
            v2 = tl.load(KRS32 + idx2, mask=m2, other=0)
            tl.store(KR32 + idx2, v2, mask=m2)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # ---------------- KDA layers 0..2 ----------------
    for layer in range(3):
        anrm = ANORM + layer * D
        mnrm = MNORM + layer * D
        rw = ROUTER + layer * (64 * D)
        s_ptr = _pick3(S0, S1, S2, layer)
        kwb = KW + layer.to(tl.int64) * (1152 * NQKVG)
        ksb = KS + layer * (GD * NQKVG)
        kzb = KZ + layer * (GD * NQKVG)

        # --- K1: qkvg projections (+ beta logits, + commit layer input) ---
        rstd = _rstd(SCR, OFF_XB, OFF_MACC)
        _prequant(SCR, XQP, XSP, OFF_XB, OFF_MACC, rstd, anrm, 0)
        NGC1 = GD // GPT1
        NT1G = 64 * NGC1
        NT1 = NT1G + 4 + 2
        for t in range(pid, NT1, NPROG):
            if t < NT1G:
                nb = t // NGC1
                gc = t % NGC1
                _gemv_task(kwb, ksb, kzb, NQKVG, 0, gc, nb, qp16, psb,
                           SCR + OFF_QRAW, NQKVG, 1.0, GPT1)
            elif t < NT1G + 4:
                bt = t - NT1G
                rows = bt * 8 + tl.arange(0, 8)
                acc8 = tl.zeros([8], tl.float32)
                for ch in range(GD):
                    xn = tl.load(psb + PXN + ch * 128 + r128)
                    bw = tl.load(KBETA + layer * (32 * D) + rows[:, None] * D
                                 + ch * 128 + r128[None, :]).to(tl.float32)
                    acc8 += tl.sum(bw * xn[None, :], 1)
                tl.store(SCR + OFF_BETA + rows, acc8.to(tl.bfloat16).to(tl.float32))
            else:
                ct = t - NT1G - 4
                co = ct * 2048 + tl.arange(0, 2048)
                cm = co < D
                av = tl.load(SCR + OFF_XB + co, mask=cm, other=0.0, cache_modifier=".cg")
                bv = tl.load(SCR + OFF_MACC + co, mask=cm, other=0.0, cache_modifier=".cg")
                x2 = (av + bv.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16).to(tl.float32)
                tl.store(SCR + OFF_XA + co, x2, mask=cm)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

        # --- K2: short conv (inline) + gated delta-rule state update ---
        wq_l = _pick3(CQ0, CQ1, CQ2, layer)
        wk_l = _pick3(CK0, CK1, CK2, layer)
        wv_l = _pick3(CV0, CV1, CV2, layer)
        for t in range(pid, 256, NPROG):
            h = t // 8
            dvb = t % 8
            dv = dvb * 16 + tl.arange(0, 16)
            hc = h * 128 + r128
            cwq = KCONV + (layer * 3) * 16384 + h * 128
            rq = tl.load(SCR + OFF_QRAW + hc,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            oq = tl.load(wq_l + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwq + r128) \
                + tl.load(wq_l + 4096 + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwq + 4096 + r128) \
                + tl.load(wq_l + 8192 + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwq + 8192 + r128) \
                + rq * tl.load(cwq + 12288 + r128)
            qq = (oq * tl.sigmoid(oq)).to(tl.bfloat16).to(tl.float32) * KDA_SCALE
            cwk = KCONV + (layer * 3 + 1) * 16384 + h * 128
            rk = tl.load(SCR + OFF_QRAW + 4096 + hc,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            ok_ = tl.load(wk_l + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwk + r128) \
                + tl.load(wk_l + 4096 + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwk + 4096 + r128) \
                + tl.load(wk_l + 8192 + hc, cache_modifier=".cg").to(tl.float32) * tl.load(cwk + 8192 + r128) \
                + rk * tl.load(cwk + 12288 + r128)
            kk = (ok_ * tl.sigmoid(ok_)).to(tl.bfloat16).to(tl.float32)
            vc = h * 128 + dv
            cwv = KCONV + (layer * 3 + 2) * 16384 + vc
            rv = tl.load(SCR + OFF_QRAW + 8192 + vc,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            ov = tl.load(wv_l + vc, cache_modifier=".cg").to(tl.float32) * tl.load(cwv) \
                + tl.load(wv_l + 4096 + vc, cache_modifier=".cg").to(tl.float32) * tl.load(cwv + 4096) \
                + tl.load(wv_l + 8192 + vc, cache_modifier=".cg").to(tl.float32) * tl.load(cwv + 8192) \
                + rv * tl.load(cwv + 12288)
            vv = (ov * tl.sigmoid(ov)).to(tl.bfloat16).to(tl.float32)
            gr = tl.load(SCR + OFF_QRAW + 12288 + hc,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            dec = tl.sigmoid(-gr)
            sp = s_ptr + h * 16384
            S = tl.load(sp + r128[:, None] * 128 + dv[None, :], cache_modifier=".cg")
            bl = tl.load(SCR + OFF_BETA + h, cache_modifier=".cg")
            beta = tl.sigmoid(bl)
            Sd = S * dec[:, None]
            pred = tl.sum(Sd * kk[:, None], 0)
            Sn = Sd + (beta * kk)[:, None] * (vv - pred)[None, :]
            tl.store(sp + r128[:, None] * 128 + dv[None, :], Sn)
            oo = tl.sum(Sn * qq[:, None], 0)
            tl.store(SCR + OFF_OATT + h * 128 + dv, oo)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

        # --- K3: o_proj into AACC (+ zero GU/MACC/LOGITS, + shift conv windows) ---
        _oproj(pid, SCR, XQP, XQP16, XSP, KOW + layer.to(tl.int64) * (2048 * 2304),
               KOS + layer * (32 * 2304), KOZ + layer * (32 * 2304),
               wq_l, wk_l, wv_l, 12, GPTO, NPROG)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

        # --- K4a: router logits + shared-expert gate/up (+ zero QRAW) ---
        _moe_logits(pid, SCR, XQP, XQP16, XSP, mnrm, rw, GUW, GUS, GUZ, layer, GPT4A, NPROG)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

        # --- K4b: routed gate/up (+ commit XB) ---
        _moe_gu(pid, SCR, XQP16, XSP, GUW, GUS, GUZ, layer, GPT4, NPROG)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

        # --- K5: MoE down into MACC (+ zero AACC) ---
        _moe_down(pid, SCR, XQP, XQP16, XSP, DWW, DWS, DWZ, layer, GPT5, NPROG)
        phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # ---------------- MLA layer 3 ----------------
    anrm = ANORM + 3 * D
    mnrm = MNORM + 3 * D
    rw = ROUTER + 3 * (64 * D)
    L1 = pos + 1

    # --- M1: q + kv_a projections (+ zero CTX/GMAX/DEN, + commit XA) ---
    rstd = _rstd(SCR, OFF_XB, OFF_MACC)
    _prequant(SCR, XQP, XSP, OFF_XB, OFF_MACC, rstd, anrm, 0)
    NGCM = GD // GPTM1
    NTMG = 27 * NGCM
    NTM = NTMG + 9 + 2
    for t in range(pid, NTM, NPROG):
        if t < NTMG:
            nb = t // NGCM
            gc = t % NGCM
            _gemv_task(MQW, MQS, MQZ, NMQ, 0, gc, nb, qp16, psb,
                       SCR + OFF_QRAW, NMQ, 1.0, GPTM1)
        elif t < NTMG + 9:
            zt = t - NTMG
            if zt < 8:
                zo = zt * 2048 + tl.arange(0, 2048)
                tl.store(SCR + OFF_CTX + zo, tl.zeros([2048], tl.float32))
            else:
                tl.store(SCR + OFF_GMAX + r32, tl.full([32], -1e30, tl.float32))
                tl.store(SCR + OFF_DEN + r32, tl.zeros([32], tl.float32))
        else:
            ct = t - NTMG - 9
            co = ct * 2048 + tl.arange(0, 2048)
            cm = co < D
            av = tl.load(SCR + OFF_XB + co, mask=cm, other=0.0, cache_modifier=".cg")
            bv = tl.load(SCR + OFF_MACC + co, mask=cm, other=0.0, cache_modifier=".cg")
            x2 = (av + bv.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16).to(tl.float32)
            tl.store(SCR + OFF_XA + co, x2, mask=cm)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M2: absorbed q_eff, rope, cache append ---
    r32b = tl.arange(0, 32)
    for t in range(pid, 258, NPROG):
        if t < 256:
            h = t // 8
            cb = (t % 8) // 2
            ch = t % 2
            qn = tl.load(SCR + OFF_QRAW + h * 192 + r128,
                         cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            s1v = tl.load(KVBS + cb * 8192 + h * 256 + r128).to(tl.float32)
            z1v = tl.load(KVBZ + cb * 8192 + h * 256 + r128).to(tl.float32)
            sq = s1v * qn
            t0 = tl.sum(z1v * sq)
            w1t = tl.load(KVBW + (cb * 64 + ch * 32 + r32b)[:, None] * 8192
                          + h * 256 + r128[None, :])
            qe1 = tl.sum((w1t & 15).to(tl.float32) * sq[None, :], 1) - t0
            qo1 = tl.sum((w1t >> 4).to(tl.float32) * sq[None, :], 1) - t0
            tl.store(QEFFB + h * 512 + cb * 128 + ch * 64 + 2 * r32b, qe1.to(tl.bfloat16))
            tl.store(QEFFB + h * 512 + cb * 128 + ch * 64 + 2 * r32b + 1, qo1.to(tl.bfloat16))
        elif t == 256:
            pq = SCR + OFF_QRAW + hh[:, None] * 192 + 128 + 2 * r32[None, :]
            qe2 = tl.load(pq, cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            qo2 = tl.load(pq + 1, cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            inv2 = tl.load(RINV + r32)
            ang2 = pos * inv2
            cs2 = tl.cos(ang2)[None, :]
            sn2 = tl.sin(ang2)[None, :]
            re2 = qe2 * cs2 - qo2 * sn2
            ro2 = qo2 * cs2 + qe2 * sn2
            tl.store(QROPEB + hh[:, None] * 64 + 2 * r32[None, :], re2.to(tl.bfloat16))
            tl.store(QROPEB + hh[:, None] * 64 + 2 * r32[None, :] + 1, ro2.to(tl.bfloat16))
        elif t == 257:
            ke3 = tl.load(SCR + OFF_QRAW + 6656 + 2 * r32,
                          cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            ko3 = tl.load(SCR + OFF_QRAW + 6656 + 2 * r32 + 1,
                          cache_modifier=".cg").to(tl.bfloat16).to(tl.float32)
            inv3 = tl.load(RINV + r32)
            ang3 = pos * inv3
            cs3 = tl.cos(ang3)
            sn3 = tl.sin(ang3)
            tl.store(KR + pos.to(tl.int64) * 64 + 2 * r32, (ke3 * cs3 - ko3 * sn3).to(tl.bfloat16))
            tl.store(KR + pos.to(tl.int64) * 64 + 2 * r32 + 1, (ko3 * cs3 + ke3 * sn3).to(tl.bfloat16))
            o512 = tl.arange(0, 512)
            cv3 = tl.load(SCR + OFF_QRAW + 6144 + o512, cache_modifier=".cg")
            tl.store(CKV + pos.to(tl.int64) * 512 + o512, cv3.to(tl.bfloat16))
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M3a: attention scores + per-head max ---
    NLB = (L1 + 31) // 32
    for t in range(pid, NLB, NPROG):
        rows = t * 32 + r32
        rm = rows < L1
        acc = tl.zeros([32, 32], tl.float32)
        for cc in range(4):
            qe = tl.load(QEFFB + hh[:, None] * 512 + cc * 128 + r128[None, :],
                         cache_modifier=".cg")
            cvt = tl.load(CKV + rows[:, None].to(tl.int64) * 512 + cc * 128 + r128[None, :],
                          mask=rm[:, None], other=0, cache_modifier=".cg")
            acc = tl.dot(qe, tl.trans(cvt), acc)
        qr = tl.load(QROPEB + hh[:, None] * 64 + r64[None, :], cache_modifier=".cg")
        krt = tl.load(KR + rows[:, None].to(tl.int64) * 64 + r64[None, :],
                      mask=rm[:, None], other=0, cache_modifier=".cg")
        acc = tl.dot(qr, tl.trans(krt), acc)
        sc = acc * MLA_SCALE
        scm = tl.where(rm[None, :], sc, -1e30)
        tl.atomic_max(SCR + OFF_GMAX + hh, tl.max(scm, 1))
        tl.store(SCORE + rows[:, None] * 32 + hh[None, :], tl.trans(scm), mask=rm[:, None])
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M3c: softmax-weighted context accumulation ---
    NSC = (L1 + 127) // 128
    for t in range(pid, 4 * NSC + 2, NPROG):
        if t >= 4 * NSC:
            zo5 = (t - 4 * NSC) * 2048 + tl.arange(0, 2048)
            tl.store(SCR + OFF_OATT + zo5, tl.zeros([2048], tl.float32))
        else:
            lsp = t % NSC
            cb = t // NSC
            m = tl.load(SCR + OFF_GMAX + hh, cache_modifier=".cg")
            acc = tl.zeros([32, 128], tl.float32)
            dsum = tl.zeros([32], tl.float32)
            for j in range(4):
                rows = lsp * 128 + j * 32 + r32
                rm = rows < L1
                sct = tl.load(SCORE + rows[:, None] * 32 + hh[None, :],
                              mask=rm[:, None], other=-1e30, cache_modifier=".cg")
                p = tl.exp(sct - m[None, :])
                p = tl.where(rm[:, None], p, 0.0)
                if cb == 0:
                    dsum += tl.sum(p, 0)
                cvt = tl.load(CKV + rows[:, None].to(tl.int64) * 512 + cb * 128 + r128[None, :],
                              mask=rm[:, None], other=0, cache_modifier=".cg")
                acc = tl.dot(tl.trans(p.to(tl.bfloat16)), cvt, acc)
            tl.atomic_add(SCR + OFF_CTX + hh[:, None] * 512 + cb * 128 + r128[None, :], acc)
            if cb == 0:
                tl.atomic_add(SCR + OFF_DEN + hh, dsum)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M5: absorbed value projection (partials into zeroed OATT) ---
    for t in range(pid, 256, NPROG):
        h = t // 8
        jb = (t % 8) // 4
        cc = t % 4
        den = tl.load(SCR + OFF_DEN + h, cache_modifier=".cg")
        cols = h * 256 + 128 + jb * 64 + r64
        cx = tl.load(SCR + OFF_CTX + h * 512 + cc * 128 + r128, cache_modifier=".cg") / den
        cxe, cxo = tl.split(tl.reshape(cx, (64, 2)))
        w = tl.load(KVBW + (cc * 64 + r64)[:, None] * 8192 + cols[None, :])
        sv = tl.load(KVBS + cc * 8192 + cols).to(tl.float32)
        zv = tl.load(KVBZ + cc * 8192 + cols).to(tl.float32)
        du = tl.sum((w & 15).to(tl.float32) * cxe[:, None], 0) \
             + tl.sum((w >> 4).to(tl.float32) * cxo[:, None], 0)
        acc5 = sv * du - sv * zv * tl.sum(cx)
        tl.atomic_add(SCR + OFF_OATT + h * 128 + jb * 64 + r64, acc5)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M6: o_proj into AACC (+ zero GU/MACC/LOGITS) ---
    _oproj(pid, SCR, XQP, XQP16, XSP, MOW, MOS, MOZ, CQ0, CK0, CV0, 0, GPTO, NPROG)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M7a: router logits + shared gate/up (+ zero QRAW) ---
    _moe_logits(pid, SCR, XQP, XQP16, XSP, mnrm, rw, GUW, GUS, GUZ, 3, GPT4A, NPROG)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M7b: routed gate/up (+ commit XB) ---
    _moe_gu(pid, SCR, XQP16, XSP, GUW, GUS, GUZ, 3, GPT4, NPROG)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- M8: MoE down into MACC (+ zero AACC) ---
    _moe_down(pid, SCR, XQP, XQP16, XSP, DWW, DWS, DWZ, 3, GPT5, NPROG)
    phase, sid = _bar(BAR, ABORT, TIMES, phase, sid, NPROG)

    # --- F: emit hidden = bf16(XB + bf16(MACC)) ---
    if pid == 0:
        for part in tl.static_range(2):
            o2 = part * 2048 + tl.arange(0, 2048)
            xm = o2 < D
            av = tl.load(SCR + OFF_XB + o2, mask=xm, other=0.0, cache_modifier=".cg")
            bv = tl.load(SCR + OFF_MACC + o2, mask=xm, other=0.0, cache_modifier=".cg")
            xv = (av + bv.to(tl.bfloat16).to(tl.float32)).to(tl.bfloat16)
            tl.store(XOUT + o2, xv, mask=xm)
        tl.store(TIMES + sid, tl.extra.cuda.globaltimer())


# ------------------------------------------------------------------------- #
# module skeletons (state_dict layout identical to the reference)
# ------------------------------------------------------------------------- #
class QuantLinear(nn.Module):
    def __init__(self, in_f: int, out_f: int, group: int = GROUP_SIZE):
        super().__init__()
        self.in_f, self.out_f, self.group = in_f, out_f, group
        ng = in_f // group
        self.register_buffer("w_q", torch.zeros(in_f // 2, out_f, dtype=torch.uint8))
        self.register_buffer("scales", torch.zeros(ng, out_f, dtype=torch.bfloat16))
        self.register_buffer("zeros", torch.zeros(ng, out_f, dtype=torch.bfloat16))


class QuantExperts(nn.Module):
    def __init__(self, n: int, in_f: int, out_f: int, group: int = GROUP_SIZE):
        super().__init__()
        self.n, self.in_f, self.out_f, self.group = n, in_f, out_f, group
        ng = in_f // group
        self.register_buffer("w_q", torch.zeros(n, in_f // 2, out_f, dtype=torch.uint8))
        self.register_buffer("scales", torch.zeros(n, ng, out_f, dtype=torch.bfloat16))
        self.register_buffer("zeros", torch.zeros(n, ng, out_f, dtype=torch.bfloat16))


class KDA(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        H, Dk, d = cfg.kda_heads, cfg.kda_head_dim, cfg.hidden
        self.q_proj = QuantLinear(d, H * Dk, cfg.group)
        self.k_proj = QuantLinear(d, H * Dk, cfg.group)
        self.v_proj = QuantLinear(d, H * Dk, cfg.group)
        self.g_proj = QuantLinear(d, H * Dk, cfg.group)
        self.beta_proj = nn.Linear(d, H, bias=False, dtype=cfg.dtype)
        self.conv_w = nn.Parameter(torch.empty(3, H * Dk, cfg.short_conv, dtype=cfg.dtype))
        self.o_proj = QuantLinear(H * Dk, d, cfg.group)


class MLA(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        H, d = cfg.mla_heads, cfg.hidden
        self.q_proj = QuantLinear(d, H * (cfg.qk_nope + cfg.qk_rope), cfg.group)
        self.kv_a = QuantLinear(d, cfg.kv_lora + cfg.qk_rope, cfg.group)
        self.kv_b = QuantLinear(cfg.kv_lora, H * (cfg.qk_nope + cfg.v_head), cfg.group)
        self.o_proj = QuantLinear(H * cfg.v_head, d, cfg.group)


class MoE(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        d, m, E = cfg.hidden, cfg.moe_inter, cfg.n_experts
        self.router = nn.Linear(d, E, bias=False, dtype=cfg.dtype)
        self.gate = QuantExperts(E, d, m, cfg.group)
        self.up = QuantExperts(E, d, m, cfg.group)
        self.down = QuantExperts(E, m, d, cfg.group)
        self.s_gate = QuantExperts(cfg.n_shared, d, m, cfg.group)
        self.s_up = QuantExperts(cfg.n_shared, d, m, cfg.group)
        self.s_down = QuantExperts(cfg.n_shared, m, d, cfg.group)


class Block(nn.Module):
    def __init__(self, cfg, kind):
        super().__init__()
        self.kind = kind
        self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
        self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
        self.attn = KDA(cfg) if kind == "K" else MLA(cfg)
        self.moe = MoE(cfg)


def _pick_gpt(nb, gk, nprog):
    """Largest group-chunk (task depth <= 6 groups) keeping >= ~0.75 waves."""
    for g in (6, 4, 3, 2):
        if gk % g == 0 and nb * (gk // g) >= (3 * nprog) // 4:
            return g
    return 1


class Model(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern)
        self._ready = False
        self.register_load_state_dict_post_hook(Model._invalidate)

    @staticmethod
    def _invalidate(module, incompatible_keys):
        module._ready = False

    # ------------------------------------------------------------------ #
    def _prepare(self):
        cfg = self.cfg
        assert cfg.hidden == 2304 and cfg.kda_heads == 32 and cfg.kda_head_dim == 128
        assert cfg.n_experts == 64 and cfg.moe_inter == 1024 and cfg.n_shared == 1
        assert cfg.kv_lora == 512 and cfg.qk_nope == 128 and cfg.qk_rope == 64
        dev = self.blocks[0].attn_norm.device
        assert dev.type == "cuda"
        self._dev = dev
        blks = self.blocks
        kda = [b.attn for b in blks[:3]]
        mla = blks[3].attn

        self._kw = torch.stack([torch.cat([a.q_proj.w_q, a.k_proj.w_q,
                                           a.v_proj.w_q, a.g_proj.w_q], 1) for a in kda]).contiguous()
        self._ks = torch.stack([torch.cat([a.q_proj.scales, a.k_proj.scales,
                                           a.v_proj.scales, a.g_proj.scales], 1) for a in kda]).contiguous()
        self._kz = torch.stack([torch.cat([a.q_proj.zeros, a.k_proj.zeros,
                                           a.v_proj.zeros, a.g_proj.zeros], 1) for a in kda]).contiguous()
        self._kbeta = torch.stack([a.beta_proj.weight for a in kda]).contiguous()
        self._kconv = torch.stack([a.conv_w for a in kda]).float().permute(0, 1, 3, 2).contiguous()
        self._kow = torch.stack([a.o_proj.w_q for a in kda]).contiguous()
        self._kos = torch.stack([a.o_proj.scales for a in kda]).contiguous()
        self._koz = torch.stack([a.o_proj.zeros for a in kda]).contiguous()
        self._anorm = torch.stack([b.attn_norm for b in blks]).float().contiguous()
        self._mnorm = torch.stack([b.moe_norm for b in blks]).float().contiguous()
        self._router = torch.stack([b.moe.router.weight for b in blks]).contiguous()
        self._guw = torch.stack([torch.cat([torch.cat([b.moe.gate.w_q, b.moe.up.w_q], 2),
                                            torch.cat([b.moe.s_gate.w_q, b.moe.s_up.w_q], 2)], 0)
                                 for b in blks]).contiguous()
        self._gus = torch.stack([torch.cat([torch.cat([b.moe.gate.scales, b.moe.up.scales], 2),
                                            torch.cat([b.moe.s_gate.scales, b.moe.s_up.scales], 2)], 0)
                                 for b in blks]).contiguous()
        self._guz = torch.stack([torch.cat([torch.cat([b.moe.gate.zeros, b.moe.up.zeros], 2),
                                            torch.cat([b.moe.s_gate.zeros, b.moe.s_up.zeros], 2)], 0)
                                 for b in blks]).contiguous()
        self._dww = torch.stack([torch.cat([b.moe.down.w_q, b.moe.s_down.w_q], 0)
                                 for b in blks]).contiguous()
        self._dws = torch.stack([torch.cat([b.moe.down.scales, b.moe.s_down.scales], 0)
                                 for b in blks]).contiguous()
        self._dwz = torch.stack([torch.cat([b.moe.down.zeros, b.moe.s_down.zeros], 0)
                                 for b in blks]).contiguous()
        self._mqw = torch.cat([mla.q_proj.w_q, mla.kv_a.w_q], 1).contiguous()
        self._mqs = torch.cat([mla.q_proj.scales, mla.kv_a.scales], 1).contiguous()
        self._mqz = torch.cat([mla.q_proj.zeros, mla.kv_a.zeros], 1).contiguous()
        self._kvbw = mla.kv_b.w_q.contiguous()
        self._kvbs = mla.kv_b.scales.contiguous()
        self._kvbz = mla.kv_b.zeros.contiguous()
        self._mow = mla.o_proj.w_q.contiguous()
        self._mos = mla.o_proj.scales.contiguous()
        self._moz = mla.o_proj.zeros.contiguous()
        self._rinv = (1.0 / (cfg.rope_theta ** (torch.arange(0, cfg.qk_rope, 2,
                      device=dev, dtype=torch.float32) / cfg.qk_rope))).contiguous()

        self._scr = torch.zeros(SCR_SIZE, dtype=torch.float32, device=dev)
        self._qeffb = torch.zeros(32 * 512, dtype=torch.bfloat16, device=dev)
        self._qropeb = torch.zeros(32 * 64, dtype=torch.bfloat16, device=dev)
        self._xout = torch.zeros(2304, dtype=torch.bfloat16, device=dev)
        self._abort = torch.zeros(1, dtype=torch.int32, device=dev)
        self._times = torch.zeros(64, dtype=torch.int64, device=dev)
        self._ckv = None
        self._kr = None
        self._score = None
        self._cap = 0
        self._len = -1

        sm = torch.cuda.get_device_properties(dev).multi_processor_count
        for mult, nwarp, mreg in ((2, 8, 128), (4, 4, 128), (3, 4, 168),
                                  (2, 4, 128), (2, 4, None), (1, 8, None)):
            nprog = sm * mult
            self._nprog = nprog
            self._nwarp = nwarp
            self._mreg = mreg
            self._gpt1 = _pick_gpt(64, 18, nprog)
            self._gpt4 = _pick_gpt(64, 18, nprog)
            self._gpt4a = _pick_gpt(8, 18, nprog)
            self._gpt5 = _pick_gpt(81, 8, nprog)
            self._gpto = _pick_gpt(9, 32, nprog)
            self._gptm1 = _pick_gpt(27, 18, nprog)
            self._xqp = torch.zeros(nprog * PRIV_Q, dtype=torch.int8, device=dev)
            self._xsp = torch.zeros(nprog * PRIV_F, dtype=torch.float32, device=dev)
            self._bar = torch.zeros(1, dtype=torch.int64, device=dev)
            self._base = 0
            if self._probe():
                break
        else:
            raise RuntimeError("megakernel: no co-resident grid configuration found")
        self._ready = True

    def _probe(self):
        """Dummy launch on throwaway state; verifies grid co-residency and
        calibrates the per-step barrier count."""
        dev = self._dev
        dummy_L = 384
        s = [torch.zeros(32, 128, 128, dtype=torch.float32, device=dev) for _ in range(3)]
        cw = [torch.zeros(3, 4096, dtype=torch.bfloat16, device=dev) for _ in range(9)]
        ckv = torch.zeros(dummy_L + 8, 512, dtype=torch.bfloat16, device=dev)
        kr = torch.zeros(dummy_L + 8, 64, dtype=torch.bfloat16, device=dev)
        src_c = torch.zeros(dummy_L, 512, dtype=torch.bfloat16, device=dev)
        src_k = torch.zeros(dummy_L, 64, dtype=torch.bfloat16, device=dev)
        score = torch.zeros((dummy_L + 8) * 32, dtype=torch.float32, device=dev)
        xin = torch.zeros(2304, dtype=torch.bfloat16, device=dev)
        self._abort.zero_()
        torch.cuda.synchronize()
        try:
            self._launch(xin, dummy_L, 1, s, cw, ckv, kr, src_c, src_k, score)
            torch.cuda.synchronize()
        except Exception as e:
            import os
            if os.environ.get("MEGA_DEBUG"):
                import traceback
                traceback.print_exc()
            self._last_err = e
            return False
        if int(self._abort.item()) != 0:
            return False
        cnt = int(self._bar.item())
        if cnt <= 0 or cnt % self._nprog != 0:
            return False
        self._nbar = cnt // self._nprog
        self._base = cnt
        self._scr.zero_()
        torch.cuda.synchronize()
        return True

    def _launch(self, xin, pos, copy_flag, s, cw, ckv, kr, src_c, src_k, score):
        _mega[(self._nprog,)](
            pos, copy_flag, self._base,
            self._bar, self._abort, self._times,
            xin, self._xout, self._scr, self._xqp, self._xqp.view(torch.int16), self._xsp,
            self._qeffb, self._qropeb, score,
            ckv, kr, ckv.view(torch.int32), kr.view(torch.int32),
            src_c.view(torch.int32), src_k.view(torch.int32),
            s[0], s[1], s[2],
            cw[0], cw[1], cw[2], cw[3], cw[4], cw[5], cw[6], cw[7], cw[8],
            self._kw, self._ks, self._kz, self._kbeta, self._kconv,
            self._kow, self._kos, self._koz,
            self._anorm, self._mnorm, self._router,
            self._guw, self._gus, self._guz, self._dww, self._dws, self._dwz,
            self._mqw, self._mqs, self._mqz, self._kvbw, self._kvbs, self._kvbz,
            self._mow, self._mos, self._moz,
            self._rinv,
            NPROG=self._nprog, GPT1=self._gpt1, GPT4=self._gpt4,
            GPT4A=self._gpt4a, GPT5=self._gpt5, GPTO=self._gpto,
            GPTM1=self._gptm1,
            num_warps=self._nwarp,
            **({"maxnreg": self._mreg} if self._mreg else {}),
        )

    # ------------------------------------------------------------------ #
    def step(self, hidden, state):
        if not self._ready:
            self._prepare()
        st3 = state[3]
        L = st3["c_kv"].shape[0]
        copy_flag = 0
        src_c, src_k = st3["c_kv"], st3["k_rope"]
        if (self._ckv is None or L != self._len
                or st3["c_kv"].data_ptr() != self._ckv.data_ptr()
                or st3["k_rope"].data_ptr() != self._kr.data_ptr()):
            if L + 1 > self._cap:
                self._cap = L + 1 + 256
                self._ckv = torch.empty(self._cap, 512, dtype=torch.bfloat16, device=self._dev)
                self._kr = torch.empty(self._cap, 64, dtype=torch.bfloat16, device=self._dev)
                self._score = torch.empty(self._cap * 32, dtype=torch.float32, device=self._dev)
            copy_flag = 1
        s = [state[0]["S"], state[1]["S"], state[2]["S"]]
        cw = [state[0]["cq"], state[0]["ck"], state[0]["cv"],
              state[1]["cq"], state[1]["ck"], state[1]["cv"],
              state[2]["cq"], state[2]["ck"], state[2]["cv"]]
        self._launch(hidden, L, copy_flag, s, cw, self._ckv, self._kr,
                     src_c, src_k, self._score)
        self._base += self._nbar * self._nprog
        st3["c_kv"] = self._ckv[:L + 1]
        st3["k_rope"] = self._kr[:L + 1]
        self._len = L + 1
        return self._xout, state

20260721_071218_or-fable_anthropic_claude-fable-5_02_kimi_linear_decode