KernelBench mega · RTX PRO 6000

Kimi-Linear Decode Kimi K3 (256k)

cleandid not score

manually audited: clean

RTX PRO 6000 kinetic-0715 cell (14.11x, geomean over ctx 2048/8192/16384). Genuine single-launch persistent Triton megakernel: one @triton.jit `mega` kernel invoked exactly once per step() (NBLK=376 persistent blocks, spin grid-wide barriers between statically scheduled stages). No torch.cuda. CUDAGraph, no torch.compile, no per-op kernel loop, no getattr/importlib obfuscation, no sidecar modules. Computes the full real op in-kernel: 3x KDA (fused int4 dequant-GEMV q/k/v/g via int8 tensor cores with exact per-group fp32 dequant correction, short-conv + silu with conv-window state update, gated-delta recurrence matching reference S update, o_proj + residual), MLA in absorbed form (q_abs against the 576-d latent cache, rope on q/k, cache append, two-pass chunked softmax + PV, absorbed Wv, o_proj + residual), and per-layer MoE with router softmax, iterative top-8 (tie-hardened), weight renorm * routed_scaling, all 8 routed experts PLUS the shared expert (9-slot task loops: 2592=9x288 gate/up, 1296=9x144 down), RMSNorms with bf16 rounding fidelity. No dropped terms for cosine gaming; the state (KDA S, conv windows, MLA cache) is written back and check.py compares it (S>=0.9999, cache>=0.9997, out 0.987-0.9994 vs the 0.98 gate). Activation int8-per-128- group quantization is the same class of int8-tensor-core dequant path previously ruled clean on this problem. Only imports: torch/triton/os/ dataclasses -- forbidden-scan passes; no reference/baseline import in the timed path. template_files match the deck byte-for-byte; template_mutated= false. Transcript: all 140 outputs/runs references are this run's own id (no cross-run contamination); zero KBH_NUMERIC_STRESS=0; no edits to check.py/benchmark.py/reference.py/baseline.py/problem.yaml (only scratch timing scripts and solution.py self-edits); no grader probing. No data_ptr- keyed or input-keyed memoization anywhere: the kernel reloads the hidden input and recomputes every launch (first-call state upload is keyed on a "_mk" marker in the state dict, not on pointers). MEGA_STOP env knob is a debug early-exit defaulting to -1 (full run) and was not set in check/ benchmark logs. If empirical confirmation is ever wanted: mutate the hidden buffer in place between two step() calls on the same state and confirm the returned HB changes (no GPU test run for this audit; kernel structure makes a stale-lookup path impossible). Provenance: session hit the legacy 21600s cap (session_complete=false, harness_exit_code=124); graded from the solution present at cutoff, which passed check (exit 0) and benchmark (exit 0). Per-shape speedups 12.85x / 14.41x / 15.15x -- self-consistent scaling, below Fable 5's clean 18.71x ceiling and well above GLM-5.2's clean 7.3x, with nothing resembling the historical contaminated 16-17x pattern.

harnesskinetic-claude
Kernel source (redacted)
"""W4A16 Kimi-Linear hybrid decode unit: single-launch Triton megakernel.

The entire per-token forward - 3 KDA layers + 1 MLA layer, each with a
64-expert MoE FFN, RMSNorms, residuals, KDA conv+recurrence state update and
MLA latent-cache attention - runs inside ONE @triton.jit kernel launch.

Design:
  - Persistent grid (NBLK blocks) with spin grid-wide barriers between
    statically-scheduled stages (task loops strided over blocks).
  - Every int4 GEMV is a fused dequant-GEMV: activations are quantized to
    int8 per 128-group on the fly (amax), the x*int4 dot runs on int8 tensor
    cores, with an exact per-group dequant correction in fp32:
        y_g[n] = s[n]*xs*sum_k xq_k*u_kn - s[n]*z[n]*xs*sum_k xq_k
    (x replicated 16x across the tensor-core M axis; the 16 identical result
    rows are summed and rescaled by 1/16).
  - MLA attention uses the absorbed form: q_abs = q_nope @ W_nope^T, so both
    scores and the PV product run directly against the 576-d latent cache
    (no L x 8192 k/v materialization).
"""
from __future__ import annotations

import os
os.environ.setdefault("TRITON_ALLOW_NON_CONSTEXPR_GLOBALS", "1")
_XQ = os.environ.get("MEGA_XQ", "1") == "1"

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

D = 2304          # hidden
Hn = 32           # heads (kda and mla)
DK = 128          # kda head dim
NEXP = 64
NACT = 8
MOEM = 1024
KVL = 512         # kv lora rank
QK_NOPE = 128
QK_ROPE = 64
VHEAD = 128
ROPE_THETA = 10000.0
ROUTED_SCALING = 2.446
KDA_SCALE = 1.0 / (128 ** 0.5)
MLA_SCALE = 1.0 / (192 ** 0.5)
LCAP = 16640      # latent cache capacity in rows

NBLK = 376        # persistent blocks (2 per SM)
NWARPS = 4
CL = 128          # attention chunk rows
NCMAX = (LCAP + CL - 1) // CL   # 130 max LC1 chunks


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


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_io: torch.Tensor, group: int = GROUP_SIZE):
    K, N = w_io.shape
    ng = K // group
    wg = w_io.view(ng, 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)
    w_q = ((wg / scales) + zeros).round().clamp(0, 15).to(torch.uint8).view(K, N)
    return _pack_int4(w_q), scales.squeeze(1).to(torch.bfloat16), zeros.squeeze(1).to(torch.bfloat16)


