kernelbench.com

KernelBench mega · RTX PRO 6000

Kimi-Linear Decode Qwen 3.8 Max

faileddid not score

The final isolated RTX_PRO_6000 regrade of the exact archived source timed out again after 7,200 seconds. result.json records correct=false, failure_reason=check_timeout, check_exit_code=124, check_elapsed_seconds=7200, peak_fraction=null, and no benchmark exit code or elapsed time. Its regrade object identifies the sequential-isolated run on host brev-6h4h8jy86, GPU 0, NVIDIA RTX PRO 6000 Blackwell Server Edition. The replacement check.log contains only environment setup and no seed result, PASS, numerical mismatch, exception, or forbidden-import failure. Correctness therefore remains unverified rather than numerically rejected, and there is no publishable benchmark metric. Static inspection still finds that the authentic one-launch Triton megakernel launches 2*SM CTAs while sizing three per-CTA MLA scratch arrays for only SM CTAs. Both the original 1,800-second timeout and the final isolated 7,200-second timeout are consistent with the source's pathological Triton compilation history and are not evidence of reward hacking.

Kernel source (redacted)
"""Fused W4A16 decode megakernel for the Kimi-Linear hybrid unit (batch=1).

The entire per-token forward -- 3 KDA layers + 1 MLA layer, each with its
64-expert (top-8 + 1 shared) MoE FFN, RMSNorms, residuals, short causal conv,
KDA recurrent-state update, absorbed MLA latent-cache attention, router and
expert GEMVs -- runs as ONE persistent Triton kernel launch.

Design
------
* Persistent CTA grid (one CTA per SM).  The work is split into 31 phases;
  dependent phases are separated by a GPU-wide atomic spin barrier.  All CTAs
  are resident for the whole launch, so the barrier is deadlock-free.
  Barrier counters are epoch-based (target = launch_index * GRID, never
  reset), which removes any reset race.
* Every big projection is a fused int4-dequant GEMV: the packed uint8 weights
  are streamed once, unpacked and dequantized in-register per 128-wide group,
  accumulating in fp32.  No bf16 weight is ever materialized.
  Group dequant uses  y = s*(dot(xw, q) - z*sum(xw))  so the zero-point needs
  only one reduction of the activation per group.
  At init the int4 weights are repacked into tile-contiguous layout
  (tile, k-byte, col-in-tile) so each CTA streams one contiguous DRAM chunk
  (row-strided nibble rows thrash GDDR7 row buffers and cap at ~230 GB/s).
* RMSNorm is folded into each GEMV: the normalized activation is rounded to
  bf16 exactly like the reference (which feeds bf16 rmsnorm outputs to the
  linears), and the rsqrt factor is applied per element before the rounding.
* KDA short conv (kernel 4) + conv-window update are fused into the q/k/v
  GEMV epilogue of the tile that produces each channel.
* MLA is run in *absorbed* form: kv_b is never applied to the whole cache.
  QW = kv_b_nope^T q_nope (512x32) is computed once, scores come from one
  tensor-core pass over the latent cache, and the output uses
  o = (p^T c_kv) kv_b_v -- two small GEMVs instead of an L x 8192 matmul.
* Expert down projections run one (expert, output-tile) per CTA into fp32
  partials, summed by a tiny finalize phase.
* The MLA latent cache grows by geometric reallocation; the copy of the old
  rows is fused into the MLA block's first phase and amortizes to ~0.

No CUDA graphs, no torch.compile, no per-op kernel loops: step() launches
exactly one kernel.
"""
from __future__ import annotations

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

# -------------------------------------------------------------------------- #
# constants (mirror reference.py)
# -------------------------------------------------------------------------- #
GROUP = 128
MAX_ROWS = 17408               # max MLA cache rows scratch is sized for
GROW = 128                     # cache over-allocation on realloc

# kernel-visible constants (must be tl.constexpr globals)
EPS = tl.constexpr(1.0e-6)
HID = tl.constexpr(2304)
KDA_H = tl.constexpr(32)
KDA_D = tl.constexpr(128)
KDA_C = tl.constexpr(32 * 128)         # 4096
SC_KDA = tl.constexpr(128 ** -0.5)
MLA_H = tl.constexpr(32)
KV_LORA = tl.constexpr(512)
QK_NOPE = tl.constexpr(128)
QK_ROPE = tl.constexpr(64)
V_HEAD = tl.constexpr(128)
SC_MLA = tl.constexpr((128 + 64) ** -0.5)
LOG2_THETA = tl.constexpr(13.287712379549449)  # log2(10000)
N_EXPERTS = tl.constexpr(64)
N_ACTIVE = tl.constexpr(8)
MOE_M = tl.constexpr(1024)
ROUTED_SCALING = tl.constexpr(2.446)

# tile widths of the repacked int4 tables.  32 everywhere: with num_warps=1
# (the only warp count whose GEMV codegen avoids smem layout conversions),
# BN=32 x grid-stride gave the best measured streaming (763 GB/s standalone).
B_QKV = tl.constexpr(32)
B_O = tl.constexpr(32)
B_GU = tl.constexpr(32)
B_D = tl.constexpr(32)
B_KVB = tl.constexpr(32)

