kernelbench.com

KernelBench mega · H100

Kimi-Linear Decode Qwen 3.8 Max

wrongdid not score

Preserve correct=false and failure_reason=check_failed. This is an authentic one-launch Triton megakernel and a genuine optimization attempt, not a constant-output, graph-replay, or evaluator-gaming path. The official checker accepted the first four seed/context trials, then rejected seed 2 at context 2048 because mla_cache cosine was 0.0382, below 0.98. The failure is a concrete host-side state-cache identity bug: Model.step keys a persistent capacity buffer only by id(state), so a newly allocated checker state can reuse a dead list's Python id and silently inherit old MLA cache rows. No benchmark ran and no performance number is publishable.

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

The entire per-token forward -- 3 KDA layers + 1 MLA layer, each followed by a
64-expert (top-8 + 1 shared) MoE FFN, with RMSNorms and residuals -- runs as a
SINGLE persistent Triton kernel launch.  No CUDA graphs, no torch.compile, no
per-op kernel loops: one @triton.jit grid, invoked exactly once per step().

Design
------
The GPU is launched as `grid = #SMs` resident CTAs that walk a static task
graph in phases; phases are separated by global spin barriers (an atomic
counter per phase slot, monotonically increasing target).  Inside a phase, CTAs
take tasks in grid-stride order.

* int4 dequant-GEMV tiles.  The packed uint8 weight tile is loaded straight
  from a flat int4 arena, unpacked with bit ops, dequantized per-group (128)
  exactly like the reference (w = (q - z) * s, bf16 rounding), and dotted with
  the activation in fp32.  The bf16 weight is never materialized: each weight
  byte is read exactly once per token.
* RMSNorm is folded into the GEMVs that consume it (each task recomputes the
  scalar rsqrt(mean(x^2)+eps) from x, which it streams anyway).
* KDA: one task per head fuses the kernel-4 causal depthwise conv (with the
  window-state update) and the gated-delta recurrence over S[h] (128x128 fp32).
* MLA: "absorbed" per-head decode attention.  Instead of materializing
  cache @ W_kv_b (L x 8192), we absorb W_kv_b into the query per head:
  u_h = W_b_nope_h^T q_nope_h (512-d), and score every cached row with
  score[l,h] = c_l . u_h + kro_l . qrope_h.  A first streaming pass computes
  per-(chunk,h) softmax statistics, a reduce pass combines them, and a second
  streaming pass accumulates the softmax-weighted latent c_bar_h, which is then
  projected through the value half of W_kv_b and o_proj.  The cache is
  streamed twice (O(L*512)) instead of writing/reading the L x 8192 kv_b
  output (O(L*8192)) that the naive path materializes.
* MoE: router bf16 GEMV -> softmax/top-8 inside the kernel -> expert GEMVs
  gather from a packed expert arena (64 routed + the shared expert appended as
  expert 64), so selection is pure pointer arithmetic.  gate/up are fused per
  output tile (silu(x@Wg)*(x@Wu)); down-projections atomicAdd their weighted
  outputs onto the running residual.