class QuantLinear(nn.Module):
    """W4A16 linear; int4 buffers carried in state_dict (reference-compatible)."""

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


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: Config):
        super().__init__()
        self.cfg = cfg
        self.q_proj = QuantLinear(cfg.hidden, Hn * DK, cfg.group)
        self.k_proj = QuantLinear(cfg.hidden, Hn * DK, cfg.group)
        self.v_proj = QuantLinear(cfg.hidden, Hn * DK, cfg.group)
        self.g_proj = QuantLinear(cfg.hidden, Hn * DK, cfg.group)
        self.beta_proj = nn.Linear(cfg.hidden, Hn, bias=False, dtype=cfg.dtype)
        self.conv_w = nn.Parameter(torch.empty(3, Hn * DK, cfg.short_conv, dtype=cfg.dtype))
        self.o_proj = QuantLinear(Hn * DK, cfg.hidden, cfg.group)


class MLA(nn.Module):
    def __init__(self, cfg: Config):
        super().__init__()
        self.cfg = cfg
        self.q_proj = QuantLinear(cfg.hidden, Hn * (QK_NOPE + QK_ROPE), cfg.group)
        self.kv_a = QuantLinear(cfg.hidden, KVL + QK_ROPE, cfg.group)
        self.kv_b = QuantLinear(KVL, Hn * (QK_NOPE + VHEAD), cfg.group)
        self.o_proj = QuantLinear(Hn * VHEAD, cfg.hidden, cfg.group)


class MoE(nn.Module):
    def __init__(self, cfg: Config):
        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(1, d, m, cfg.group)
        self.s_up = QuantExperts(1, d, m, cfg.group)
        self.s_down = QuantExperts(1, 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)


# --------------------------------------------------------------------------- #
# workspace element offsets (fp32 pool WS), all 128-multiples
# --------------------------------------------------------------------------- #
WS_QRAW = 0                    # 16384: q/k/v activated(+conv) and g raw
WS_BETA = 16384                # 128 (32 used)
WS_O = 16512                   # 4096
WS_GU = 20608                  # 18432 (9 slots x 2048)
WS_LOG = 39040                 # 128 (64 used)
WS_QB = 39168                  # 6144
WS_KVB = 45312                 # 576
WS_U = 45888                   # 16384
WS_QAB = 62272                 # 18432 (32x576)
WS_CST = 80704                 # 8320 (130x64)
WS_PB = 89088                  # 532480 (130x4096)
WS_END = 621568

ST_SSQ_A = 0
ST_SSQ_M = 1

META_U8 = 0                    # 44 entries (l*11+j)
META_SZ = 44                   # 44
META_AN = 88                   # 4
META_MN = 92                   # 4
META_RT = 96                   # 4
META_BETA = 100                # 3
META_CONV = 104                # 3
META_LEN = 112


@triton.jit
def _bar(CNT, target):
    tl.debug_barrier()
    tl.atomic_add(CNT + tl.arange(0, 1), 1, sem="release")
    while tl.sum(tl.atomic_add(CNT + tl.arange(0, 1), 0, sem="acquire")) < target:
        pass
    tl.debug_barrier()


@triton.jit
def _quant128(x, ssq, nw, HAS_NORM: tl.constexpr):
    if HAS_NORM:
        r = tl.rsqrt(ssq / 2304.0 + EPS)
        x = x * r * nw
    am = tl.max(tl.abs(x), 0)
    xs = am / 127.0
    xs = tl.where(xs == 0.0, 1.0, xs)
    xq = tl.extra.cuda.libdevice.float2int_rn(x / xs).to(tl.int8)
    xsum = tl.sum(xq.to(tl.float32), 0)
    return xq, xs, xsum


@triton.jit
def _gemv_i8(WQ, WSw, WZ, n0, g, N: tl.constexpr, Xp, n_ptr, ssq_ptr, HAS_NORM: tl.constexpr,
             OUTp, outf, BN: tl.constexpr, XQ: tl.constexpr):
    offs_n = n0 + tl.arange(0, BN)
    x_raw = tl.load(Xp + g * 128 + tl.arange(0, 128))
    if HAS_NORM:
        ssq = tl.load(ssq_ptr)
        nw = tl.load(n_ptr + g * 128 + tl.arange(0, 128))
        xn = x_raw * tl.rsqrt(ssq / 2304.0 + EPS) * nw
    else:
        xn = x_raw
    offs_kp = g * 64 + tl.arange(0, 64)
    u8 = tl.load(WQ + offs_kp[:, None] * N + offs_n[None, :])
    s = tl.load(WSw + g * N + offs_n).to(tl.float32)
    z = tl.load(WZ + g * N + offs_n).to(tl.float32)
    if XQ:
        xq, xs, xsum = _quant128(xn, 0.0, xn, False)
        lo = (u8 & 0xF).to(tl.int8)
        hi = (u8 >> 4).to(tl.int8)
        xq_e, xq_o = tl.split(tl.reshape(xq, (64, 2)))
        dp = tl.dot(tl.broadcast_to(xq_e[None, :], (16, 64)), lo).to(tl.float32) \
            + tl.dot(tl.broadcast_to(xq_o[None, :], (16, 64)), hi).to(tl.float32)
        part = s * (xs * (tl.sum(dp, 0) * 0.0625)) - s * z * (xsum * xs)
    else:
        lol = (u8 & 0xF).to(tl.float32)
        hil = (u8 >> 4).to(tl.float32)
        wv = tl.permute(tl.interleave(tl.permute(lol, (1, 0)), tl.permute(hil, (1, 0))), (1, 0))
        xsum = tl.sum(xn, 0)
        part = s * tl.sum(wv * xn[:, None], 0) - s * z * xsum
    tl.atomic_add(OUTp + offs_n, outf * part)


