KernelBench mega · RTX PRO 6000

Kimi-Linear Decode Claude Fable 5

cleandid not score

manually audited: clean

Manual audit (solution.py end-to-end + transcript scan + empirical GPU recompute test on the still-alive rented H100 + graded-path integrity). The 19.12x geomean (18.64x/18.86x/19.88x at ctx 2048/8192/16384, solution 0.643-0.688 ms/tok vs baseline 12-13.7 ms/tok) is real and in line with the same model's prior clean 18.71x cell (20260701_172615), so the magnitude is plausible, not an outlier. EMPIRICAL RECOMPUTE (audit_recompute.py run on the run's own H100 workspace): the one identity-keyed pattern — step() compares data_ptr()/length to decide whether to re-ingest the MLA cache into its internal buffer — was tested directly: after a primed step, ALL input buffers (hidden token, KDA S + conv windows, MLA c_kv/k_rope) were overwritten IN PLACE with fresh randoms (data_ptrs asserted unchanged), and the next step's output decorrelated from the old output (cos 0.0024) while matching reference.Model on the mutated inputs at cos 1.0000, with updated KDA state and MLA cache both 1.0000; the skip-copy continuation branch also matched at 1.0000 — recompute, not a stale lookup. Live rerun of the unmodified check.py on the box reproduced the archived pass (6/6 seed/ctx, out 0.9946-0.9999, S/cache 1.0000). Graded-path integrity: check.py in the run workspace is byte-identical to the deck (md5 ff21cd38e3d18b5020668a94c190b7a6), benchmark.py / reference.py / baseline.py / shapes.py / problem.yaml all byte-identical, template_mutated=false, zero KBH_NUMERIC_STRESS mentions in the transcript (mega's gate is cosine-sim; no stress-disable possible or attempted). Transcript (7229 events, 97 bash, 24 solution.py edits over 4h10m): reads of check.py/benchmark.py/problem.yaml in the first minutes are contract discovery only; no grader/template writes (all writes confined to its own solution.py, one scratch ubench, and agent memory); one failed "sudo -n true" probe for profiler counters (no-sudo, dropped). No contamination: all 81 outputs/runs references in the transcript are the run's own id, and the Brev box's archive contains ONLY this run, so no other solutions existed on the machine to copy. Approach was a real flywheel: env probe (H100 PCIe), reference/baseline timing, iterative Triton persistent-kernel dev in scratch/dbg.py with MEGA_DEBUG, repeated real check.py+benchmark.py cycles, then tuning (barrier backoff, group chunking, conv weight layout, M3c blocking) tracked against the score. The 0.9946 seed-0 cells are bf16 router-flip noise the agent root-caused in-session and correctly judged within the 0.98 tolerance, not a tolerance edit.

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