State handling: KDA S and conv windows are updated in place.  The MLA latent
cache grows; step() keeps an over-allocated capacity buffer per state object
and the kernel copies old rows into it on the first step for that state
(inside the same megakernel), then appends in place forever after.  All timed
steady-state steps are exactly one kernel launch.
"""
from __future__ import annotations

from dataclasses import dataclass, field

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"]
EPS = 1.0e-6
GROUP_SIZE = 128


@dataclass(frozen=True)
class Config:
    hidden: int = 2304
    kda_heads: int = 32
    kda_head_dim: int = 128
    short_conv: int = 4
    mla_heads: int = 32
    kv_lora: int = 512
    qk_nope: int = 128
    qk_rope: int = 64
    v_head: int = 128
    rope_theta: float = 10000.0
    n_experts: int = 64
    n_active: int = 8
    n_shared: int = 1
    moe_inter: int = 1024
    routed_scaling: float = 2.446
    group: int = 128
    pattern: tuple = ("K", "K", "K", "M")
    dtype: torch.dtype = field(default=torch.bfloat16)


def build_config(shape: dict) -> Config:
    return Config(n_experts=int(shape.get("n_experts", 64)))


# --------------------------------------------------------------------------- #
# weight containers (identical state_dict layout to the reference)
# --------------------------------------------------------------------------- #
def _pack_int4(w_q: torch.Tensor) -> torch.Tensor:
    lo = w_q[0::2] & 0xF
    hi = w_q[1::2] & 0xF
    return (lo | (hi << 4)).contiguous()


def _quantize(w: torch.Tensor, group: int):
    K, N = w.shape
    wg = w.view(K // group, group, N).float()
    wmin = wg.min(dim=1, keepdim=True).values
    wmax = wg.max(dim=1, keepdim=True).values
    scales = (wmax - wmin).clamp_min(1e-8) / 15.0
    zeros = (-wmin / scales).round().clamp(0, 15)
    wq = ((wg / scales) + zeros).round().clamp(0, 15).to(torch.uint8).view(K, N)
    return _pack_int4(wq), scales.squeeze(1).to(torch.bfloat16), zeros.squeeze(1).to(torch.bfloat16)


class QuantLinear(nn.Module):
    def __init__(self, in_f: int, out_f: int, group: int = GROUP_SIZE):
        super().__init__()
        assert in_f % group == 0 and in_f % 2 == 0
        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))

    def init_random(self, gen: torch.Generator, std: float = 0.02) -> None:
        w = torch.randn(self.in_f, self.out_f, generator=gen) * std
        wq, s, z = _quantize(w, self.group)
        self.w_q.copy_(wq)
        self.scales.copy_(s)
        self.zeros.copy_(z)


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))

    def init_random(self, gen: torch.Generator, std: float = 0.02) -> None:
        for e in range(self.n):
            w = torch.randn(self.in_f, self.out_f, generator=gen) * std
            wq, s, z = _quantize(w, self.group)
            self.w_q[e].copy_(wq)
            self.scales[e].copy_(s)
            self.zeros[e].copy_(z)


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


# --------------------------------------------------------------------------- #
# sizes / layout constants (bytes for int4 arenas, elements for bf16 arenas)
# --------------------------------------------------------------------------- #
D = 2304
C4 = 4096                    # kda_heads * kda_head_dim
H = 32                       # heads (KDA and MLA both)
MLA_Q_OUT = 6144             # 32 * (128 + 64)
MLA_KVA_OUT = 576            # 512 + 64
MLA_KVB_IN = 512
MLA_KVB_OUT = 8192           # 32 * (128 + 128)
MOE_M = 1024
N_EXP = 65                   # 64 routed + 1 shared (shared appended as expert 64)

KDA_LIN_B = (D // 2) * C4                    # packed bytes of one 2304x4096 linear
KDA_ATTN_B = 5 * KDA_LIN_B                   # q, k, v, g, o
KDA_LIN_SC = (D // GROUP_SIZE) * C4          # scale elements of one linear
KDA_ATTN_SC = 5 * KDA_LIN_SC

MLA_Q_B = (D // 2) * MLA_Q_OUT
MLA_Q_SC = (D // GROUP_SIZE) * MLA_Q_OUT
MLA_KVA_B = (D // 2) * MLA_KVA_OUT
MLA_KVA_SC = (D // GROUP_SIZE) * MLA_KVA_OUT
MLA_KVB_B = (MLA_KVB_IN // 2) * MLA_KVB_OUT
MLA_KVB_SC = (MLA_KVB_IN // GROUP_SIZE) * MLA_KVB_OUT
MLA_O_B = (C4 // 2) * D
MLA_O_SC = (C4 // GROUP_SIZE) * D
MLA_KVA_OFF = MLA_Q_B
MLA_KVB_OFF = MLA_KVA_OFF + MLA_KVA_B
MLA_O_OFF = MLA_KVB_OFF + MLA_KVB_B
MLA_Q_SC_OFF = 0
MLA_KVA_SC_OFF = MLA_Q_SC
MLA_KVB_SC_OFF = MLA_KVA_SC_OFF + MLA_KVA_SC
MLA_O_SC_OFF = MLA_KVB_SC_OFF + MLA_KVB_SC

EXP_B = (D // 2) * MOE_M                     # bytes per gate/up expert
DOWN_B = (MOE_M // 2) * D                    # bytes per down expert (same value)
MOE_SEC = N_EXP * EXP_B                      # one section (gate | up | down)
MOE_BLOCK_B = 3 * MOE_SEC
EXP_SC = (D // GROUP_SIZE) * MOE_M           # scale elems per gate/up expert
DOWN_SC = (MOE_M // GROUP_SIZE) * D          # scale elems per down expert (same)
MOE_SC_SEC = N_EXP * EXP_SC
MOE_SC_BLOCK = 3 * MOE_SC_SEC

KDA_SCALE = 0.08838834764831845              # 128 ** -0.5
MLA_SCALE = 0.07216878364870323              # 192 ** -0.5
ROUTED_SCALING = 2.446
ROPE_THETA = 10000.0

# bf16 workspace layout (elements)
WSB_QKV = 0          # KDA q/k/v/g (4x4096) or MLA q (6144) + kv_a (576)
WSB_O = 16384        # KDA o (4096) or MLA o_v (4096)
WSB_HE = 24576       # expert hidden, 9 x 1024
WSB_LEN = 33792

# fp32 workspace layout (elements)
WSF_XN = 0           # running hidden (residual stream), 2304
WSF_HB = 2304        # post-attention hidden h, 2304
WSF_BETA = 4608      # KDA beta logits, 32
WSF_LOGITS = 4640    # router logits, 64
WSF_WGT = 4704       # expert weights, 9
# MLA absorbed attention
WSF_U = 4800         # u[h,j], H*512
WSF_QROPE = WSF_U + H * 512
WSF_MPART = WSF_QROPE + H * 64
MLA_MAXCH = 512
WSF_LPART = WSF_MPART + MLA_MAXCH * H
WSF_MH = WSF_LPART + MLA_MAXCH * H
WSF_LH = WSF_MH + H
WSF_CBAR = WSF_LH + H
WSF_LEN = WSF_CBAR + H * 512

MLA_CL = 64          # MLA cache rows per attention-chunk task
NPHASE = 80          # barrier slots: b*16 + local

# expose the layout constants to Triton as constexpr globals (host-side
# names that torch needs stay plain python ints)
_HOST_ONLY = {"OP_TYPE", "HARDWARE_REQUIRED", "GROUP_SIZE", "WSB_LEN",
              "WSF_LEN", "NPHASE", "MLA_KVB_IN"}
for _name, _val in list(globals().items()):
    if _name.isupper() and _name not in _HOST_ONLY and isinstance(_val, (int, float)):
        globals()[_name] = tl.constexpr(_val)


# --------------------------------------------------------------------------- #
# Triton helpers (all inlined into the single megakernel)
# --------------------------------------------------------------------------- #
@triton.jit
def _rms_scale(xp):
    acc = 0.0
    for i in range(D // 128):
        v = tl.load(xp + i * 128 + tl.arange(0, 128)).to(tl.float32)
        acc += tl.sum(v * v)
    return 1.0 / tl.sqrt(acc / D + EPS)


@triton.jit
def _bar(bars, ph, target):
    tl.atomic_add(bars + ph, 1, sem="acq_rel")
    cur = tl.atomic_add(bars + ph, 0, sem="acquire")
    while cur < target:
        cur = tl.atomic_add(bars + ph, 0, sem="acquire")


@triton.jit
def _gemv_w4(wp, sp, zp, xp, nwp, rp, yp, nscale, oscale,
             K: tl.constexpr, BN: tl.constexpr, NSTRIDE: tl.constexpr,
             HASN: tl.constexpr, OUT_BF16: tl.constexpr,
             ADD_RESID: tl.constexpr, RESID_BF16: tl.constexpr,
             ATOMIC_OUT: tl.constexpr):
    """y[n0:n0+BN] = sum_k xn[k] * dequant(W)[k, n0:n0+BN];  xn = x (*norm*scale).

    wp/sp/zp already point at column n0 of the packed matrix / its scale rows.
    """
    on = tl.arange(0, BN)
    acc = tl.zeros((BN,), tl.float32)
    for g in range(K // 128):
        kb = g * 128
        ke = kb + 2 * tl.arange(0, 64)
        xe = tl.load(xp + ke).to(tl.float32)
        xo = tl.load(xp + ke + 1).to(tl.float32)
        if HASN:
            nwe = tl.load(nwp + ke).to(tl.float32)
            nwo = tl.load(nwp + ke + 1).to(tl.float32)
            xe = xe * nwe * nscale
            xo = xo * nwo * nscale
        pb = wp + (g * 64 + tl.arange(0, 64)[:, None]) * NSTRIDE + on[None, :]
        w8 = tl.load(pb)
        lo = (w8 & 0xF).to(tl.bfloat16)
        hi = ((w8 >> 4) & 0xF).to(tl.bfloat16)
        sv = tl.load(sp + g * NSTRIDE + on)
        zv = tl.load(zp + g * NSTRIDE + on)
        wlo = (lo - zv[None, :]) * sv[None, :]
        whi = (hi - zv[None, :]) * sv[None, :]
        acc += tl.sum(xe[:, None] * wlo.to(tl.float32), 0)
        acc += tl.sum(xo[:, None] * whi.to(tl.float32), 0)
    if ADD_RESID:
        if RESID_BF16:
            acc += tl.load(rp + on).to(tl.float32)
        else:
            acc += tl.load(rp + on)
    if ATOMIC_OUT:
        tl.atomic_add(yp + on, acc * oscale)
    elif OUT_BF16:
        tl.store(yp + on, acc.to(tl.bfloat16))
    else:
        tl.store(yp + on, acc)


@triton.jit
def _gemv2_w4(wp1, sp1, zp1, wp2, sp2, zp2, xp, nwp, hp, nscale,
              K: tl.constexpr, BN: tl.constexpr, NSTRIDE: tl.constexpr):
    """h = silu(x @ W1) * (x @ W2) for one output tile (MoE gate/up fused)."""
    on = tl.arange(0, BN)
    a1 = tl.zeros((BN,), tl.float32)
    a2 = tl.zeros((BN,), tl.float32)
    for g in range(K // 128):
        kb = g * 128
        ke = kb + 2 * tl.arange(0, 64)
        xe = tl.load(xp + ke).to(tl.float32)
        xo = tl.load(xp + ke + 1).to(tl.float32)
        nwe = tl.load(nwp + ke).to(tl.float32)
        nwo = tl.load(nwp + ke + 1).to(tl.float32)
        xe = xe * nwe * nscale
        xo = xo * nwo * nscale
        idx = (g * 64 + tl.arange(0, 64)[:, None]) * NSTRIDE + on[None, :]
        w8a = tl.load(wp1 + idx)
        s1 = tl.load(sp1 + g * NSTRIDE + on)
        z1 = tl.load(zp1 + g * NSTRIDE + on)
        wlo = ((w8a & 0xF).to(tl.bfloat16) - z1[None, :]) * s1[None, :]
        whi = (((w8a >> 4) & 0xF).to(tl.bfloat16) - z1[None, :]) * s1[None, :]
        a1 += tl.sum(xe[:, None] * wlo.to(tl.float32), 0)
        a1 += tl.sum(xo[:, None] * whi.to(tl.float32), 0)
        w8b = tl.load(wp2 + idx)
        s2 = tl.load(sp2 + g * NSTRIDE + on)
        z2 = tl.load(zp2 + g * NSTRIDE + on)
        wlo = ((w8b & 0xF).to(tl.bfloat16) - z2[None, :]) * s2[None, :]
        whi = (((w8b >> 4) & 0xF).to(tl.bfloat16) - z2[None, :]) * s2[None, :]
        a2 += tl.sum(xe[:, None] * wlo.to(tl.float32), 0)
        a2 += tl.sum(xo[:, None] * whi.to(tl.float32), 0)
    hv = a1 * tl.sigmoid(a1) * a2
    tl.store(hp + on, hv.to(tl.bfloat16))


@triton.jit
def _gemv_bf16(wp, xp, nwp, yp, nscale,
               OUT: tl.constexpr, NSTRIDE: tl.constexpr, HASN: tl.constexpr):
    """Dense bf16 GEMV for the tiny weights: y[n] = sum_k W[n,k] * xn[k]."""
    on = tl.arange(0, OUT)
    acc = tl.zeros((OUT,), tl.float32)
    for k0 in range(0, D, 128):
        kk = k0 + tl.arange(0, 128)
        xk = tl.load(xp + kk).to(tl.float32)
        if HASN:
            nw = tl.load(nwp + kk).to(tl.float32)
            xk = xk * nw * nscale
        wt = tl.load(wp + on[:, None] * NSTRIDE + kk[None, :]).to(tl.float32)
        acc += tl.sum(wt * xk[None, :], 1)
    tl.store(yp + on, acc)


@triton.jit
def _mla_u(wp, sp, zp, qp, up, hh, j0):
    """u[hh, j0:j0+64] = sum_d q_nope[hh,d] * W_b[j, hh*256+d].

    Output index j is the packed input dimension of W_b, so 64 output rows are
    32 packed rows (even row = low nibble, odd row = high nibble).  The column
    window is head hh's nope block [hh*256, hh*256+128).
    """
    col0 = hh * 256
    jr = tl.arange(0, 32)
    cc = tl.arange(0, 128)
    w8 = tl.load(wp + (j0 // 2 + jr)[:, None] * MLA_KVB_OUT + col0 + cc[None, :])
    sv = tl.load(sp + (j0 // 128) * MLA_KVB_OUT + col0 + cc)
    zv = tl.load(zp + (j0 // 128) * MLA_KVB_OUT + col0 + cc)
    wlo = ((w8 & 0xF).to(tl.bfloat16) - zv[None, :]) * sv[None, :]
    whi = (((w8 >> 4) & 0xF).to(tl.bfloat16) - zv[None, :]) * sv[None, :]
    q = tl.load(qp + hh * 192 + cc).to(tl.float32)
    ue = tl.sum(wlo.to(tl.float32) * q[None, :], 1)
    uo = tl.sum(whi.to(tl.float32) * q[None, :], 1)
    tl.store(up + hh * 512 + j0 + 2 * tl.arange(0, 32), ue)
    tl.store(up + hh * 512 + j0 + 2 * tl.arange(0, 32) + 1, uo)


# --------------------------------------------------------------------------- #
# the megakernel
# --------------------------------------------------------------------------- #
@triton.jit(do_not_specialize=["L", "do_copy", "step", "stop_at"])
def _kimi_mega(
    hidden, out,
    attn_wq, attn_sc, attn_z,
    moe_wq, moe_sc, moe_z,
    beta_w, router_w, conv_w, norm_w,
    S0, S1, S2,
    cq0, ck0, cv0, cq1, ck1, cv1, cq2, ck2, cv2,
    ckv_src, kro_src, ckv_buf, kro_buf,
    wsf, wsb, wsi, bars, dbg,
    L, do_copy, step, stop_at,
    G: tl.constexpr,
):
    pid = tl.program_id(0)
    target = (step + 1) * G

    b = 0
    stopped = 0
    while b < 4 and stopped == 0:
        is_mla = b == 3
        abase = b * KDA_ATTN_B          # byte offset of this block's attn int4 (MLA last)
        asc = b * KDA_ATTN_SC           # element offset of this block's attn scales/zeros
        nbase = b * 4608                # norm arena: [attn_norm, moe_norm] per block
        mbase = b * MOE_BLOCK_B         # byte offset of this block's MoE int4
        msc = b * MOE_SC_BLOCK          # element offset of this block's MoE scales/zeros

        # state pointers for this block (S / conv windows; MLA leaves them unused)
        if b == 0:
            Sp = S0
            cqp = cq0
            ckp = ck0
            cvp = cv0
        elif b == 1:
            Sp = S1
            cqp = cq1
            ckp = ck1
            cvp = cv1
        else:
            Sp = S2
            cqp = cq2
            ckp = ck2
            cvp = cv2

        # rmsnorm scale of the block input x (block 0: bf16 hidden; else fp32 x_next)
        if b == 0:
            nscale = _rms_scale(hidden)
        else:
            nscale = _rms_scale(wsf + WSF_XN)

        # ---------------- phase A: projections on rmsnorm(x) ---------------- #
        t = pid
        if is_mla:
            nta = 57
        else:
            nta = 129
            if b == 0 and do_copy == 1:
                nta += (L + 1023) // 1024
        while t < nta:
            if is_mla:
                # MLA is never block 0, so its input is the fp32 residual stream
                if t < 48:
                    n0 = t * 128
                    _gemv_w4(attn_wq + abase + n0, attn_sc + asc + n0,
                             attn_z + asc + n0, wsf + WSF_XN, norm_w + nbase,
                             wsf + WSF_XN, wsb + WSB_QKV + n0,
                             nscale, 0.0, K=D, BN=128, NSTRIDE=MLA_Q_OUT,
                             HASN=True, OUT_BF16=True, ADD_RESID=False,
                             RESID_BF16=False, ATOMIC_OUT=False)
                else:
                    n0 = (t - 48) * 64
                    _gemv_w4(attn_wq + abase + MLA_KVA_OFF + n0,
                             attn_sc + asc + MLA_KVA_SC_OFF + n0,
                             attn_z + asc + MLA_KVA_SC_OFF + n0,
                             wsf + WSF_XN, norm_w + nbase, wsf + WSF_XN,
                             wsb + WSB_QKV + 6144 + n0,
                             nscale, 0.0, K=D, BN=64, NSTRIDE=MLA_KVA_OUT,
                             HASN=True, OUT_BF16=True, ADD_RESID=False,
                             RESID_BF16=False, ATOMIC_OUT=False)
            else:
                if t < 128:
                    mat = t // 32
                    n0 = (t % 32) * 128
                    if b == 0:
                        _gemv_w4(attn_wq + abase + mat * KDA_LIN_B + n0,
                                 attn_sc + asc + mat * KDA_LIN_SC + n0,
                                 attn_z + asc + mat * KDA_LIN_SC + n0,
                                 hidden, norm_w + nbase, wsf + WSF_XN,
                                 wsb + WSB_QKV + mat * 4096 + n0,
                                 nscale, 0.0, K=D, BN=128, NSTRIDE=C4,
                                 HASN=True, OUT_BF16=True, ADD_RESID=False,
                                 RESID_BF16=False, ATOMIC_OUT=False)
                    else:
                        _gemv_w4(attn_wq + abase + mat * KDA_LIN_B + n0,
                                 attn_sc + asc + mat * KDA_LIN_SC + n0,
                                 attn_z + asc + mat * KDA_LIN_SC + n0,
                                 wsf + WSF_XN, norm_w + nbase, wsf + WSF_XN,
                                 wsb + WSB_QKV + mat * 4096 + n0,
                                 nscale, 0.0, K=D, BN=128, NSTRIDE=C4,
                                 HASN=True, OUT_BF16=True, ADD_RESID=False,
                                 RESID_BF16=False, ATOMIC_OUT=False)
                elif t == 128:
                    # beta logits (bf16 dense weight, row-major (32, 2304))
                    if b == 0:
                        _gemv_bf16(beta_w + b * 73728, hidden, norm_w + nbase,
                                   wsf + WSF_BETA, nscale, OUT=32,
                                   NSTRIDE=D, HASN=True)
                    else:
                        _gemv_bf16(beta_w + b * 73728, wsf + WSF_XN,
                                   norm_w + nbase, wsf + WSF_BETA, nscale,
                                   OUT=32, NSTRIDE=D, HASN=True)
                else:
                    # one-time cache copy into the capacity buffer (block 0 only)
                    ct = t - 129
                    r0 = ct * 1024
                    for i in range(0, 1024 * 512, 8192):
                        off = r0 * 512 + i + tl.arange(0, 8192)
                        mm = off < L * 512
                        vv = tl.load(ckv_src + off, mask=mm)
                        tl.store(ckv_buf + off, vv, mask=mm)
                    for i in range(0, 1024 * 64, 4096):
                        off = r0 * 64 + i + tl.arange(0, 4096)
                        mm = off < L * 64
                        vv = tl.load(kro_src + off, mask=mm)
                        tl.store(kro_buf + off, vv, mask=mm)
            t += G
        _bar(bars, b * 16 + 0, target)
        if stop_at == 100 + b:
            v = tl.load(wsb + WSB_QKV + tl.arange(0, 4096)).to(tl.float32)
            tl.store(dbg + tl.arange(0, 4096), v)

        # ---------------- phase B ---------------- #
        t = pid
        if not is_mla:
            # KDA: per-head conv + gated-delta recurrence
            while t < 32:
                h = t
                c0 = h * 128
                col = tl.arange(0, 128)
                qc = tl.zeros((128,), tl.float32)
                kc = tl.zeros((128,), tl.float32)
                vc = tl.zeros((128,), tl.float32)
                for idx in tl.static_range(3):
                    val = tl.load(wsb + WSB_QKV + idx * 4096 + c0 + col).to(tl.float32)
                    if idx == 0:
                        csp = cqp + c0 + col
                    elif idx == 1:
                        csp = ckp + c0 + col
                    else:
                        csp = cvp + c0 + col
                    p0 = tl.load(csp + 0 * 4096).to(tl.float32)
                    p1 = tl.load(csp + 1 * 4096).to(tl.float32)
                    p2 = tl.load(csp + 2 * 4096).to(tl.float32)
                    cwb = conv_w + (b * 3 + idx) * (4096 * 4) + (c0 + col)[:, None] * 4
                    cw = tl.load(cwb + tl.arange(0, 4)[None, :]).to(tl.float32)
                    # split of (128,2,2) yields ([w0,w2], [w1,w3]) per channel
                    cwa, cwb2 = tl.split(tl.reshape(cw, (128, 2, 2)))
                    cw0, cw2 = tl.split(cwa)
                    cw1, cw3 = tl.split(cwb2)
                    ov = p0 * cw0 + p1 * cw1 + p2 * cw2 + val * cw3
                    sv = ov * tl.sigmoid(ov)
                    tl.store(csp + 0 * 4096, p1.to(tl.bfloat16))
                    tl.store(csp + 1 * 4096, p2.to(tl.bfloat16))
                    tl.store(csp + 2 * 4096, val.to(tl.bfloat16))
                    if idx == 0:
                        qc = sv
                    elif idx == 1:
                        kc = sv
                    else:
                        vc = sv
                gr = tl.load(wsb + WSB_QKV + 3 * 4096 + c0 + col).to(tl.float32)
                spg = tl.where(gr > 20.0, gr, tl.log(1.0 + tl.exp(gr)))
                dec = tl.exp(-spg)
                be = tl.sigmoid(tl.load(wsf + WSF_BETA + h))
                qh = qc * KDA_SCALE
                soff = Sp + h * 16384 + col[:, None] * 128 + col[None, :]
                S = tl.load(soff)
                Sd = S * dec[:, None]
                pred = tl.sum(Sd * kc[:, None], 0)
                pq = tl.sum(Sd * qh[:, None], 0)
                kq = tl.sum(kc * qh, 0)
                dv = vc - pred
                oh = pq + be * kq * dv
                Sn = Sd + be * kc[:, None] * dv[None, :]
                tl.store(soff, Sn)
                tl.store(wsb + WSB_O + c0 + col, oh.to(tl.bfloat16))
                t += G
        else:
            # MLA: absorbed u[h], roped q_rope[h], append new token to caches
            ntu = H * 8 + 1        # u tasks (h, j-tile of 64) + 1 append task
            while t < ntu:
                if t < H * 8:
                    hh = t // 8
                    j0 = (t % 8) * 64
                    _mla_u(attn_wq + abase + MLA_KVB_OFF,
                           attn_sc + asc + MLA_KVB_SC_OFF,
                           attn_z + asc + MLA_KVB_SC_OFF,
                           wsb + WSB_QKV, wsf + WSF_U, hh, j0)
                else:
                    qpb = wsb + WSB_QKV
                    ii = tl.arange(0, 32)
                    inv = tl.exp(-(ii.to(tl.float32) * 2.0 / 64) * tl.log(ROPE_THETA))
                    ang = L.to(tl.float32) * inv
                    cosp = tl.cos(ang)
                    sinp = tl.sin(ang)
                    # roped q_rope[h] for every head
                    qrt = tl.load(qpb + tl.arange(0, 32)[:, None] * 192 + 128 +
                                  ii[None, :] * 2).to(tl.float32)
                    qro = tl.load(qpb + tl.arange(0, 32)[:, None] * 192 + 128 +
                                  ii[None, :] * 2 + 1).to(tl.float32)
                    qre = qrt * cosp[None, :] - qro * sinp[None, :]
                    qro2 = qro * cosp[None, :] + qrt * sinp[None, :]
                    tl.store(wsf + WSF_QROPE + tl.arange(0, 32)[:, None] * 64 +
                             ii[None, :] * 2, qre)
                    tl.store(wsf + WSF_QROPE + tl.arange(0, 32)[:, None] * 64 +
                             ii[None, :] * 2 + 1, qro2)
                    # rope the new k_rope and append both rows to the caches
                    krt = tl.load(qpb + 6144 + 512 + ii * 2).to(tl.float32)
                    kro = tl.load(qpb + 6144 + 512 + ii * 2 + 1).to(tl.float32)
                    kre = krt * cosp - kro * sinp
                    kro2 = kro * cosp + krt * sinp
                    tl.store(kro_buf + L * 64 + ii * 2, kre.to(tl.bfloat16))
                    tl.store(kro_buf + L * 64 + ii * 2 + 1, kro2.to(tl.bfloat16))
                    cvv = tl.load(qpb + 6144 + tl.arange(0, 512))
                    tl.store(ckv_buf + L * 512 + tl.arange(0, 512), cvv)
                t += G
        _bar(bars, b * 16 + 1, target)
        if stop_at == 200 + b:
            v = tl.load(wsb + WSB_O + tl.arange(0, 4096)).to(tl.float32)
            tl.store(dbg + tl.arange(0, 4096), v)
        if stop_at == 300 + b:
            v = tl.load(wsf + WSF_U + tl.arange(0, 16384))
            tl.store(dbg + tl.arange(0, 16384), v)

        # ---------------- phase C ---------------- #
        if not is_mla:
            # KDA: o_proj (4096 -> 2304) + residual -> h
            t = pid
            while t < 36:
                n0 = t * 64
                if b == 0:
                    _gemv_w4(attn_wq + abase + 4 * KDA_LIN_B + n0,
                             attn_sc + asc + 4 * KDA_LIN_SC + n0,
                             attn_z + asc + 4 * KDA_LIN_SC + n0,
                             wsb + WSB_O, norm_w + nbase, hidden + n0,
                             wsf + WSF_HB + n0, 0.0, 0.0,
                             K=C4, BN=64, NSTRIDE=D,
                             HASN=False, OUT_BF16=False, ADD_RESID=True,
                             RESID_BF16=True, ATOMIC_OUT=False)
                else:
                    _gemv_w4(attn_wq + abase + 4 * KDA_LIN_B + n0,
                             attn_sc + asc + 4 * KDA_LIN_SC + n0,
                             attn_z + asc + 4 * KDA_LIN_SC + n0,
                             wsb + WSB_O, norm_w + nbase, wsf + WSF_XN + n0,
                             wsf + WSF_HB + n0, 0.0, 0.0,
                             K=C4, BN=64, NSTRIDE=D,
                             HASN=False, OUT_BF16=False, ADD_RESID=True,
                             RESID_BF16=False, ATOMIC_OUT=False)
                t += G
            _bar(bars, b * 16 + 2, target)
            if stop_at == 2 * b:
                v = tl.load(wsf + WSF_HB + tl.arange(0, 4096),
                            mask=tl.arange(0, 4096) < D, other=0.0)
                tl.store(out + tl.arange(0, 4096), v.to(tl.bfloat16),
                         mask=tl.arange(0, 4096) < D)
                stopped = 1
        else:
            # MLA: scoring pass over the latent cache (all heads per chunk)
            nch = (L + 1 + MLA_CL - 1) // MLA_CL
            t = pid
            while t < nch:
                l0 = t * MLA_CL
                rows = l0 + tl.arange(0, MLA_CL)
                rm = rows < L + 1
                # u (H,512) and roped q_rope (H,64), kept live for the dots
                u = tl.load(wsf + WSF_U + tl.arange(0, H)[:, None] * 512 +
                            tl.arange(0, 512)[None, :])
                qr = tl.load(wsf + WSF_QROPE + tl.arange(0, H)[:, None] * 64 +
                             tl.arange(0, 64)[None, :])
                ckk = tl.load(ckv_buf + rows[:, None] * 512 + tl.arange(0, 512)[None, :],
                              mask=rm[:, None], other=0.0).to(tl.float32)
                krr = tl.load(kro_buf + rows[:, None] * 64 + tl.arange(0, 64)[None, :],
                              mask=rm[:, None], other=0.0).to(tl.float32)
                sc_nope = tl.dot(ckk, tl.trans(u), input_precision="tf32")
                sc_rope = tl.dot(krr, tl.trans(qr), input_precision="tf32")
                scv = (sc_nope + sc_rope) * MLA_SCALE
                scv = tl.where(rm[:, None], scv, -1.0e30)
                mch = tl.max(scv, 0)                       # (H,)
                pch = tl.exp(scv - mch[None, :])           # (CL, H)
                lch = tl.sum(pch, 0)                       # (H,)
                tl.store(wsf + WSF_MPART + t * H + tl.arange(0, H), mch)
                tl.store(wsf + WSF_LPART + t * H + tl.arange(0, H), lch)
                t += G
            _bar(bars, b * 16 + 2, target)

            # reduce per-head softmax stats across chunks; zero c_bar
            t = pid
            while t < H:
                cid = tl.arange(0, MLA_MAXCH)
                cm = cid < nch
                mvec = tl.load(wsf + WSF_MPART + cid * H + t, mask=cm, other=-1.0e30)
                lvec = tl.load(wsf + WSF_LPART + cid * H + t, mask=cm, other=0.0)
                mstar = tl.max(mvec, 0)
                f = tl.exp(mvec - mstar)
                ltot = tl.sum(lvec * f, 0)
                tl.store(wsf + WSF_MH + t, mstar)
                tl.store(wsf + WSF_LH + t, ltot)
                tl.store(wsf + WSF_CBAR + t * 512 + tl.arange(0, 512),
                         tl.zeros((512,), tl.float32))
                t += G
            _bar(bars, b * 16 + 3, target)

            # c_bar pass: p = softmax(score); c_bar_h += sum_l p[l,h] * c_l
            t = pid
            while t < nch:
                l0 = t * MLA_CL
                rows = l0 + tl.arange(0, MLA_CL)
                rm = rows < L + 1
                u = tl.load(wsf + WSF_U + tl.arange(0, H)[:, None] * 512 +
                            tl.arange(0, 512)[None, :])
                qr = tl.load(wsf + WSF_QROPE + tl.arange(0, H)[:, None] * 64 +
                             tl.arange(0, 64)[None, :])
                mh = tl.load(wsf + WSF_MH + tl.arange(0, H))
                lh = tl.load(wsf + WSF_LH + tl.arange(0, H))
                ckk = tl.load(ckv_buf + rows[:, None] * 512 + tl.arange(0, 512)[None, :],
                              mask=rm[:, None], other=0.0).to(tl.float32)
                krr = tl.load(kro_buf + rows[:, None] * 64 + tl.arange(0, 64)[None, :],
                              mask=rm[:, None], other=0.0).to(tl.float32)
                sc_nope = tl.dot(ckk, tl.trans(u), input_precision="tf32")
                sc_rope = tl.dot(krr, tl.trans(qr), input_precision="tf32")
                scv = (sc_nope + sc_rope) * MLA_SCALE
                scv = tl.where(rm[:, None], scv, -1.0e30)
                p = tl.exp(scv - mh[None, :]) / lh[None, :]       # (CL, H)
                cb = tl.dot(tl.trans(p), ckk, input_precision="tf32")  # (H, 512)
                tl.atomic_add(wsf + WSF_CBAR + tl.arange(0, H)[:, None] * 512 +
                              tl.arange(0, 512)[None, :], cb)
                t += G
            _bar(bars, b * 16 + 4, target)
            if stop_at == 400 + b:
                v = tl.load(wsf + WSF_CBAR + tl.arange(0, 16384))
                tl.store(dbg + tl.arange(0, 16384), v)

            # o = c_bar @ W_b_v (per head), grouped dequant-GEMV -> o_v (4096)
            t = pid
            while t < H * 2:
                hh = t // 2
                dv0 = (t % 2) * 64
                # o[hh, dv] = sum_j cbar[hh,j] * W_b[j, hh*256 + 128 + dv]
                col0 = hh * 256 + 128 + dv0
                acc = tl.zeros((64,), tl.float32)
                on = tl.arange(0, 64)
                for g in range(0, 512, 128):
                    cb = tl.load(wsf + WSF_CBAR + hh * 512 + g + 2 * tl.arange(0, 64))
                    cb2 = tl.load(wsf + WSF_CBAR + hh * 512 + g + 2 * tl.arange(0, 64) + 1)
                    w8 = tl.load(attn_wq + abase + MLA_KVB_OFF +
                                 (g // 2 + tl.arange(0, 64)[:, None]) * MLA_KVB_OUT +
                                 col0 + on[None, :])
                    sv = tl.load(attn_sc + asc + MLA_KVB_SC_OFF +
                                 (g // 128) * MLA_KVB_OUT + col0 + on)
                    zv = tl.load(attn_z + asc + MLA_KVB_SC_OFF +
                                 (g // 128) * MLA_KVB_OUT + col0 + on)
                    wlo = ((w8 & 0xF).to(tl.bfloat16) - zv[None, :]) * sv[None, :]
                    whi = (((w8 >> 4) & 0xF).to(tl.bfloat16) - zv[None, :]) * sv[None, :]
                    acc += tl.sum(wlo.to(tl.float32) * cb[:, None], 0)
                    acc += tl.sum(whi.to(tl.float32) * cb2[:, None], 0)
                tl.store(wsb + WSB_O + hh * 128 + dv0 + on, acc.to(tl.bfloat16))
                t += G
            _bar(bars, b * 16 + 5, target)

            # o_proj (4096 -> 2304) + residual -> h
            t = pid
            while t < 36:
                n0 = t * 64
                _gemv_w4(attn_wq + abase + MLA_O_OFF + n0,
                         attn_sc + asc + MLA_O_SC_OFF + n0,
                         attn_z + asc + MLA_O_SC_OFF + n0,
                         wsb + WSB_O, norm_w + nbase, wsf + WSF_XN + n0,
                         wsf + WSF_HB + n0, 0.0, 0.0,
                         K=C4, BN=64, NSTRIDE=D,
                         HASN=False, OUT_BF16=False, ADD_RESID=True,
                         RESID_BF16=False, ATOMIC_OUT=False)
                t += G
            _bar(bars, b * 16 + 6, target)
            if stop_at == 2 * b:
                v = tl.load(wsf + WSF_HB + tl.arange(0, 4096),
                            mask=tl.arange(0, 4096) < D, other=0.0)
                tl.store(out + tl.arange(0, 4096), v.to(tl.bfloat16),
                         mask=tl.arange(0, 4096) < D)
                stopped = 1

        # ---------------- MoE (both kinds) ---------------- #
        moe_base_phase = 7 if is_mla else 3
        if stopped == 0:
            ns2 = _rms_scale(wsf + WSF_HB)

            # router (bf16 dense (64, 2304)), 4 tasks of 16 outputs
            t = pid
            while t < 4:
                _gemv_bf16(router_w + b * 147456 + t * 16 * D, wsf + WSF_HB,
                           norm_w + nbase + 2304, wsf + WSF_LOGITS + t * 16,
                           ns2, OUT=16, NSTRIDE=D, HASN=True)
                t += G
            _bar(bars, b * 16 + moe_base_phase, target)

            # softmax + top-8 + renormalize; seed x_next with h
            if pid == 0:
                lg = tl.load(wsf + WSF_LOGITS + tl.arange(0, 64))
                mx = tl.max(lg, 0)
                pp = tl.exp(lg - mx)
                pp = pp / tl.sum(pp, 0)
                for j in tl.static_range(8):
                    m = tl.max(pp, 0)
                    i = tl.argmax(pp, 0)
                    tl.store(wsi + j, i)
                    tl.store(wsf + WSF_WGT + j, m)
                    pp = tl.where(tl.arange(0, 64) == i, -1.0, pp)
                wsum = 0.0
                for j in tl.static_range(8):
                    wsum += tl.load(wsf + WSF_WGT + j)
                for j in tl.static_range(8):
                    wj = tl.load(wsf + WSF_WGT + j) / (wsum + 1e-9) * ROUTED_SCALING
                    tl.store(wsf + WSF_WGT + j, wj)
                tl.store(wsi + 8, 64)                 # shared expert lives at slot 64
                tl.store(wsf + WSF_WGT + 8, 1.0)
                hv = tl.load(wsf + WSF_HB + tl.arange(0, 4096),
                             mask=tl.arange(0, 4096) < D, other=0.0)
                tl.store(wsf + WSF_XN + tl.arange(0, 4096), hv,
                         mask=tl.arange(0, 4096) < D)
            _bar(bars, b * 16 + moe_base_phase + 1, target)

            # expert gate/up: 9 experts x 16 tiles of 64
            t = pid
            while t < 144:
                slot = t // 16
                n0 = (t % 16) * 64
                e = tl.load(wsi + slot).to(tl.int64)
                gb = moe_wq + mbase + e * EXP_B + n0
                gsc = moe_sc + msc + e * EXP_SC + n0
                gzp = moe_z + msc + e * EXP_SC + n0
                ub = moe_wq + mbase + MOE_SEC + e * EXP_B + n0
                usc = moe_sc + msc + MOE_SC_SEC + e * EXP_SC + n0
                uzp = moe_z + msc + MOE_SC_SEC + e * EXP_SC + n0
                _gemv2_w4(gb, gsc, gzp, ub, usc, uzp, wsf + WSF_HB,
                          norm_w + nbase + 2304, wsb + WSB_HE + slot * 1024 + n0,
                          ns2, K=D, BN=64, NSTRIDE=MOE_M)
                t += G
            _bar(bars, b * 16 + moe_base_phase + 2, target)

            # expert down: 9 experts x 18 tiles of 128, weighted atomic accumulate
            t = pid
            while t < 162:
                slot = t // 18
                n0 = (t % 18) * 128
                e = tl.load(wsi + slot).to(tl.int64)
                we = tl.load(wsf + WSF_WGT + slot)
                db = moe_wq + mbase + 2 * MOE_SEC + e * DOWN_B + n0
                dsc = moe_sc + msc + 2 * MOE_SC_SEC + e * DOWN_SC + n0
                dzp = moe_z + msc + 2 * MOE_SC_SEC + e * DOWN_SC + n0
                _gemv_w4(db, dsc, dzp, wsb + WSB_HE + slot * 1024,
                         norm_w + nbase + 2304, wsf + WSF_XN,
                         wsf + WSF_XN + n0, 0.0, we,
                         K=MOE_M, BN=128, NSTRIDE=D,
                         HASN=False, OUT_BF16=False, ADD_RESID=False,
                         RESID_BF16=False, ATOMIC_OUT=True)
                t += G
            _bar(bars, b * 16 + moe_base_phase + 3, target)
            if stop_at == 2 * b + 1:
                v = tl.load(wsf + WSF_XN + tl.arange(0, 4096),
                            mask=tl.arange(0, 4096) < D, other=0.0)
                tl.store(out + tl.arange(0, 4096), v.to(tl.bfloat16),
                         mask=tl.arange(0, 4096) < D)
                stopped = 1
        b += 1

    # final: fp32 residual stream -> bf16 output hidden
    if stopped == 0 and pid == 0:
        v = tl.load(wsf + WSF_XN + tl.arange(0, 4096),
                    mask=tl.arange(0, 4096) < D, other=0.0)
        tl.store(out + tl.arange(0, 4096), v.to(tl.bfloat16),
                 mask=tl.arange(0, 4096) < D)


# --------------------------------------------------------------------------- #
# model
# --------------------------------------------------------------------------- #
class Model(nn.Module):
    def __init__(self, cfg: Config):
        super().__init__()
        self.cfg = cfg
        self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern)
        self.reset_parameters()
        self._ready = False
        self._state_cache: dict = {}
        self._step_i = 0
        self._stop_at = -1          # debug hook: stop after block b (2b attn / 2b+1 moe)
        self._grid = 0
        self._kda_idxs = [i for i, k in enumerate(cfg.pattern) if k == "K"]
        self._mla_idx = cfg.pattern.index("M")

    def reset_parameters(self):
        g = torch.Generator(device="cpu").manual_seed(1234)
        for mod in self.modules():
            if isinstance(mod, (QuantLinear, QuantExperts)):
                mod.init_random(g)
            elif isinstance(mod, nn.Linear):
                nn.init.normal_(mod.weight, 0.0, 0.02, generator=g)
            elif isinstance(mod, KDA):
                nn.init.normal_(mod.conv_w, 0.0, 0.1, generator=g)

    # -- arena construction (once, on the first step after weights are loaded) --
    def _build(self, device):
        bl = self.blocks
        attn_wq, attn_sc, attn_z = [], [], []
        for i, kind in enumerate(self.cfg.pattern):
            attn = bl[i].attn
            if kind == "K":
                mats = [attn.q_proj, attn.k_proj, attn.v_proj, attn.g_proj, attn.o_proj]
            else:
                mats = [attn.q_proj, attn.kv_a, attn.kv_b, attn.o_proj]
            for m in mats:
                attn_wq.append(m.w_q.flatten())
                attn_sc.append(m.scales.flatten())
                attn_z.append(m.zeros.flatten())
        self._attn_wq = torch.cat(attn_wq).contiguous()
        self._attn_sc = torch.cat(attn_sc).contiguous()
        self._attn_z = torch.cat(attn_z).contiguous()

        moe_wq, moe_sc, moe_z = [], [], []
        for i in range(len(self.cfg.pattern)):
            moe = bl[i].moe
            for routed, shared in ((moe.gate, moe.s_gate), (moe.up, moe.s_up),
                                   (moe.down, moe.s_down)):
                moe_wq.append(routed.w_q.flatten())
                moe_wq.append(shared.w_q.flatten())
                moe_sc.append(routed.scales.flatten())
                moe_sc.append(shared.scales.flatten())
                moe_z.append(routed.zeros.flatten())
                moe_z.append(shared.zeros.flatten())
        self._moe_wq = torch.cat(moe_wq).contiguous()
        self._moe_sc = torch.cat(moe_sc).contiguous()
        self._moe_z = torch.cat(moe_z).contiguous()

        self._beta_w = torch.cat(
            [bl[i].attn.beta_proj.weight.flatten() for i in self._kda_idxs]
        ).contiguous()
        self._router_w = torch.cat(
            [bl[i].moe.router.weight.flatten() for i in range(len(bl))]
        ).contiguous()
        self._conv_w = torch.cat(
            [bl[i].attn.conv_w.flatten() for i in self._kda_idxs]
        ).contiguous()
        norms = []
        for i in range(len(bl)):
            norms.append(bl[i].attn_norm)
            norms.append(bl[i].moe_norm)
        self._norm_w = torch.cat([n.flatten() for n in norms]).contiguous()

        self._wsf = torch.zeros(WSF_LEN, dtype=torch.float32, device=device)
        self._wsb = torch.zeros(WSB_LEN, dtype=torch.bfloat16, device=device)
        self._wsi = torch.zeros(16, dtype=torch.int32, device=device)
        self._bars = torch.zeros(NPHASE, dtype=torch.int32, device=device)
        self._dbg_buf = torch.zeros(16384, dtype=torch.float32, device=device)
        self._grid = torch.cuda.get_device_properties(device).multi_processor_count

    @torch.no_grad()
    def step(self, hidden, state):
        device = hidden.device
        if not self._ready:
            self._build(device)
            self._ready = True

        mla = self._mla_idx
        st_m = state[mla]
        L = st_m["c_kv"].shape[0]
        key = id(state)
        sc = self._state_cache.get(key)
        if sc is None or sc["cap"] <= L:
            cap = L + 4096
            ckv = torch.empty(cap, MLA_KVB_IN, dtype=torch.bfloat16, device=device)
            kro = torch.empty(cap, 64, dtype=torch.bfloat16, device=device)
            src_ckv, src_kro = st_m["c_kv"], st_m["k_rope"]
            do_copy = 1
            sc = {"cap": cap, "ckv": ckv, "kro": kro, "len": L}
            self._state_cache[key] = sc
        else:
            ckv, kro = sc["ckv"], sc["kro"]
            src_ckv, src_kro = ckv, kro
            do_copy = 0

        out = torch.empty_like(hidden)
        kda = self._kda_idxs
        _kimi_mega[(self._grid,)](
            hidden, out,
            self._attn_wq, self._attn_sc, self._attn_z,
            self._moe_wq, self._moe_sc, self._moe_z,
            self._beta_w, self._router_w, self._conv_w, self._norm_w,
            state[kda[0]]["S"], state[kda[1]]["S"], state[kda[2]]["S"],
            state[kda[0]]["cq"], state[kda[0]]["ck"], state[kda[0]]["cv"],
            state[kda[1]]["cq"], state[kda[1]]["ck"], state[kda[1]]["cv"],
            state[kda[2]]["cq"], state[kda[2]]["ck"], state[kda[2]]["cv"],
            src_ckv, src_kro, ckv, kro,
            self._wsf, self._wsb, self._wsi, self._bars, self._dbg_buf,
            L, do_copy, self._step_i, self._stop_at,
            G=self._grid, num_warps=8,
        )
        sc["len"] = L + 1
        st_m["c_kv"] = ckv[: L + 1]
        st_m["k_rope"] = kro[: L + 1]
        self._step_i += 1
        return out, state


# --------------------------------------------------------------------------- #
# state / input builders (same contract as reference.py)
# --------------------------------------------------------------------------- #
def init_state(cfg: Config, context_len: int, seed: int) -> list:
    dev = torch.device("cuda:0")
    g = torch.Generator(device=dev).manual_seed(seed)
    Hh, Dk = cfg.kda_heads, cfg.kda_head_dim
    C = Hh * Dk
    state = []
    for kind in cfg.pattern:
        if kind == "K":
            state.append({
                "S": torch.randn(Hh, Dk, Dk, device=dev, generator=g) * 0.05,
                "cq": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1,
                "ck": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1,
                "cv": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1,
            })
        else:
            state.append({
                "c_kv": torch.randn(context_len, cfg.kv_lora, device=dev, generator=g, dtype=cfg.dtype) * 0.1,
                "k_rope": torch.randn(context_len, cfg.qk_rope, device=dev, generator=g, dtype=cfg.dtype) * 0.1,
            })
    return state


def init_token(cfg: Config, seed: int) -> torch.Tensor:
    dev = torch.device("cuda:0")
    g = torch.Generator(device=dev).manual_seed(seed + 1)
    return torch.randn(cfg.hidden, device=dev, generator=g, dtype=cfg.dtype) * 0.25


if __name__ == "__main__":
    cfg = build_config({"n_experts": 64})
    m = Model(cfg).cuda().eval()
    st = init_state(cfg, context_len=2048, seed=0)
    h = init_token(cfg, seed=0)
    for _ in range(4):
        h, st = m.step(h, st)
    torch.cuda.synchronize()
    print(f"ok: out {tuple(h.shape)} finite {torch.isfinite(h).all().item()} | "
          f"MLA cache {st[3]['c_kv'].shape[0]}")

20260803_194401_or-fable_qwen_qwen3.8-max_02_kimi_linear_decode