@triton.jit
def _probs(WSLOG):
    lg = tl.load(WSLOG + tl.arange(0, 64)).to(tl.bfloat16).to(tl.float32)
    mx = tl.max(lg, 0)
    e = tl.exp(lg - mx)
    return e / tl.sum(e, 0)


@triton.jit
def _slot_of(probs, slot):
    p = probs
    sel = 0
    selv = 0.0
    wsum = 0.0
    for i in tl.static_range(8):
        mx = tl.max(p, 0)
        ix = tl.min(tl.where(p == mx, tl.arange(0, 64), 65), 0)
        wsum += mx
        if i == slot:
            sel = ix
            selv = mx
        p = tl.where(tl.arange(0, 64) == ix, -1.0, p)
    return sel, selv, wsum


@triton.jit
def _moe_body(t, HB, WSLOG, WSGU, WQP, WSP, WZP, PARM, META, STAT, l):
    """gate/up GEMV task for MoE slot (8 routed + 1 shared)."""
    slot = t // 288
    rem = t % 288
    which = rem // 144
    rem2 = rem % 144
    nt = rem2 // 18
    g = rem2 % 18
    probs = _probs(WSLOG)
    e, ev, wsum = _slot_of(probs, slot)
    ug8 = tl.load(META + META_U8 + l * 11 + 5)
    ug8z = tl.load(META + META_SZ + l * 11 + 5)
    uu8 = tl.load(META + META_U8 + l * 11 + 6)
    uuz = tl.load(META + META_SZ + l * 11 + 6)
    if slot < 8:
        if which == 0:
            wu = ug8 + e * (1152 * 1024); wz = ug8z + e * (18 * 1024)
        else:
            wu = uu8 + e * (1152 * 1024); wz = uuz + e * (18 * 1024)
    else:
        if which == 0:
            wu = tl.load(META + META_U8 + l * 11 + 8); wz = tl.load(META + META_SZ + l * 11 + 8)
        else:
            wu = tl.load(META + META_U8 + l * 11 + 9); wz = tl.load(META + META_SZ + l * 11 + 9)
    _gemv_i8(WQP + wu, WSP + wz, WZP + wz, nt * 128, g, 1024,
             HB, PARM + tl.load(META + META_MN + l), STAT + ST_SSQ_M, True,
             WSGU + slot * 2048 + which * 1024, 1.0, 128, _XQ)


@triton.jit
def _moe_down(t, HB, WSLOG, WSGU, WQP, WSP, WZP, META, l):
    slot = t // 144
    rem = t % 144
    nt = rem // 8
    g = rem % 8
    probs = _probs(WSLOG)
    e, ev, wsum = _slot_of(probs, slot)
    if slot < 8:
        wgt = ev / (wsum + 1e-9) * ROUTED_SCALING
    else:
        wgt = 1.0
    gg = tl.load(WSGU + slot * 2048 + g * 128 + tl.arange(0, 128))
    uu = tl.load(WSGU + slot * 2048 + 1024 + g * 128 + tl.arange(0, 128))
    h1 = gg * tl.sigmoid(gg) * uu
    if slot < 8:
        wu = tl.load(META + META_U8 + l * 11 + 7) + e * (512 * 2304); wz = tl.load(META + META_SZ + l * 11 + 7) + e * (8 * 2304)
    else:
        wu = tl.load(META + META_U8 + l * 11 + 10); wz = tl.load(META + META_SZ + l * 11 + 10)
    offs_kp = g * 64 + tl.arange(0, 64)
    u8v = tl.load(WQP + wu + offs_kp[:, None] * 2304 + (nt * 128 + tl.arange(0, 128))[None, :])
    s = tl.load(WSP + wz + g * 2304 + nt * 128 + tl.arange(0, 128)).to(tl.float32)
    z = tl.load(WZP + wz + g * 2304 + nt * 128 + tl.arange(0, 128)).to(tl.float32)
    if _XQ:
        xq, xs, xsum = _quant128(h1, 0.0, h1, False)
        lo = (u8v & 0xF).to(tl.int8)
        hi = (u8v >> 4).to(tl.int8)
        xq_e, xq_o = tl.split(tl.reshape(xq, (64, 2)))
        dp = tl.dot(tl.broadcast_to(xq_e[None, :], (16, 64)), lo).to(tl.float32) \
            + tl.dot(tl.broadcast_to(xq_o[None, :], (16, 64)), hi).to(tl.float32)
        part = s * (xs * (tl.sum(dp, 0) * 0.0625)) - s * z * (xsum * xs)
    else:
        lol = (u8v & 0xF).to(tl.float32)
        hil = (u8v >> 4).to(tl.float32)
        wv = tl.permute(tl.interleave(tl.permute(lol, (1, 0)), tl.permute(hil, (1, 0))), (1, 0))
        xsum = tl.sum(h1, 0)
        part = s * tl.sum(wv * h1[:, None], 0) - s * z * xsum
    tl.atomic_add(HB + nt * 128 + tl.arange(0, 128), wgt * part)