# tile strides (elements) of the repacked tables
TS_QKV = tl.constexpr((2304 // 2) * 32)
TS_O = tl.constexpr((4096 // 2) * 32)
TS_GU = tl.constexpr((2304 // 2) * 32)
TS_D = tl.constexpr((1024 // 2) * 32)
TS_KVB = tl.constexpr((512 // 2) * 32)
TS_QKV_SC = tl.constexpr((2304 // 128) * 32)
TS_O_SC = tl.constexpr((4096 // 128) * 32)
TS_GU_SC = tl.constexpr((2304 // 128) * 32)
TS_D_SC = tl.constexpr((1024 // 128) * 32)
TS_KVB_SC = tl.constexpr((512 // 128) * 32)
NT_QKV = tl.constexpr(4096 // 32)      # 128 tiles per projection
NT_O = tl.constexpr(2304 // 32)        # 72
NT_GU = tl.constexpr(1024 // 32)       # 32
NT_D = tl.constexpr(2304 // 32)        # 72
NT_KVB = tl.constexpr(8192 // 32)      # 256


# -------------------------------------------------------------------------- #
# Triton helpers
# -------------------------------------------------------------------------- #
@triton.jit
def _bar(bar_ptr, ph, TGT):
    """GPU-wide barrier on a monotonically increasing arrival counter.

    Slot `ph` gains GRID arrivals per launch, so the target for launch #E is
    E*GRID.  Epoch-based (never reset), so there is no reset race.
    """
    tl.atomic_add(bar_ptr + ph, 1, sem="release", scope="gpu")
    v = tl.atomic_add(bar_ptr + ph, 0, sem="acquire", scope="gpu")
    while v < TGT:
        v = tl.atomic_add(bar_ptr + ph, 0, sem="acquire", scope="gpu")


@triton.jit
def _gemv4(x_ptr, wq_t, sc_t, z_t, K: tl.constexpr, BN: tl.constexpr):
    """Fused int4-dequant GEMV on one repacked tile.

    wq_t/sc_t/z_t point at the tile's (K//2, BN) / (K//128, BN) slices.
    y[0:BN] = x @ dequant(W_tile).
    """
    acc = tl.zeros([BN], dtype=tl.float32)
    cn = tl.arange(0, BN)
    jj = tl.arange(0, 64)
    for g in range(K // 128):
        kg = g * 128
        xl = tl.load(x_ptr + kg + tl.arange(0, 128)).to(tl.float32)
        xe, xo = tl.split(tl.reshape(xl, (64, 2)))
        xs = tl.sum(xe) + tl.sum(xo)
        b = tl.load(wq_t + (kg // 2 + jj[:, None]) * BN + cn[None, :])
        lo = (b & 15).to(tl.float32)
        hi = (b >> 4).to(tl.float32)
        part = tl.sum(xe[:, None] * lo + xo[:, None] * hi, 0)
        s = tl.load(sc_t + g * BN + cn).to(tl.float32)
        z = tl.load(z_t + g * BN + cn).to(tl.float32)
        acc += s * (part - xs * z)
    return acc


@triton.jit
def _gemv4_norm(x_ptr, nw_ptr, wq_t, sc_t, z_t,
                K: tl.constexpr, BN: tl.constexpr):
    """_gemv4 with RMSNorm folded into the activation.

    Matches the reference numerics: the normalized activation is rounded to
    bf16 before the dot (the reference feeds a bf16 rmsnorm output into the
    linear).  Two passes over x: one for sum(x^2), one for the dot.
    """
    ss = 0.0
    jj = tl.arange(0, 64)
    for g in range(K // 128):
        kg = g * 128
        xl = tl.load(x_ptr + kg + tl.arange(0, 128)).to(tl.float32)
        ss += tl.sum(xl * xl)
    r = tl.rsqrt(ss / K + EPS)
    acc = tl.zeros([BN], dtype=tl.float32)
    cn = tl.arange(0, BN)
    for g in range(K // 128):
        kg = g * 128
        xl = tl.load(x_ptr + kg + tl.arange(0, 128)).to(tl.float32)
        nl = tl.load(nw_ptr + kg + tl.arange(0, 128)).to(tl.float32)
        xe, xo = tl.split(tl.reshape(xl, (64, 2)))
        ne, no = tl.split(tl.reshape(nl, (64, 2)))
        we = ((xe * r) * ne).to(tl.bfloat16).to(tl.float32)
        wo = ((xo * r) * no).to(tl.bfloat16).to(tl.float32)
        xs = tl.sum(we) + tl.sum(wo)
        b = tl.load(wq_t + (kg // 2 + jj[:, None]) * BN + cn[None, :])
        lo = (b & 15).to(tl.float32)
        hi = (b >> 4).to(tl.float32)
        part = tl.sum(we[:, None] * lo + wo[:, None] * hi, 0)
        s = tl.load(sc_t + g * BN + cn).to(tl.float32)
        z = tl.load(z_t + g * BN + cn).to(tl.float32)
        acc += s * (part - xs * z)
    return acc


# -------------------------------------------------------------------------- #
# MoE: 4 phases (router | top-k + gate/up | down partials | finalize)
# -------------------------------------------------------------------------- #
@triton.jit
def _moe_block(pid, TGT, pb, HIDDEN, BAR, MNORM, ROUT,
               WGQ, WGS, WGZ, WUQ, WUS, WUZ, WDQ, WDS, WDZ,
               LOGITS, EIDX, EW, EWR, HE, PART,
               LAST: tl.constexpr):
    # ---------------- phase 0: router logits --------------------------------
    if pid < 32:
        e0 = pid * 2
        ss = 0.0
        for k0 in range(0, HID, 256):
            kk = k0 + tl.arange(0, 256)
            x = tl.load(HIDDEN + kk).to(tl.float32)
            ss += tl.sum(x * x)
        r = tl.rsqrt(ss / HID + EPS)
        acc2 = tl.zeros([2], dtype=tl.float32)
        for k0 in range(0, HID, 256):
            kk = k0 + tl.arange(0, 256)
            x = tl.load(HIDDEN + kk).to(tl.float32)
            nw = tl.load(MNORM + kk).to(tl.float32)
            xn = ((x * r) * nw).to(tl.bfloat16).to(tl.float32)
            wrow = tl.load(ROUT + (e0 + tl.arange(0, 2))[:, None] * HID
                           + kk[None, :]).to(tl.float32)
            acc2 += tl.sum(xn[None, :] * wrow, 1)
        tl.store(LOGITS + e0 + tl.arange(0, 2),
                 acc2.to(tl.bfloat16).to(tl.float32))
    _bar(BAR, pb + 0, TGT)

    # ---------------- phase 1: top-k + expert gate/up GEMVs ------------------
    # (N_ACTIVE+1)*NT_GU = 288 units <= GRID, straight-line
    if pid < (N_ACTIVE + 1) * NT_GU:
        s = pid // NT_GU
        t = pid % NT_GU
        # redundant top-8 selection from the 64 router logits
        lg = tl.load(LOGITS + tl.arange(0, 64))
        lgm = tl.max(lg, 0)
        p = tl.exp(lg - lgm)
        p = p / tl.sum(p, 0)
        wsum = 0.0
        for j in range(N_ACTIVE):
            m = tl.max(p, 0)
            idx = tl.argmax(p, 0)
            tl.store(EIDX + j, idx)
            tl.store(EWR + j, m)
            wsum += m
            p = tl.where(tl.arange(0, 64) == idx, -1.0, p)
        for j in range(N_ACTIVE):
            raw = tl.load(EWR + j)
            tl.store(EW + j, raw / (wsum + 1e-9) * ROUTED_SCALING)
        tl.store(EIDX + N_ACTIVE, N_EXPERTS)  # shared expert = entry 64 of the
        tl.store(EW + N_ACTIVE, 1.0)          # concatenated expert tables
        eid = tl.load(EIDX + s)
        eid64 = eid.to(tl.int64)
        tt_ = eid64 * NT_GU + t
        gv = _gemv4_norm(HIDDEN, MNORM, WGQ + tt_ * TS_GU,
                         WGS + tt_ * TS_GU_SC, WGZ + tt_ * TS_GU_SC,
                         K=HID, BN=B_GU)
        uv = _gemv4_norm(HIDDEN, MNORM, WUQ + tt_ * TS_GU,
                         WUS + tt_ * TS_GU_SC, WUZ + tt_ * TS_GU_SC,
                         K=HID, BN=B_GU)
        hev = gv * tl.sigmoid(gv) * uv
        tl.store(HE + s * MOE_M + t * B_GU + tl.arange(0, B_GU), hev)
    _bar(BAR, pb + 1, TGT)

    # ---------------- phase 2: expert down GEMV partials ---------------------
    GRID = tl.num_programs(0)
    for u in range(pid, (N_ACTIVE + 1) * NT_D, GRID):
        s = u // NT_D
        t = u % NT_D
        eid = tl.load(EIDX + s)
        wgt = tl.load(EW + s)
        tt_ = eid.to(tl.int64) * NT_D + t
        acc = _gemv4(HE + s * MOE_M, WDQ + tt_ * TS_D, WDS + tt_ * TS_D_SC,
                     WDZ + tt_ * TS_D_SC, K=MOE_M, BN=B_D)
        tl.store(PART + s * HID + t * B_D + tl.arange(0, B_D), wgt * acc)
    _bar(BAR, pb + 2, TGT)

    # ---------------- phase 3: sum partials + residual -----------------------
    if pid < NT_O:
        n0 = pid * B_O
        cn = n0 + tl.arange(0, B_O)
        total = tl.load(HIDDEN + cn).to(tl.float32)
        for s in range(N_ACTIVE + 1):
            total += tl.load(PART + s * HID + cn)
        tl.store(HIDDEN + cn, total.to(tl.bfloat16))
    if LAST == 0:
        _bar(BAR, pb + 3, TGT)


# -------------------------------------------------------------------------- #
# KDA attention: 3 phases (qkvg+beta+conv | recurrence | o_proj + residual)
# -------------------------------------------------------------------------- #
@triton.jit
def _kda_attn(pid, TGT, pb, HIDDEN, BAR, ACT, BETA, O4K,
              ANORM, WQKV, SCKV, ZKV, CW, BETA_W, S, CQ, CK, CV,
              WO, SO, ZO):
    # ---------------- phase 0: q/k/v/g + beta GEMVs, fused short conv -------
    GRID = tl.num_programs(0)
    for u in range(pid, 4 * NT_QKV + 4, GRID):
        if u < 4 * NT_QKV:
            seg = u // NT_QKV
            t = u % NT_QKV
            n0 = t * B_QKV
            tb = u * TS_QKV
            stb = u * TS_QKV_SC
            y = _gemv4_norm(HIDDEN, ANORM, WQKV + tb, SCKV + stb, ZKV + stb,
                            K=HID, BN=B_QKV)
            if seg < 3:
                # short causal conv (kernel 4) on this tile's channels
                cols = n0 + tl.arange(0, B_QKV)
                yb = y.to(tl.bfloat16).to(tl.float32)  # ref rounds here
                if seg == 0:
                    WP = CQ
                elif seg == 1:
                    WP = CK
                else:
                    WP = CV
                w0 = tl.load(WP + cols).to(tl.float32)
                w1 = tl.load(WP + KDA_C + cols).to(tl.float32)
                w2 = tl.load(WP + 2 * KDA_C + cols).to(tl.float32)
                wc0 = tl.load(CW + seg * KDA_C * 4 + cols * 4 + 0).to(tl.float32)
                wc1 = tl.load(CW + seg * KDA_C * 4 + cols * 4 + 1).to(tl.float32)
                wc2 = tl.load(CW + seg * KDA_C * 4 + cols * 4 + 2).to(tl.float32)
                wc3 = tl.load(CW + seg * KDA_C * 4 + cols * 4 + 3).to(tl.float32)
                tt = w0 * wc0 + w1 * wc1 + w2 * wc2 + yb * wc3
                yc = tt * tl.sigmoid(tt)
                tl.store(ACT + seg * KDA_C + cols, yc.to(tl.bfloat16))
                tl.store(WP + cols, w1.to(tl.bfloat16))
                tl.store(WP + KDA_C + cols, w2.to(tl.bfloat16))
                tl.store(WP + 2 * KDA_C + cols, y.to(tl.bfloat16))
            else:
                tl.store(ACT + 3 * KDA_C + n0 + tl.arange(0, B_QKV),
                         y.to(tl.bfloat16))
        else:
            # beta = sigmoid(rmsnorm(x) @ beta_w^T); 4 tiles x 8 outputs
            o0 = (u - 4 * NT_QKV) * 8
            ss = 0.0
            for k0 in range(0, HID, 256):
                kk = k0 + tl.arange(0, 256)
                x = tl.load(HIDDEN + kk).to(tl.float32)
                ss += tl.sum(x * x)
            r = tl.rsqrt(ss / HID + EPS)
            accb = tl.zeros([8], dtype=tl.float32)
            for k0 in range(0, HID, 256):
                kk = k0 + tl.arange(0, 256)
                x = tl.load(HIDDEN + kk).to(tl.float32)
                nw = tl.load(ANORM + kk).to(tl.float32)
                xn = ((x * r) * nw).to(tl.bfloat16).to(tl.float32)
                wrow = tl.load(BETA_W + (o0 + tl.arange(0, 8))[:, None] * HID
                               + kk[None, :]).to(tl.float32)
                accb += tl.sum(xn[None, :] * wrow, 1)
            tl.store(BETA + o0 + tl.arange(0, 8), tl.sigmoid(accb))
    _bar(BAR, pb + 0, TGT)

    # ---------------- phase 1: gated-delta recurrence, v-sliced -------------
    if pid < 128:
        h = pid // 4
        v0 = (pid % 4) * 32
        vv = v0 + tl.arange(0, 32)
        beta_h = tl.load(BETA + h)
        pred = tl.zeros([32], dtype=tl.float32)
        for d0 in range(0, KDA_D, 32):
            dd = d0 + tl.arange(0, 32)
            graw = tl.load(ACT + 3 * KDA_C + h * KDA_D + dd).to(tl.float32)
            decay = tl.sigmoid(-graw)                     # exp(-softplus(g))
            kblk = tl.load(ACT + KDA_C + h * KDA_D + dd).to(tl.float32)
            St = tl.load(S + h * KDA_D * KDA_D + dd[:, None] * KDA_D + vv[None, :])
            Sd = St * decay[:, None]
            tl.store(S + h * KDA_D * KDA_D + dd[:, None] * KDA_D + vv[None, :], Sd)
            pred += tl.sum(Sd * kblk[:, None], 0)
        v32 = tl.load(ACT + 2 * KDA_C + h * KDA_D + vv).to(tl.float32)
        err = v32 - pred
        o = tl.zeros([32], dtype=tl.float32)
        for d0 in range(0, KDA_D, 32):
            dd = d0 + tl.arange(0, 32)
            qblk = tl.load(ACT + h * KDA_D + dd).to(tl.float32) * SC_KDA
            kblk = tl.load(ACT + KDA_C + h * KDA_D + dd).to(tl.float32)
            Sd = tl.load(S + h * KDA_D * KDA_D + dd[:, None] * KDA_D + vv[None, :])
            Sn = Sd + beta_h * kblk[:, None] * err[None, :]
            tl.store(S + h * KDA_D * KDA_D + dd[:, None] * KDA_D + vv[None, :], Sn)
            o += tl.sum(Sn * qblk[:, None], 0)
        tl.store(O4K + h * KDA_D + vv, o.to(tl.bfloat16))
    _bar(BAR, pb + 1, TGT)

    # ---------------- phase 2: o_proj GEMV + residual -----------------------
    if pid < NT_O:
        acc = _gemv4(O4K, WO + pid * TS_O, SO + pid * TS_O_SC,
                     ZO + pid * TS_O_SC, K=KDA_C, BN=B_O)
        n0 = pid * B_O
        x32 = tl.load(HIDDEN + n0 + tl.arange(0, B_O)).to(tl.float32)
        tl.store(HIDDEN + n0 + tl.arange(0, B_O), (acc + x32).to(tl.bfloat16))
    _bar(BAR, pb + 2, TGT)


# -------------------------------------------------------------------------- #
# MLA attention: 6 phases
# -------------------------------------------------------------------------- #
@triton.jit
def _mla_attn(pid, GRID, TGT, pb, HIDDEN, BAR, QKA, QW, QR, SCORES, PMAX,
              PSUM, PVPART, PVSC, O4K, ANORM,
              WQ, SQ, ZQ, WA, SA, ZA, WB, SB, ZB, WO, SO, ZO,
              COLD, CNEW, KOLD, KNEW, L, R, COPY):
    # ---------------- phase 0: cache copy (on realloc) + q/kv_a GEMVs -------
    if COPY == 1:
        ctot = L * KV_LORA
        for off in range(pid * 2048, ctot, GRID * 2048):
            ii = off + tl.arange(0, 2048)
            tl.store(CNEW + ii, tl.load(COLD + ii, mask=ii < ctot, other=0.0),
                     mask=ii < ctot)
        ktot = L * QK_ROPE
        for off in range(pid * 1024, ktot, GRID * 1024):
            ii = off + tl.arange(0, 1024)
            tl.store(KNEW + ii, tl.load(KOLD + ii, mask=ii < ktot, other=0.0),
                     mask=ii < ktot)
    # NQ + 18 = 210 units <= GRID, straight-line
    NQ: tl.constexpr = MLA_H * (QK_NOPE + QK_ROPE) // 32   # 192 q tiles
    if pid < NQ + 18:
        if pid < NQ:
            y = _gemv4_norm(HIDDEN, ANORM, WQ + pid * TS_QKV,
                            SQ + pid * TS_QKV_SC, ZQ + pid * TS_QKV_SC,
                            K=HID, BN=B_QKV)
            tl.store(QKA + pid * B_QKV + tl.arange(0, B_QKV), y.to(tl.bfloat16))
        else:
            t = pid - NQ
            y = _gemv4_norm(HIDDEN, ANORM, WA + t * TS_QKV, SA + t * TS_QKV_SC,
                            ZA + t * TS_QKV_SC, K=HID, BN=B_QKV)
            tl.store(QKA + MLA_H * (QK_NOPE + QK_ROPE) + t * B_QKV
                     + tl.arange(0, B_QKV), y.to(tl.bfloat16))
    _bar(BAR, pb + 0, TGT)

    # ---------------- phase 1: rope + cache append + absorb QW --------------
    if pid < 4:
        ii = tl.arange(0, 32)
        inv = tl.exp2(-(2.0 * ii / QK_ROPE) * LOG2_THETA)
        ang = L.to(tl.float32) * inv
        c32 = tl.cos(ang)
        s32 = tl.sin(ang)
        for hh in range(8):
            h = pid * 8 + hh
            qb = QKA + h * (QK_NOPE + QK_ROPE) + QK_NOPE
            qe = tl.load(qb + 2 * ii).to(tl.float32)
            qo = tl.load(qb + 2 * ii + 1).to(tl.float32)
            tl.store(QR + h * QK_ROPE + 2 * ii, (qe * c32 - qo * s32).to(tl.bfloat16))
            tl.store(QR + h * QK_ROPE + 2 * ii + 1, (qo * c32 + qe * s32).to(tl.bfloat16))
        if pid == 0:
            # append new token's latent + rope key to the cache
            for c0 in range(0, KV_LORA, 256):
                cc = c0 + tl.arange(0, 256)
                tl.store(CNEW + L * KV_LORA + cc,
                         tl.load(QKA + MLA_H * (QK_NOPE + QK_ROPE) + cc))
            kb = QKA + MLA_H * (QK_NOPE + QK_ROPE) + KV_LORA
            ke = tl.load(kb + 2 * ii).to(tl.float32)
            ko = tl.load(kb + 2 * ii + 1).to(tl.float32)
            tl.store(KNEW + L * QK_ROPE + 2 * ii, (ke * c32 - ko * s32).to(tl.bfloat16))
            tl.store(KNEW + L * QK_ROPE + 2 * ii + 1, (ko * c32 + ke * s32).to(tl.bfloat16))
    elif pid < 132:
        # QW[c, h] = sum_d q_nope[h,d] * kv_b[c, h*256+d]  (absorbed).
        # kv_b is repacked in 32-col tiles; head h occupies tiles h*8..h*8+7
        # (nope half = tiles 0..3, v half = tiles 4..7 within the head).
        qt = pid - 4
        h = qt // 4
        cg = qt % 4                      # c-group of 128 latents
        c0 = cg * 128
        j64 = tl.arange(0, 64)
        d32 = tl.arange(0, 32)
        acce = tl.zeros([64], dtype=tl.float32)
        acco = tl.zeros([64], dtype=tl.float32)
        for ds in range(4):
            q32 = tl.load(QKA + h * (QK_NOPE + QK_ROPE) + ds * 32 + d32).to(tl.float32)
            soff = (h * 8 + ds) * TS_KVB_SC + cg * B_KVB
            srow = tl.load(SB + soff + d32).to(tl.float32)
            zrow = tl.load(ZB + soff + d32).to(tl.float32)
            qs = q32 * srow
            bias = tl.sum(qs * zrow)
            b = tl.load(WB + (h * 8 + ds) * TS_KVB + (cg * 64 + j64[:, None]) * B_KVB
                        + d32[None, :])
            lo = (b & 15).to(tl.float32)
            hi = (b >> 4).to(tl.float32)
            acce += tl.sum(lo * qs[None, :], 1) - bias
            acco += tl.sum(hi * qs[None, :], 1) - bias
        tl.store(QW + (c0 + 2 * j64) * MLA_H + h, acce.to(tl.bfloat16))
        tl.store(QW + (c0 + 2 * j64 + 1) * MLA_H + h, acco.to(tl.bfloat16))
    _bar(BAR, pb + 1, TGT)

    # ---------------- phase 2: scores pass A (dot + online softmax) ---------
    # QW is split along the latent dim so each dot's operands stay small
    # (shared memory must fit 2 CTAs/SM for the GEMV phases' grid).
    c256 = tl.arange(0, 256)
    h32 = tl.arange(0, 32)
    qw0 = tl.load(QW + c256[:, None] * MLA_H + h32[None, :])
    qw1 = tl.load(QW + (256 + c256[:, None]) * MLA_H + h32[None, :])
    qr = tl.load(QR + h32[:, None] * QK_ROPE + tl.arange(0, 64)[None, :])
    qrt = tl.trans(qr)
    m_run = tl.full([32], -1.0e30, dtype=tl.float32)
    l_run = tl.zeros([32], dtype=tl.float32)
    for row0 in range(pid * 16, R, GRID * 16):
        rows = row0 + tl.arange(0, 16)
        valid = rows < R
        ca = tl.load(CNEW + rows[:, None] * KV_LORA + c256[None, :],
                     mask=valid[:, None], other=0.0)
        cb = tl.load(CNEW + rows[:, None] * KV_LORA + 256 + c256[None, :],
                     mask=valid[:, None], other=0.0)
        sn = tl.dot(ca, qw0, out_dtype=tl.float32) + tl.dot(cb, qw1, out_dtype=tl.float32)
        krr = tl.load(KNEW + rows[:, None] * QK_ROPE + tl.arange(0, 64)[None, :],
                      mask=valid[:, None], other=0.0)
        sr = tl.dot(krr, qrt, out_dtype=tl.float32)
        sc = (sn + sr) * SC_MLA
        sc = tl.where(valid[:, None], sc, -1.0e30)
        bm = tl.max(sc, 0)
        m2 = tl.maximum(m_run, bm)
        l_run = l_run * tl.exp(m_run - m2) + tl.sum(tl.exp(sc - m2[None, :]), 0)
        m_run = m2
        tl.store(SCORES + rows[:, None] * MLA_H + h32[None, :],
                 sc.to(tl.bfloat16), mask=valid[:, None])
    tl.store(PMAX + pid * MLA_H + h32, m_run)
    tl.store(PSUM + pid * MLA_H + h32, l_run)
    _bar(BAR, pb + 2, TGT)

    # ---------------- phase 3: softmax normalize + p^T c_kv partials --------
    pii = tl.arange(0, 512)
    pmask = pii < GRID
    pm = tl.load(PMAX + pii[:, None] * MLA_H + h32[None, :],
                 mask=pmask[:, None], other=-1.0e30)
    psm = tl.load(PSUM + pii[:, None] * MLA_H + h32[None, :],
                  mask=pmask[:, None], other=0.0)
    M32 = tl.max(pm, 0)
    tot = tl.sum(psm * tl.exp(pm - M32[None, :]), 0)
    inv32 = 1.0 / tot
    pv_acc = tl.zeros([32, 512], dtype=tl.float32)
    for row0 in range(pid * 16, R, GRID * 16):
        rows = row0 + tl.arange(0, 16)
        valid = rows < R
        sc = tl.load(SCORES + rows[:, None] * MLA_H + h32[None, :],
                     mask=valid[:, None], other=-1.0e30).to(tl.float32)
        p = tl.exp(sc - M32[None, :]) * inv32[None, :]
        p = tl.where(valid[:, None], p, 0.0).to(tl.bfloat16)
        cch = tl.load(CNEW + rows[:, None] * KV_LORA + tl.arange(0, 512)[None, :],
                      mask=valid[:, None], other=0.0)
        pv_acc += tl.dot(tl.trans(p), cch, out_dtype=tl.float32)
    tl.store(PVPART + pid * MLA_H * KV_LORA + tl.arange(0, 32)[:, None] * KV_LORA
             + tl.arange(0, 512)[None, :], pv_acc)
    _bar(BAR, pb + 3, TGT)

    # ---------------- phase 4: reduce partials -> pv per head ---------------
    if pid < MLA_H:
        h = pid
        pv = tl.zeros([512], dtype=tl.float32)
        for pc in range(0, 512, 32):
            pr = pc + tl.arange(0, 32)
            blk = tl.load(PVPART + pr[:, None] * MLA_H * KV_LORA + h * KV_LORA
                          + tl.arange(0, 512)[None, :],
                          mask=(pr[:, None] < GRID), other=0.0)
            pv += tl.sum(blk, 0)
        tl.store(PVSC + pid * 512 + tl.arange(0, 512), pv.to(tl.bfloat16))
    _bar(BAR, pb + 4, TGT)

    # ---------------- phase 5: o-GEMV (kv_b v halves) ------------------------
    for u in range(pid, MLA_H * 4, GRID):
        h = u // 4
        ds = u % 4
        tile = h * 8 + 4 + ds            # head h's v half = tiles h*8+4..+7
        acc = _gemv4(PVSC + h * 512, WB + tile * TS_KVB, SB + tile * TS_KVB_SC,
                     ZB + tile * TS_KVB_SC, K=KV_LORA, BN=B_KVB)
        tl.store(O4K + h * V_HEAD + ds * B_KVB + tl.arange(0, B_KVB),
                 acc.to(tl.bfloat16))
    _bar(BAR, pb + 5, TGT)

    # ---------------- phase 6: o_proj GEMV + residual ------------------------
    for u in range(pid, NT_O, GRID):
        acc = _gemv4(O4K, WO + u * TS_O, SO + u * TS_O_SC,
                     ZO + u * TS_O_SC, K=MLA_H * V_HEAD, BN=B_O)
        n0 = u * B_O
        x32 = tl.load(HIDDEN + n0 + tl.arange(0, B_O)).to(tl.float32)
        tl.store(HIDDEN + n0 + tl.arange(0, B_O), (acc + x32).to(tl.bfloat16))
    _bar(BAR, pb + 6, TGT)


# -------------------------------------------------------------------------- #
# the megakernel
# -------------------------------------------------------------------------- #
@triton.jit(do_not_specialize=["L", "R", "COPY", "EPOCH"])
def _mega_kernel(HIDDEN, BAR,
                 ACT, BETA, O4K, LOGITS, EIDX, EW, EWR, HE, PART,
                 QKA, QW, QR, SCORES, PMAX, PSUM, PVPART, PVSC,
                 AN0, MN0, WQKV0, SCKV0, ZKV0, CW0, BETA_W0,
                 S0, CQ0, CK0, CV0, WO0, SO0, ZO0,
                 ROUT0, WGQ0, WGS0, WGZ0, WUQ0, WUS0, WUZ0, WDQ0, WDS0, WDZ0,
                 AN1, MN1, WQKV1, SCKV1, ZKV1, CW1, BETA_W1,
                 S1, CQ1, CK1, CV1, WO1, SO1, ZO1,
                 ROUT1, WGQ1, WGS1, WGZ1, WUQ1, WUS1, WUZ1, WDQ1, WDS1, WDZ1,
                 AN2, MN2, WQKV2, SCKV2, ZKV2, CW2, BETA_W2,
                 S2, CQ2, CK2, CV2, WO2, SO2, ZO2,
                 ROUT2, WGQ2, WGS2, WGZ2, WUQ2, WUS2, WUZ2, WDQ2, WDS2, WDZ2,
                 AN3, MN3, WQ3, SQ3, ZQ3, WA3, SA3, ZA3,
                 WB3, SB3, ZB3, WO3, SO3, ZO3,
                 ROUT3, WGQ3, WGS3, WGZ3, WUQ3, WUS3, WUZ3, WDQ3, WDS3, WDZ3,
                 COLD, CNEW, KOLD, KNEW, L, R, COPY, EPOCH, TS):
    pid = tl.program_id(0)
    GRID = tl.num_programs(0)
    TGT = EPOCH * GRID
    _kda_attn(pid, TGT, 0, HIDDEN, BAR, ACT, BETA, O4K, AN0,
              WQKV0, SCKV0, ZKV0, CW0, BETA_W0,
              S0, CQ0, CK0, CV0, WO0, SO0, ZO0)
    _moe_block(pid, TGT, 3, HIDDEN, BAR, MN0, ROUT0,
               WGQ0, WGS0, WGZ0, WUQ0, WUS0, WUZ0, WDQ0, WDS0, WDZ0,
               LOGITS, EIDX, EW, EWR, HE, PART, LAST=0)
    _kda_attn(pid, TGT, 7, HIDDEN, BAR, ACT, BETA, O4K, AN1,
              WQKV1, SCKV1, ZKV1, CW1, BETA_W1,
              S1, CQ1, CK1, CV1, WO1, SO1, ZO1)
    _moe_block(pid, TGT, 10, HIDDEN, BAR, MN1, ROUT1,
               WGQ1, WGS1, WGZ1, WUQ1, WUS1, WUZ1, WDQ1, WDS1, WDZ1,
               LOGITS, EIDX, EW, EWR, HE, PART, LAST=0)
    _kda_attn(pid, TGT, 14, HIDDEN, BAR, ACT, BETA, O4K, AN2,
              WQKV2, SCKV2, ZKV2, CW2, BETA_W2,
              S2, CQ2, CK2, CV2, WO2, SO2, ZO2)
    _moe_block(pid, TGT, 17, HIDDEN, BAR, MN2, ROUT2,
               WGQ2, WGS2, WGZ2, WUQ2, WUS2, WUZ2, WDQ2, WDS2, WDZ2,
               LOGITS, EIDX, EW, EWR, HE, PART, LAST=0)
    _mla_attn(pid, GRID, TGT, 21, HIDDEN, BAR, QKA, QW, QR, SCORES, PMAX,
              PSUM, PVPART, PVSC, O4K, AN3,
              WQ3, SQ3, ZQ3, WA3, SA3, ZA3,
              WB3, SB3, ZB3, WO3, SO3, ZO3,
              COLD, CNEW, KOLD, KNEW, L, R, COPY)
    _moe_block(pid, TGT, 28, HIDDEN, BAR, MN3, ROUT3,
               WGQ3, WGS3, WGZ3, WUQ3, WUS3, WUZ3, WDQ3, WDS3, WDZ3,
               LOGITS, EIDX, EW, EWR, HE, PART, LAST=1)


# -------------------------------------------------------------------------- #
# modules mirroring reference.py's state_dict layout exactly
# -------------------------------------------------------------------------- #
class QuantLinear(nn.Module):
    def __init__(self, in_f: int, out_f: int, group: int = GROUP):
        super().__init__()
        self.in_f, self.out_f, self.group = in_f, out_f, group
        self.register_buffer("w_q", torch.zeros(in_f // 2, out_f, dtype=torch.uint8))
        self.register_buffer("scales", torch.zeros(in_f // group, out_f, dtype=torch.bfloat16))
        self.register_buffer("zeros", torch.zeros(in_f // group, out_f, dtype=torch.bfloat16))


class QuantExperts(nn.Module):
    def __init__(self, n: int, in_f: int, out_f: int, group: int = GROUP):
        super().__init__()
        self.n, self.in_f, self.out_f, self.group = n, in_f, out_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, in_f // group, out_f, dtype=torch.bfloat16))
        self.register_buffer("zeros", torch.zeros(n, in_f // group, out_f, dtype=torch.bfloat16))


class KDA(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        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__()
        self.cfg = cfg
        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__()
        self.cfg = cfg
        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: str):
        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)


# -------------------------------------------------------------------------- #
# int4 repacking: (K//2, N) -> (N//BN, K//2, BN) tile-contiguous
# -------------------------------------------------------------------------- #
def _repack(w_q: torch.Tensor, scales: torch.Tensor, zeros: torch.Tensor,
            bn: int):
    K2, N = w_q.shape
    t = N // bn
    w = w_q.view(K2, t, bn).permute(1, 0, 2).contiguous()
    s = scales.view(scales.shape[0], t, bn).permute(1, 0, 2).contiguous()
    z = zeros.view(zeros.shape[0], t, bn).permute(1, 0, 2).contiguous()
    return w, s, z


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._epoch = 0

    # ------------------------------------------------------------------ #
    def _prepare(self, device):
        """One-time scratch allocation + weight repacking."""
        assert self.cfg.n_experts == N_EXPERTS.value
        assert self.cfg.hidden == HID.value
        bf = torch.bfloat16
        f = torch.float32
        grid = torch.cuda.get_device_properties(device).multi_processor_count
        self._grid = 2 * grid
        self._scratch = dict(
            act=torch.zeros(4 * 4096, dtype=bf, device=device),
            beta=torch.zeros(32, dtype=f, device=device),
            o4k=torch.zeros(4096, dtype=bf, device=device),
            logits=torch.zeros(64, dtype=f, device=device),
            eidx=torch.zeros(9, dtype=torch.int32, device=device),
            ew=torch.zeros(9, dtype=f, device=device),
            ewr=torch.zeros(8, dtype=f, device=device),
            he=torch.zeros(9 * 1024, dtype=f, device=device),
            part=torch.zeros(9 * 2304, dtype=f, device=device),
            qka=torch.zeros(6144 + 576, dtype=bf, device=device),
            qw=torch.zeros(512 * 32, dtype=bf, device=device),
            qr=torch.zeros(32 * 64, dtype=bf, device=device),
            scores=torch.zeros(MAX_ROWS * 32, dtype=bf, device=device),
            pmax=torch.zeros(grid * 32, dtype=f, device=device),
            psum=torch.zeros(grid * 32, dtype=f, device=device),
            pvpart=torch.zeros(grid * 32 * 512, dtype=f, device=device),
            pvsc=torch.zeros(32 * 512, dtype=bf, device=device),
            bar=torch.zeros(40, dtype=torch.int32, device=device),
            ts=torch.zeros(2 * 40 * grid, dtype=torch.int64, device=device),
        )
        # Repacked int4 tables (tile-contiguous for DRAM streaming).
        P = []
        for blk in self.blocks:
            p = {}
            if blk.kind == "K":
                a = blk.attn
                wq = torch.cat([a.q_proj.w_q, a.k_proj.w_q, a.v_proj.w_q,
                                a.g_proj.w_q], 1)
                s = torch.cat([a.q_proj.scales, a.k_proj.scales,
                               a.v_proj.scales, a.g_proj.scales], 1)
                z = torch.cat([a.q_proj.zeros, a.k_proj.zeros,
                               a.v_proj.zeros, a.g_proj.zeros], 1)
                p["qkv"] = _repack(wq, s, z, 32)
                p["o"] = _repack(a.o_proj.w_q, a.o_proj.scales,
                                 a.o_proj.zeros, 32)
            else:
                a = blk.attn
                p["q"] = _repack(a.q_proj.w_q, a.q_proj.scales,
                                 a.q_proj.zeros, 64)
                p["a"] = _repack(a.kv_a.w_q, a.kv_a.scales, a.kv_a.zeros, 32)
                p["b"] = _repack(a.kv_b.w_q, a.kv_b.scales, a.kv_b.zeros, 32)
                p["o"] = _repack(a.o_proj.w_q, a.o_proj.scales,
                                 a.o_proj.zeros, 32)
            moe = blk.moe
            for nm in ("gate", "up", "down"):
                rq = getattr(moe, nm)
                sq = getattr(moe, "s_" + nm)
                wq = torch.cat([rq.w_q, sq.w_q], 0)
                s = torch.cat([rq.scales, sq.scales], 0)
                z = torch.cat([rq.zeros, sq.zeros], 0)
                bn = 32
                E = wq.shape[0]
                K2 = wq.shape[1]
                N = wq.shape[2]
                t = N // bn
                w = wq.view(E, K2, t, bn).permute(0, 2, 1, 3).contiguous()
                sg = s.shape[1]
                ss_ = s.view(E, sg, t, bn).permute(0, 2, 1, 3).contiguous()
                zz = z.view(E, sg, t, bn).permute(0, 2, 1, 3).contiguous()
                p[nm] = (w, ss_, zz)
            P.append(p)
        self._P = P
        self._ready = True

    # ------------------------------------------------------------------ #
    def step(self, hidden, state):
        if not self._ready:
            self._prepare(hidden.device)
        self._epoch += 1
        sc = self._scratch
        dev = hidden.device
        st = state[3]
        c = st["c_kv"]
        kr = st["k_rope"]
        L = c.shape[0]
        assert L + 1 <= MAX_ROWS, "cache exceeds preallocated scratch"
        # grow the latent cache only when the backing storage is full
        ccap = c.untyped_storage().nbytes() // c.element_size()
        if c.storage_offset() == 0 and c.is_contiguous() and ccap >= (L + 1) * 512:
            cnew, copy = c, 0
        else:
            cnew, copy = torch.empty(L + GROW, 512, dtype=c.dtype, device=dev), 1
        kcap = kr.untyped_storage().nbytes() // kr.element_size()
        if kr.storage_offset() == 0 and kr.is_contiguous() and kcap >= (L + 1) * 64:
            knew = kr
        else:
            knew = torch.empty(L + GROW, 64, dtype=kr.dtype, device=dev)
        st["c_kv"] = cnew[: L + 1]
        st["k_rope"] = knew[: L + 1]

        b = self.blocks
        P = self._P
        kargs = []
        for i in (0, 1, 2):
            m = b[i].moe
            kargs += [
                b[i].attn_norm, b[i].moe_norm,
                P[i]["qkv"][0], P[i]["qkv"][1], P[i]["qkv"][2],
                b[i].attn.conv_w, b[i].attn.beta_proj.weight,
                state[i]["S"], state[i]["cq"], state[i]["ck"], state[i]["cv"],
                P[i]["o"][0], P[i]["o"][1], P[i]["o"][2],
                m.router.weight,
                P[i]["gate"][0], P[i]["gate"][1], P[i]["gate"][2],
                P[i]["up"][0], P[i]["up"][1], P[i]["up"][2],
                P[i]["down"][0], P[i]["down"][1], P[i]["down"][2],
            ]
        m = b[3].moe
        margs = [
            b[3].attn_norm, b[3].moe_norm,
            P[3]["q"][0], P[3]["q"][1], P[3]["q"][2],
            P[3]["a"][0], P[3]["a"][1], P[3]["a"][2],
            P[3]["b"][0], P[3]["b"][1], P[3]["b"][2],
            P[3]["o"][0], P[3]["o"][1], P[3]["o"][2],
            m.router.weight,
            P[3]["gate"][0], P[3]["gate"][1], P[3]["gate"][2],
            P[3]["up"][0], P[3]["up"][1], P[3]["up"][2],
            P[3]["down"][0], P[3]["down"][1], P[3]["down"][2],
        ]
        _mega_kernel[(self._grid,)](
            hidden, sc["bar"],
            sc["act"], sc["beta"], sc["o4k"], sc["logits"], sc["eidx"],
            sc["ew"], sc["ewr"], sc["he"], sc["part"],
            sc["qka"], sc["qw"], sc["qr"], sc["scores"], sc["pmax"],
            sc["psum"], sc["pvpart"], sc["pvsc"],
            *kargs, *margs,
            c, cnew, kr, knew, L, L + 1, copy, self._epoch, sc["ts"],
            num_warps=1, num_stages=1,
        )
        return hidden, state

20260803_214551_or-fable_qwen_qwen3.8-max_02_kimi_linear_decode