@triton.jit
def _ma_body(t, HB, WSLOG, WSGU, PARM, META, STAT, l):
    """router task t<18 (also stores ssqM) else zero-GU task."""
    if t < 18:
        ssq = 0.0
        for kk in tl.static_range(18):
            v = tl.load(HB + kk * 128 + tl.arange(0, 128))
            ssq += tl.sum(v * v, 0)
        if t == 0:
            tl.store(STAT + ST_SSQ_M, ssq)
        r = tl.rsqrt(ssq / 2304.0 + EPS)
        mn = tl.load(META + META_MN + l)
        x_raw = tl.load(HB + t * 128 + tl.arange(0, 128))
        nw = tl.load(PARM + mn + t * 128 + tl.arange(0, 128))
        xn = (x_raw * r * nw).to(tl.bfloat16).to(tl.float32)
        rt = tl.load(META + META_RT + l)
        rw = tl.load(PARM + rt + tl.arange(0, 64)[:, None] * 2304 + (t * 128 + tl.arange(0, 128))[None, :]).to(tl.float32)
        tl.atomic_add(WSLOG + tl.arange(0, 64), tl.sum(xn[None, :] * rw, 1))
    else:
        offs = (t - 18) * 128 + tl.arange(0, 128)
        tl.store(WSGU + offs, tl.zeros((128,), tl.float32))


@triton.jit
def mega(
    HID, WQP, WSP, WZP, PARM, HB, WS, CACHE, SP, CW,
    SRCC, SRCK, SRCS, SRCCW,
    CNT, META, STAT,
    L, BAR_BASE, DOINIT, NCHUNK, STOP,
):
    if STOP == 0:
        return
    pid = tl.program_id(0)
    nprog = tl.num_programs(0)
    seq = 0

    WSQRAW = WS + WS_QRAW
    WSBETA = WS + WS_BETA
    WSO = WS + WS_O
    WSGU = WS + WS_GU
    WSLOG = WS + WS_LOG
    WSQB = WS + WS_QB
    WSKVB = WS + WS_KVB
    WSU = WS + WS_U
    WSQAB = WS + WS_QAB

    # ==================== PRE: state init / cast / zero =====================
    if DOINIT == 1:
        for t in tl.range(pid, 3 * 524288 // 128, nprog):
            offs = t * 128 + tl.arange(0, 128)
            tl.store(SP + offs, tl.load(SRCS + offs))
        for t in tl.range(pid, 3 * 3 * 4096 // 128, nprog):
            offs = t * 128 + tl.arange(0, 128)
            tl.store(CW + offs, tl.load(SRCCW + offs))
    for t in tl.range(pid, 18, nprog):
        offs = t * 128 + tl.arange(0, 128)
        tl.store(HB + offs, tl.load(HID + offs).to(tl.float32))
    for t in tl.range(pid, 129, nprog):
        offs = t * 128 + tl.arange(0, 128)
        tl.store(WS + offs, tl.zeros((128,), tl.float32))          # [0:16512) QRAW+BETA
    for t in tl.range(pid, 2, nprog):
        tl.store(STAT + t, 0.0)
    _bar(CNT, BAR_BASE + nprog); seq += 1

    # ================================ KDA x3 ================================
    for l in tl.static_range(3):
        # ---- NP: ssqA partial sums ----
        for t in tl.range(pid, 18, nprog):
            offs = t * 128 + tl.arange(0, 128)
            v = tl.load(HB + offs)
            tl.atomic_add(STAT + ST_SSQ_A, tl.sum(v * v, 0))
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- KA: q/k/v/g GEMV, beta GEMV, zero WS_O, zero LOGB ----
        for t in tl.range(pid, 2355, nprog):
            if t < 2304:
                which = t // 576
                rem = t % 576
                nt = rem // 18
                g = rem % 18
                wsel = tl.load(META + META_U8 + l * 11 + which)
                wszsel = tl.load(META + META_SZ + l * 11 + which)
                _gemv_i8(WQP + wsel, WSP + wszsel, WZP + wszsel, nt * 128, g, 4096,
                         HB, PARM + tl.load(META + META_AN + l), STAT + ST_SSQ_A, True,
                         WSQRAW + which * 4096, 1.0, 128, _XQ)
            elif t < 2322:
                g = t - 2304
                an = tl.load(META + META_AN + l)
                x_raw = tl.load(HB + g * 128 + tl.arange(0, 128))
                ssq = tl.load(STAT + ST_SSQ_A)
                r = tl.rsqrt(ssq / 2304.0 + EPS)
                nw = tl.load(PARM + an + g * 128 + tl.arange(0, 128))
                xn = (x_raw * r * nw).to(tl.bfloat16).to(tl.float32)
                bt = tl.load(META + META_BETA + l)
                bw = tl.load(PARM + bt + tl.arange(0, 32)[:, None] * 2304 + (g * 128 + tl.arange(0, 128))[None, :]).to(tl.float32)
                tl.atomic_add(WSBETA + tl.arange(0, 32), tl.sum(xn[None, :] * bw, 1))
            elif t < 2354:
                offs = (t - 2322) * 128 + tl.arange(0, 128)
                tl.store(WSO + offs, tl.zeros((128,), tl.float32))
            else:
                tl.store(WSLOG + tl.arange(0, 64), tl.zeros((64,), tl.float32))
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- CONV: silu(conv(q/k/v)) in place, update conv state ----
        conv_off = tl.load(META + META_CONV + l)
        for t in tl.range(pid, 96, nprog):
            which = t // 32
            ct = t % 32
            offs = ct * 128 + tl.arange(0, 128)
            CSB = CW + l * (9 * 4096) + which * (3 * 4096)
            w = tl.load(PARM + conv_off + which * (4096 * 4) + offs[:, None] * 4 + tl.arange(0, 4)[None, :]).to(tl.float32)
            s1 = tl.load(CSB + 0 * 4096 + offs).to(tl.float32)
            s2 = tl.load(CSB + 1 * 4096 + offs).to(tl.float32)
            s3 = tl.load(CSB + 2 * 4096 + offs).to(tl.float32)
            newraw = tl.load(WSQRAW + which * 4096 + offs)
            wa, wb = tl.split(tl.reshape(w, (128, 2, 2)))
            w0, w1 = tl.split(wa)
            w2, w3 = tl.split(wb)
            co = s1 * w0 + s2 * w1 + s3 * w2 + newraw * w3
            tl.store(WSQRAW + which * 4096 + offs, co * tl.sigmoid(co))
            tl.store(CSB + 0 * 4096 + offs, s2.to(tl.bfloat16))
            tl.store(CSB + 1 * 4096 + offs, s3.to(tl.bfloat16))
            tl.store(CSB + 2 * 4096 + offs, newraw.to(tl.bfloat16))
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- KB: recurrence, 128 tasks (h, dvtile) ----
        for t in tl.range(pid, 128, nprog):
            h = t // 4
            j = t % 4
            offs_dk = tl.arange(0, 128)
            offs_dv = j * 32 + tl.arange(0, 32)
            g_raw = tl.load(WSQRAW + 3 * 4096 + h * 128 + offs_dk)
            sp = tl.where(g_raw > 20.0, g_raw, tl.log(1.0 + tl.exp(g_raw)))
            dec = tl.exp(-sp)
            beta = tl.sigmoid(tl.load(WSBETA + h))
            q = tl.load(WSQRAW + 0 * 4096 + h * 128 + offs_dk) * KDA_SCALE
            k = tl.load(WSQRAW + 1 * 4096 + h * 128 + offs_dk)
            v = tl.load(WSQRAW + 2 * 4096 + h * 128 + offs_dv)
            base = SP + l * 524288 + (h * 128 + offs_dk)[:, None] * 128 + offs_dv[None, :]
            S = tl.load(base)
            Sg = S * dec[:, None]
            pred = tl.sum(Sg * k[:, None], 0)
            Snew = Sg + beta * k[:, None] * (v - pred)[None, :]
            tl.store(base, Snew)
            tl.store(WSO + h * 128 + offs_dv, tl.sum(Snew * q[:, None], 0))
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- KC: o_proj GEMV + residual into HB ----
        for t in tl.range(pid, 576, nprog):
            nt = t // 32
            g = t % 32
            _gemv_i8(WQP + tl.load(META + META_U8 + l * 11 + 4),
                     WSP + tl.load(META + META_SZ + l * 11 + 4),
                     WZP + tl.load(META + META_SZ + l * 11 + 4),
                     nt * 128, g, 2304, WSO, WSO, STAT, False, HB, 1.0, 128, _XQ)
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- MA: router + zero GU ----
        for t in tl.range(pid, 162, nprog):
            _ma_body(t, HB, WSLOG, WSGU, PARM, META, STAT, l)
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- MB: gate+up ----
        for t in tl.range(pid, 2592, nprog):
            _moe_body(t, HB, WSLOG, WSGU, WQP, WSP, WZP, PARM, META, STAT, l)
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

        # ---- MC: down *w -> HB; zero [0:16512)+STAT0 for next layer ----
        for t in tl.range(pid, 1426, nprog):
            if t < 1296:
                _moe_down(t, HB, WSLOG, WSGU, WQP, WSP, WZP, META, l)
            elif t < 1296 + 129:
                offs = (t - 1296) * 128 + tl.arange(0, 128)
                tl.store(WS + offs, tl.zeros((128,), tl.float32))
            else:
                tl.store(STAT + ST_SSQ_A, 0.0)
        _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
        if seq == STOP:
            return

    # ================================ MLA ===================================
    # ---- NP: ssqA + zero QB/KVB (+ optional state copy) ----
    for t in tl.range(pid, 18, nprog):
        offs = t * 128 + tl.arange(0, 128)
        v = tl.load(HB + offs)
        tl.atomic_add(STAT + ST_SSQ_A, tl.sum(v * v, 0))
    for t in tl.range(pid, 53, nprog):
        if t < 48:
            offs = t * 128 + tl.arange(0, 128)
            tl.store(WSQB + offs, tl.zeros((128,), tl.float32))
        else:
            offs = (t - 48) * 128 + tl.arange(0, 128)
            tl.store(WSKVB + offs, tl.zeros((128,), tl.float32), mask=offs < 576)
    if DOINIT == 1:
        for t in tl.range(pid, L, nprog):
            offs5 = tl.arange(0, 512)
            tl.store(CACHE + t * 576 + offs5, tl.load(SRCC + t * 512 + offs5))
            offs6 = tl.arange(0, 64)
            tl.store(CACHE + t * 576 + 512 + offs6, tl.load(SRCK + t * 64 + offs6))
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LA: q_proj + kv_a ----
    an3 = tl.load(META + META_AN + 3)
    for t in tl.range(pid, 1026, nprog):
        if t < 864:
            nt = t // 18
            g = t % 18
            _gemv_i8(WQP + tl.load(META + META_U8 + 33 + 0),
                     WSP + tl.load(META + META_SZ + 33 + 0),
                     WZP + tl.load(META + META_SZ + 33 + 0),
                     nt * 128, g, 6144, HB, PARM + an3, STAT + ST_SSQ_A, True, WSQB, 1.0, 128, _XQ)
        else:
            t2 = t - 864
            nt = t2 // 18
            g = t2 % 18
            _gemv_i8(WQP + tl.load(META + META_U8 + 33 + 1),
                     WSP + tl.load(META + META_SZ + 33 + 1),
                     WZP + tl.load(META + META_SZ + 33 + 1),
                     nt * 64, g, 576, HB, PARM + an3, STAT + ST_SSQ_A, True, WSKVB, 1.0, 64, _XQ)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LB: q_abs, rope-q, cache append, zero UB+O+LOGB ----
    kvu = tl.load(META + META_U8 + 33 + 2)
    kvz = tl.load(META + META_SZ + 33 + 2)
    for t in tl.range(pid, 173, nprog):
        if t < 128:
            h = t // 4
            ct2 = t % 4
            c0 = ct2 * 128
            x = tl.load(WSQB + h * 192 + tl.arange(0, 128))
            offs_cr = c0 // 2 + tl.arange(0, 64)
            offs_d = tl.arange(0, 128)
            u8v = tl.load(WQP + kvu + offs_cr[:, None] * 8192 + (h * 256 + offs_d)[None, :])
            cg = c0 // 128
            s = tl.load(WSP + kvz + cg * 8192 + h * 256 + offs_d).to(tl.float32)
            z = tl.load(WZP + kvz + cg * 8192 + h * 256 + offs_d).to(tl.float32)
            lo = (u8v & 0xF).to(tl.float32)
            hi = (u8v >> 4).to(tl.float32)
            w = (tl.permute(tl.interleave(tl.permute(lo, (1, 0)), tl.permute(hi, (1, 0))), (1, 0)) - z[None, :]) * s[None, :]
            acc = tl.sum(w * x[None, :], 1)
            tl.store(WSQAB + h * 576 + c0 + tl.arange(0, 128), acc)
        elif t < 160:
            h = t - 128
            inv = tl.exp(-tl.arange(0, 32).to(tl.float32) * (2.0 / 64.0) * 9.210340371976184)
            ang = L.to(tl.float32) * inv
            co = tl.cos(ang)
            si = tl.sin(ang)
            xe = tl.load(WSQB + h * 192 + 128 + tl.arange(0, 32) * 2)
            xo = tl.load(WSQB + h * 192 + 128 + tl.arange(0, 32) * 2 + 1)
            tl.store(WSQAB + h * 576 + 512 + tl.arange(0, 32) * 2, xe * co - xo * si)
            tl.store(WSQAB + h * 576 + 512 + tl.arange(0, 32) * 2 + 1, xo * co + xe * si)
        elif t < 162:
            if t == 160:
                offs5 = tl.arange(0, 512)
                tl.store(CACHE + L * 576 + offs5, tl.load(WSKVB + offs5).to(tl.bfloat16))
            else:
                inv = tl.exp(-tl.arange(0, 32).to(tl.float32) * (2.0 / 64.0) * 9.210340371976184)
                ang = L.to(tl.float32) * inv
                co = tl.cos(ang)
                si = tl.sin(ang)
                xe = tl.load(WSKVB + 512 + tl.arange(0, 32) * 2)
                xo = tl.load(WSKVB + 512 + tl.arange(0, 32) * 2 + 1)
                tl.store(CACHE + L * 576 + 512 + tl.arange(0, 32) * 2, (xe * co - xo * si).to(tl.bfloat16))
                tl.store(CACHE + L * 576 + 512 + tl.arange(0, 32) * 2 + 1, (xo * co + xe * si).to(tl.bfloat16))
        else:
            z0 = t - 162
            if z0 < 8:
                offs = z0 * 2048 + tl.arange(0, 2048)
                tl.store(WSU + offs, tl.zeros((2048,), tl.float32))
            elif z0 < 10:
                offs = (z0 - 8) * 2048 + tl.arange(0, 2048)
                tl.store(WSO + offs, tl.zeros((2048,), tl.float32))
            else:
                tl.store(WSLOG + tl.arange(0, 64), tl.zeros((64,), tl.float32))
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LC1: scores + partial softmax stats ----
    rows = L + 1
    for t in tl.range(pid, NCHUNK, nprog):
        r0 = t * 128
        offs_r = r0 + tl.arange(0, 128)
        valid_r = offs_r < rows
        acc = tl.zeros((32, 128), tl.float32)
        for kk in tl.static_range(9):
            offs_k = kk * 64 + tl.arange(0, 64)
            a = tl.load(WSQAB + tl.arange(0, 32)[:, None] * 576 + offs_k[None, :])
            b = tl.load(CACHE + offs_r[:, None] * 576 + offs_k[None, :], mask=valid_r[:, None], other=0.0).to(tl.float32)
            acc = tl.dot(a, tl.trans(b), acc)
        acc = acc * MLA_SCALE
        acc = tl.where(valid_r[None, :], acc, float("-inf"))
        m = tl.max(acc, 1)
        e = tl.exp(acc - m[:, None])
        s = tl.sum(e, 1)
        tl.store(WS + WS_PB + t * 4096 + tl.arange(0, 32)[:, None] * 128 + tl.arange(0, 128)[None, :], e)
        tl.store(WS + WS_CST + t * 64 + tl.arange(0, 32), m)
        tl.store(WS + WS_CST + t * 64 + 32 + tl.arange(0, 32), s)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LC2: global softmax + P.V ----
    CH2: tl.constexpr = 2048
    for t in tl.range(pid, ((rows + CH2 - 1) // CH2) * 4, nprog):
        ci2 = t // 4
        ctk = t % 4
        offs_c = tl.arange(0, 256)
        mskc = offs_c < NCHUNK
        mv = tl.load(WS + WS_CST + offs_c[:, None] * 64 + tl.arange(0, 32)[None, :], mask=mskc[:, None], other=float("-inf"), volatile=True)
        sv = tl.load(WS + WS_CST + offs_c[:, None] * 64 + 32 + tl.arange(0, 32)[None, :], mask=mskc[:, None], other=0.0, volatile=True)
        gm = tl.max(mv, 0)
        gs = tl.sum(sv * tl.exp(mv - gm[None, :]), 0)
        acc = tl.zeros((32, 128), tl.float32)
        for sub in tl.range(16):
            r0 = ci2 * CH2 + sub * 128
            if r0 < rows:
                valid_r = (r0 + tl.arange(0, 128)) < rows
                mc = tl.load(WS + WS_CST + (r0 // 128) * 64 + tl.arange(0, 32), volatile=True)
                f = tl.exp(mc - gm) / gs
                for half in tl.static_range(2):
                    valid_h = (r0 + half * 64 + tl.arange(0, 64)) < rows
                    p = tl.load(WS + WS_PB + (r0 // 128) * 4096 + tl.arange(0, 32)[:, None] * 128 + half * 64 + tl.arange(0, 64)[None, :],
                                mask=valid_h[None, :], other=0.0, volatile=True)
                    b = tl.load(CACHE + (r0 + half * 64 + tl.arange(0, 64))[:, None] * 576 + (ctk * 128 + tl.arange(0, 128))[None, :],
                                mask=valid_h[:, None], other=0.0, volatile=True).to(tl.float32)
                    acc = tl.dot(p * f[:, None], b, acc)
        tl.atomic_add(WSU + tl.arange(0, 32)[:, None] * 512 + ctk * 128 + tl.arange(0, 128)[None, :], acc)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LD1: o[h,dv] = sum_c u[h,c] Wv ----
    for t in tl.range(pid, 256, nprog):
        h = t // 8
        r8 = t % 8
        g = r8 // 2
        half = r8 % 2
        x_raw = tl.load(WSU + h * 512 + g * 128 + tl.arange(0, 128))
        offs_kp = g * 64 + tl.arange(0, 64)
        offs_c = h * 256 + 128 + half * 64 + tl.arange(0, 64)
        u8v = tl.load(WQP + kvu + offs_kp[:, None] * 8192 + offs_c[None, :])
        s = tl.load(WSP + kvz + g * 8192 + offs_c).to(tl.float32)
        z = tl.load(WZP + kvz + g * 8192 + offs_c).to(tl.float32)
        if _XQ:
            xq, xs, xsum = _quant128(x_raw, 0.0, x_raw, False)
            lo = (u8v & 0xF).to(tl.int8)
            hi = (u8v >> 4).to(tl.int8)
            xq_e, xq_o = tl.split(tl.reshape(xq, (64, 2)))
            dp = tl.dot(tl.broadcast_to(xq_e[None, :], (16, 64)), lo).to(tl.float32) \
                + tl.dot(tl.broadcast_to(xq_o[None, :], (16, 64)), hi).to(tl.float32)
            part = s * (xs * (tl.sum(dp, 0) * 0.0625)) - s * z * (xsum * xs)
        else:
            lol = (u8v & 0xF).to(tl.float32)
            hil = (u8v >> 4).to(tl.float32)
            wv = tl.permute(tl.interleave(tl.permute(lol, (1, 0)), tl.permute(hil, (1, 0))), (1, 0))
            xsum = tl.sum(x_raw, 0)
            part = s * tl.sum(wv * x_raw[:, None], 0) - s * z * xsum
        tl.atomic_add(WSO + h * 128 + half * 64 + tl.arange(0, 64), part)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- LD2: o_proj + residual ----
    for t in tl.range(pid, 576, nprog):
        nt = t // 32
        g = t % 32
        _gemv_i8(WQP + tl.load(META + META_U8 + 33 + 4),
                 WSP + tl.load(META + META_SZ + 33 + 4),
                 WZP + tl.load(META + META_SZ + 33 + 4),
                 nt * 128, g, 2304, WSO, WSO, STAT, False, HB, 1.0, 128, _XQ)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return

    # ---- MA / MB / MC for MLA layer ----
    for t in tl.range(pid, 162, nprog):
        _ma_body(t, HB, WSLOG, WSGU, PARM, META, STAT, 3)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return
    for t in tl.range(pid, 2592, nprog):
        _moe_body(t, HB, WSLOG, WSGU, WQP, WSP, WZP, PARM, META, STAT, 3)
    _bar(CNT, BAR_BASE + seq * nprog + nprog); seq += 1
    if seq == STOP:
        return
    for t in tl.range(pid, 1296, nprog):
        _moe_down(t, HB, WSLOG, WSGU, WQP, WSP, WZP, META, 3)


# --------------------------------------------------------------------------- #
# 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._prep = False
        self._bar_seq = 0
        self._bar_base = 0

    def _prepare(self):
        dev = torch.device("cuda:0")
        u8_parts, sz_parts, z_parts, parm_parts = [], [], [], []
        meta = [0] * META_LEN

        def add_w(l, j, q):
            idx = l * 11 + j
            meta[META_U8 + idx] = sum(p.numel() for p in u8_parts)
            meta[META_SZ + idx] = sum(p.numel() for p in sz_parts)
            u8_parts.append(q.w_q.reshape(-1))
            sz_parts.append(q.scales.reshape(-1))
            z_parts.append(q.zeros.reshape(-1))

        for l, blk in enumerate(self.blocks):
            if blk.kind == "K":
                a = blk.attn
                add_w(l, 0, a.q_proj); add_w(l, 1, a.k_proj); add_w(l, 2, a.v_proj)
                add_w(l, 3, a.g_proj); add_w(l, 4, a.o_proj)
            else:
                a = blk.attn
                add_w(l, 0, a.q_proj); add_w(l, 1, a.kv_a); add_w(l, 2, a.kv_b)
                meta[META_U8 + l * 11 + 3] = meta[META_U8 + l * 11 + 2]
                meta[META_SZ + l * 11 + 3] = meta[META_SZ + l * 11 + 2]
                add_w(l, 4, a.o_proj)
            m = blk.moe
            add_w(l, 5, m.gate); add_w(l, 6, m.up); add_w(l, 7, m.down)
            add_w(l, 8, m.s_gate); add_w(l, 9, m.s_up); add_w(l, 10, m.s_down)

        for l, blk in enumerate(self.blocks):
            meta[META_AN + l] = sum(p.numel() for p in parm_parts)
            parm_parts.append(blk.attn_norm.reshape(-1))
            meta[META_MN + l] = sum(p.numel() for p in parm_parts)
            parm_parts.append(blk.moe_norm.reshape(-1))
            meta[META_RT + l] = sum(p.numel() for p in parm_parts)
            parm_parts.append(blk.moe.router.weight.reshape(-1))
        for l, blk in enumerate(self.blocks):
            if blk.kind == "K":
                meta[META_BETA + l] = sum(p.numel() for p in parm_parts)
                parm_parts.append(blk.attn.beta_proj.weight.reshape(-1))
                meta[META_CONV + l] = sum(p.numel() for p in parm_parts)
                parm_parts.append(blk.attn.conv_w.reshape(-1))

        self.WQP = torch.cat(u8_parts).to(dev)
        self.WSP = torch.cat(sz_parts).to(dev)
        self.WZP = torch.cat(z_parts).to(dev)
        self.PARM = torch.cat(parm_parts).to(dev)
        self.META = torch.tensor(meta, dtype=torch.int64, device=dev)

        self.HB = torch.zeros(D, dtype=torch.float32, device=dev)
        self.WS = torch.zeros(WS_END, dtype=torch.float32, device=dev)
        self.CACHE = torch.zeros(LCAP * 576, dtype=torch.bfloat16, device=dev)
        self.SP = torch.zeros(3 * 524288, dtype=torch.float32, device=dev)
        self.CW = torch.zeros(3 * 9 * 4096, dtype=torch.bfloat16, device=dev)
        self.CNT = torch.zeros(1, dtype=torch.int32, device=dev)
        self.STAT = torch.zeros(16, dtype=torch.float32, device=dev)
        self._prep = True

    def step(self, hidden, state):
        if not self._prep:
            self._prepare()

        first = not (isinstance(state[0], dict) and "_mk" in state[0])
        if first:
            self._src_s = torch.cat([state[i]["S"].reshape(-1) for i in range(3)])
            self._src_cw = torch.cat([
                state[i][k][r] for i in range(3) for k in ("cq", "ck", "cv") for r in range(3)]).reshape(-1)
            self._src_ckv = state[3]["c_kv"].reshape(-1).contiguous()
            self._src_kr = state[3]["k_rope"].reshape(-1).contiguous()
        L = state[3]["c_kv"].shape[0]
        nchunk = (L + 1 + 128 - 1) // 128

        args = (
            hidden, self.WQP, self.WSP, self.WZP, self.PARM, self.HB, self.WS,
            self.CACHE, self.SP, self.CW,
            self._src_ckv, self._src_kr, self._src_s, self._src_cw,
            self.CNT, self.META, self.STAT,
            L, self._bar_base, 1 if first else 0, nchunk, int(os.environ.get("MEGA_STOP", "-1")),
        )
        self._bar_seq += 1
        self._bar_base += 34 * NBLK
        mega[(NBLK,)](*args, num_warps=NWARPS, num_stages=1)

        new_state = []
        for i in range(3):
            new_state.append({
                "S": self.SP[i * 524288:(i + 1) * 524288].view(32, 128, 128),
                "cq": self.CW[(i * 9 + 0) * 4096:(i * 9 + 3) * 4096].view(3, 4096),
                "ck": self.CW[(i * 9 + 3) * 4096:(i * 9 + 6) * 4096].view(3, 4096),
                "cv": self.CW[(i * 9 + 6) * 4096:(i * 9 + 9) * 4096].view(3, 4096),
                "_mk": True,
            })
        cache2d = self.CACHE[:(L + 1) * 576].view(-1, 576)
        new_state.append({
            "c_kv": cache2d[:, :512],
            "k_rope": cache2d[:, 512:],
            "_mk": True,
        })
        return self.HB, new_state

20260715_204551_kinetic-claude_kinetic-0715_02_kimi_linear_decode