kernelbench.com

KernelBench mega · H100

Kimi-Linear Decode Kimi K3 (256k)

14.82×geomean speedup across shapes

manually audited: clean

Clean H100 cell (14.8182x geomean speedup versus the eager baseline over contexts 2048/8192/16384). The submission is a genuine single-launch raw CUDA cooperative megakernel, not an output cache, constant answer, forbidden library call, CUDA graph, torch.compile wrapper, or per-op launch loop. It executes all four decode blocks in one kernel: three KDA blocks with fused int4 dequant GEMVs, short-convolution window updates, gated recurrent S updates, output projections and residuals; one MLA block with q/kv_a projections, RoPE, latent-cache append, absorbed kv_b attention over every live cache row, value/output projections and residual; and a router/top-8 MoE plus shared expert after every block. The private _ws identity marker only identifies state already resident in the model-owned live cache so it can avoid copying historical rows again. A foreign initialized state is copied in, continuous returned state already aliases the updated workspace, and every call appends a row, recomputes attention, mutates all three KDA states/windows, and overwrites the complete 2304-element output buffer. Therefore the pointer/identity reuse cannot replay a stale output in the checker or benchmark call structures.

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

The whole per-token forward (4 blocks: KDA,KDA,KDA,MLA; each attn + 64-expert MoE,
int4 fused dequant GEMVs, conv, recurrent state update, latent-cache attention,
router+topk, RMSNorms, residuals) is fused into ONE CUDA __global__ kernel launched
 cooperatively once per step() with grid-wide software barriers between phases.

Weights are repacked once (prepare, off the timed path) into flat blobs:
  WB: all packed int4 weights (uint8, (in//2, out) tiles concatenated)
  SB: all scales+zeros (bf16, per weight: [scales(G,N)][zeros(G,N)])
  AB: aux bf16 weights (norms, router, beta, conv)
The decode streams int4 bytes once; dequant uses the exact-integer magic bias
(0x4B000000|nibble -> fp32) so no dequantized bf16 matrix is ever materialized.

MLA uses the "absorb" form: scores are computed in latent space
  s[h,l] = (q_nope[h] @ Wk_B[h]) . c_kv[l] + q_rope[h] . k_rope[l]
  o[h]   = (sum_l p[h,l] c_kv[l]) @ Wv_B[h]
so the kv-cache is read once per step (576 bytes/token) instead of materializing
per-token k/v (which would be 16x the traffic of the weights at ctx 16k).
"""
from __future__ import annotations

import os
from dataclasses import dataclass, field

import torch
import torch.nn as nn
import torch.nn.functional as F

OP_TYPE = "kimi_linear_w4a16_decode"
HARDWARE_REQUIRED = ["RTX_PRO_6000"]
EPS = 1.0e-6
GROUP_SIZE = 128

# --------------------------------------------------------------------------- #
# dims (static for this problem)
# --------------------------------------------------------------------------- #
HID = 2304
C = 4096           # KDA channels (heads*head_dim)
NKH = 32           # kda heads
DK = 128
NMH = 32           # mla heads
KVL = 512
QKN = 128
QKR = 64
QKD = QKN + QKR
VH = 128
EXPERTS = 64
NACT = 8
NSHARED = 1
MINT = 1024
RSCALE = 2.446
GRP = 128
LMAX = 16512


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


# --------------------------------------------------------------------------- #
# layout of the flat blobs (kept in sync with the CUDA kernel)
# --------------------------------------------------------------------------- #
class Lay:
    """Byte/element offsets for every piece inside the flat blobs."""

    def __init__(self):
        # WB pieces: (name, nbytes)
        self.wb = {}
        wb_order = []
        self.sb = {}
        sb_order = []
        self.ab = {}
        ab_order = []

        def add_wb(name, nbytes):
            wb_order.append((name, nbytes))

        def add_sb(name, nelem):
            sb_order.append((name, nelem))

        def add_ab(name, nelem):
            ab_order.append((name, nelem))

        def qweight(tag, kin, kout):
            # packed int4: (kin//2) * kout bytes; scales/zeros: 2*(kin//128)*kout bf16
            add_wb(tag, (kin // 2) * kout)
            add_sb(tag, 2 * (kin // 128) * kout)

        for bi, kind in enumerate(("K", "K", "K", "M")):
            if kind == "K":
                for p in ("q", "k", "v", "g"):
                    qweight(f"b{bi}.{p}", HID, C)
                qweight(f"b{bi}.o", C, HID)
            else:
                qweight(f"b{bi}.q", HID, NMH * QKD)      # 2304 -> 6144
                qweight(f"b{bi}.kva", HID, KVL + QKR)    # 2304 -> 576
                qweight(f"b{bi}.kvb", KVL, NMH * (QKN + VH))  # 512 -> 8192
                qweight(f"b{bi}.o", NMH * VH, HID)       # 4096 -> 2304
            # MoE experts
            qweight(f"b{bi}.eg", HID, MINT * EXPERTS)    # (64, 1152, 1024) flattened contiguous ex-major
            qweight(f"b{bi}.eu", HID, MINT * EXPERTS)
            qweight(f"b{bi}.ed", MINT, HID * EXPERTS)    # (64, 512, 2304)
            qweight(f"b{bi}.sg", HID, MINT)
            qweight(f"b{bi}.su", HID, MINT)
            qweight(f"b{bi}.sd", MINT, HID)
            # aux
            add_ab(f"b{bi}.an", HID)
            add_ab(f"b{bi}.mn", HID)
            add_ab(f"b{bi}.rt", EXPERTS * HID)
            if kind == "K":
                add_ab(f"b{bi}.beta", NKH * HID)
                add_ab(f"b{bi}.conv", 3 * C * 4)

        off = 0
        for name, nbytes in wb_order:
            self.wb[name] = off
            off += nbytes
        self.wb_total = off
        off = 0
        for name, nelem in sb_order:
            self.sb[name] = off
            off += nelem
        self.sb_total = off
        off = 0
        for name, nelem in ab_order:
            self.ab[name] = off
            off += nelem
        self.ab_total = off

        # scales offset for a weight = sb[tag]; zeros at sb[tag] + (kin//128)*kout
        # NCOLS for each weight:
        self.ncols = {}
        self.kins = {}
        for bi, kind in enumerate(("K", "K", "K", "M")):
            if kind == "K":
                for p in ("q", "k", "v", "g"):
                    self.ncols[f"b{bi}.{p}"], self.kins[f"b{bi}.{p}"] = C, HID
                self.ncols[f"b{bi}.o"], self.kins[f"b{bi}.o"] = HID, C
            else:
                self.ncols[f"b{bi}.q"], self.kins[f"b{bi}.q"] = NMH * QKD, HID
                self.ncols[f"b{bi}.kva"], self.kins[f"b{bi}.kva"] = KVL + QKR, HID
                self.ncols[f"b{bi}.kvb"], self.kins[f"b{bi}.kvb"] = NMH * (QKN + VH), KVL
                self.ncols[f"b{bi}.o"], self.kins[f"b{bi}.o"] = HID, NMH * VH
            self.ncols[f"b{bi}.eg"], self.kins[f"b{bi}.eg"] = MINT, HID
            self.ncols[f"b{bi}.eu"], self.kins[f"b{bi}.eu"] = MINT, HID
            self.ncols[f"b{bi}.ed"], self.kins[f"b{bi}.ed"] = HID, MINT
            self.ncols[f"b{bi}.sg"], self.kins[f"b{bi}.sg"] = MINT, HID
            self.ncols[f"b{bi}.su"], self.kins[f"b{bi}.su"] = MINT, HID
            self.ncols[f"b{bi}.sd"], self.kins[f"b{bi}.sd"] = HID, MINT


LAY = Lay()

# scratch fp32 layout (element offsets)
SC_QKVG = 0          # 16384  (kda q,k,v,g raw)
SC_MLAQ = 16384      # 6144
SC_KV = 22528        # 640
SC_QR = 23168        # 2048
SC_QABS = 25216      # 16384
SC_CTX = 41536       # 16384
SC_O = 57920         # 4096
SC_MOEH = 62016      # 9216
SC_HACC = 71232      # 2304
SC_MOACC = 73536     # 2304
SC_LOGIT = 75840     # 64
SC_W8 = 75904        # 8
SC_IDS = 75912       # 8  (int32 view)
SC_XSUM = 75920      # 160 (per-64k-unit x sums, < 160)
SC_END = 76096
NCHUNK_MAX = 96
SC_PART = SC_END                    # NCHUNK_MAX * 32 * 514 fp32
SC_TOTAL = SC_PART + NCHUNK_MAX * 32 * 514 + 384 + 9216 * 2 + 64

# BAR (uint64) layout: [0..31] arrive, [32..63] release, [64..71] router flags,
# [72..72+8] router done counters, [80..143] work counters
BAR_ARR = 0
BAR_REL = 32
BAR_RFLAG = 64
BAR_RDONE = 72
BAR_WORK = 80
BAR_TOTAL = 144


# --------------------------------------------------------------------------- #
# quantization helpers (identical math 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 _unpack_int4(w_packed: torch.Tensor, K: int) -> torch.Tensor:
    out = torch.empty((K, w_packed.shape[1]), dtype=torch.uint8, device=w_packed.device)
    out[0::2] = w_packed & 0xF
    out[1::2] = (w_packed >> 4) & 0xF
    return out


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)


def dequant(w_q: torch.Tensor, scales: torch.Tensor, zeros: torch.Tensor, K: int, group: int) -> torch.Tensor:
    wu = _unpack_int4(w_q, K).to(torch.bfloat16)
    s = scales.repeat_interleave(group, dim=0)
    z = zeros.repeat_interleave(group, dim=0)
    return (wu - z) * s


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 weight_bf(self) -> torch.Tensor:
        return dequant(self.w_q, self.scales, self.zeros, self.in_f, self.group)


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 weight_bf(self, e: int) -> torch.Tensor:
        return dequant(self.w_q[e], self.scales[e], self.zeros[e], self.in_f, self.group)


def _rmsnorm(x: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
    xf = x.float()
    xf = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + EPS)
    return (xf * w.float()).to(x.dtype)


def _rope_cossin(pos: int, dim: int, theta: float, device):
    inv = 1.0 / (theta ** (torch.arange(0, dim, 2, device=device, dtype=torch.float32) / dim))
    ang = pos * inv
    return torch.cos(ang), torch.sin(ang)


def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
    xf = x.float()
    even, odd = xf[..., 0::2], xf[..., 1::2]
    out = torch.empty_like(xf)
    out[..., 0::2] = even * cos - odd * sin
    out[..., 1::2] = odd * cos + even * sin
    return out.to(x.dtype)


# --------------------------------------------------------------------------- #
# eager reference-path layers (debug / fallback; identical math to reference)
# --------------------------------------------------------------------------- #
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)
        self.scale = Dk ** -0.5

    def _short_conv(self, val, prev, idx):
        win = torch.cat([prev, val[None]], dim=0)
        w = self.conv_w[idx].float().transpose(0, 1)
        out = (win.float() * w).sum(0)
        return F.silu(out).to(val.dtype), win[1:]

    def _qlin(self, ql, x):
        return (x.float() @ ql.weight_bf().float()).to(torch.bfloat16)

    def step(self, x, st):
        H, Dk = self.cfg.kda_heads, self.cfg.kda_head_dim
        q = self._qlin(self.q_proj, x)
        k = self._qlin(self.k_proj, x)
        v = self._qlin(self.v_proj, x)
        q, st["cq"] = self._short_conv(q, st["cq"], 0)
        k, st["ck"] = self._short_conv(k, st["ck"], 1)
        v, st["cv"] = self._short_conv(v, st["cv"], 2)
        q = q.view(H, Dk).float() * self.scale
        k = k.view(H, Dk).float()
        v = v.view(H, Dk).float()
        g = (-F.softplus(self._qlin(self.g_proj, x).float())).view(H, Dk)
        beta = torch.sigmoid(self.beta_proj(x).float())
        S = st["S"] * g.exp()[:, :, None]
        pred = (S * k[:, :, None]).sum(1)
        S = S + beta[:, None, None] * k[:, :, None] * (v - pred)[:, None, :]
        o = (S * q[:, :, None]).sum(1)
        st["S"] = S
        return self._qlin(self.o_proj, o.reshape(H * Dk).to(torch.bfloat16))


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)
        self.scale = (cfg.qk_nope + cfg.qk_rope) ** -0.5

    def _qlin(self, ql, x):
        return (x.float() @ ql.weight_bf().float()).to(torch.bfloat16)

    def step(self, x, st):
        cfg = self.cfg
        H = cfg.mla_heads
        pos = st["c_kv"].shape[0]
        q = self._qlin(self.q_proj, x).view(H, cfg.qk_nope + cfg.qk_rope)
        q_nope = q[:, : cfg.qk_nope].float()
        q_rope = q[:, cfg.qk_nope :]
        kv = self._qlin(self.kv_a, x)
        c_kv = kv[: cfg.kv_lora]
        k_rope = kv[cfg.kv_lora :]
        cos, sin = _rope_cossin(pos, cfg.qk_rope, cfg.rope_theta, x.device)
        q_rope = _apply_rope(q_rope, cos, sin).float()
        k_rope = _apply_rope(k_rope, cos, sin)
        st["c_kv"] = torch.cat([st["c_kv"], c_kv[None]], 0)
        st["k_rope"] = torch.cat([st["k_rope"], k_rope[None]], 0)
        kvb = self._qlin(self.kv_b, st["c_kv"]).view(-1, H, cfg.qk_nope + cfg.v_head).float()
        k_nope = kvb[..., : cfg.qk_nope]
        v = kvb[..., cfg.qk_nope :]
        scores = (torch.einsum("hd,lhd->lh", q_nope, k_nope)
                  + torch.einsum("hd,ld->lh", q_rope, st["k_rope"].float())) * self.scale
        p = torch.softmax(scores, dim=0)
        o = torch.einsum("lh,lhd->hd", p, v)
        return self._qlin(self.o_proj, o.reshape(H * cfg.v_head).to(torch.bfloat16))


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)

    def _ffn(self, x, experts_g, experts_u, experts_d, e):
        h = F.silu(x.float() @ experts_g.weight_bf(e).float()) * (x.float() @ experts_u.weight_bf(e).float())
        return h @ experts_d.weight_bf(e).float()

    def step(self, x):
        cfg = self.cfg
        probs = torch.softmax(self.router(x).float(), dim=-1)
        w, idx = torch.topk(probs, cfg.n_active)
        w = w / (w.sum() + 1e-9) * cfg.routed_scaling
        out = x.new_zeros(cfg.hidden, dtype=torch.float32)
        for j in range(cfg.n_active):
            out = out + w[j] * self._ffn(x, self.gate, self.up, self.down, int(idx[j]))
        for s in range(cfg.n_shared):
            out = out + self._ffn(x, self.s_gate, self.s_up, self.s_down, s)
        return out.to(torch.bfloat16)


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)

    def step(self, x, st):
        h = x + self.attn.step(_rmsnorm(x, self.attn_norm), st)
        return h + self.moe.step(_rmsnorm(h, self.moe_norm))


# --------------------------------------------------------------------------- #
# CUDA megakernel (built by load_inline at prepare time)
# --------------------------------------------------------------------------- #
CUDA_SRC = r"""
// __MEGA_SRC_PLACEHOLDER__
"""

from mega_impl import build_cuda_source, extension  # noqa: E402


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._prepared = False
        self._ext = None
        self._spec = None
        self._gen = 0

    # -- weights arrive from the reference state dict; then repack once.
    def load_state_dict(self, *args, **kwargs):
        ret = super().load_state_dict(*args, **kwargs)
        if self.blocks[0].attn.q_proj.w_q.is_cuda:
            self.prepare()
        return ret

    _load = load_state_dict  # keep hook for older torch calling _load

    def prepare(self):
        if self._prepared:
            return
        cfg = self.cfg
        dev = self.blocks[0].attn.q_proj.w_q.device
        lay = LAY
        wb = torch.empty(lay.wb_total + 1024, dtype=torch.uint8, device=dev)
        sb = torch.empty(lay.sb_total + 1024, dtype=torch.bfloat16, device=dev)
        ab = torch.empty(lay.ab_total + 1024, dtype=torch.bfloat16, device=dev)

        def put_q(tag, ql_w2, ql_s, ql_z):
            o = lay.wb[tag]
            wv = ql_w2.reshape(-1)
            wb[o:o + wv.numel()] = wv
            so = lay.sb[tag]
            sv = ql_s.reshape(-1)
            sb[so:so + sv.numel()] = sv
            zv = ql_z.reshape(-1)
            sb[so + sv.numel():so + 2 * sv.numel()] = zv

        def put_ab(tag, t):
            o = lay.ab[tag]
            ab[o:o + t.numel()] = t.reshape(-1)

        for bi, blk in enumerate(self.blocks):
            if blk.kind == "K":
                a = blk.attn
                put_q(f"b{bi}.q", a.q_proj.w_q, a.q_proj.scales, a.q_proj.zeros)
                put_q(f"b{bi}.k", a.k_proj.w_q, a.k_proj.scales, a.k_proj.zeros)
                put_q(f"b{bi}.v", a.v_proj.w_q, a.v_proj.scales, a.v_proj.zeros)
                put_q(f"b{bi}.g", a.g_proj.w_q, a.g_proj.scales, a.g_proj.zeros)
                put_q(f"b{bi}.o", a.o_proj.w_q, a.o_proj.scales, a.o_proj.zeros)
                put_ab(f"b{bi}.beta", a.beta_proj.weight)
                put_ab(f"b{bi}.conv", a.conv_w)
            else:
                a = blk.attn
                put_q(f"b{bi}.q", a.q_proj.w_q, a.q_proj.scales, a.q_proj.zeros)
                put_q(f"b{bi}.kva", a.kv_a.w_q, a.kv_a.scales, a.kv_a.zeros)
                put_q(f"b{bi}.kvb", a.kv_b.w_q, a.kv_b.scales, a.kv_b.zeros)
                put_q(f"b{bi}.o", a.o_proj.w_q, a.o_proj.scales, a.o_proj.zeros)
            m = blk.moe
            put_q(f"b{bi}.eg", m.gate.w_q, m.gate.scales, m.gate.zeros)
            put_q(f"b{bi}.eu", m.up.w_q, m.up.scales, m.up.zeros)
            put_q(f"b{bi}.ed", m.down.w_q, m.down.scales, m.down.zeros)
            put_q(f"b{bi}.sg", m.s_gate.w_q[0], m.s_gate.scales[0], m.s_gate.zeros[0])
            put_q(f"b{bi}.su", m.s_up.w_q[0], m.s_up.scales[0], m.s_up.zeros[0])
            put_q(f"b{bi}.sd", m.s_down.w_q[0], m.s_down.scales[0], m.s_down.zeros[0])
            put_ab(f"b{bi}.an", blk.attn_norm)
            put_ab(f"b{bi}.mn", blk.moe_norm)
            put_ab(f"b{bi}.rt", m.router.weight)

        sc = torch.zeros(SC_TOTAL, dtype=torch.float32, device=dev)
        bar = torch.zeros(BAR_TOTAL, dtype=torch.int64, device=dev)
        ws_c = torch.zeros(LMAX, KVL, dtype=torch.bfloat16, device=dev)
        ws_k = torch.zeros(LMAX, QKR, dtype=torch.bfloat16, device=dev)
        x_out = torch.zeros(HID, dtype=torch.bfloat16, device=dev)

        # rope cos/sin tables, fp32, built with the exact reference formula
        theta = cfg.rope_theta
        inv = 1.0 / (theta ** (torch.arange(0, QKR, 2, device=dev, dtype=torch.float32) / QKR))
        posv = torch.arange(LMAX, device=dev, dtype=torch.float32)
        ang = torch.outer(posv, inv)
        cosc = torch.cos(ang).contiguous()
        sinc = torch.sin(ang).contiguous()

        # int64 offset vectors for the kernel (canonical piece order, see mega_impl):
        # wb/sb piece index = bi*11 + p
        #   KDA: 0 q,1 k,2 v,3 g,4 o,5 eg,6 eu,7 ed,8 sg,9 su,10 sd
        #   MLA: 0 q,1 kva,2 kvb,3 o,4 eg,5 eu,6 ed,7 sg,8 su,9 sd,(10 unused)
        # ab index = bi*5 + a: 0 attn_norm,1 moe_norm,2 router,3 beta,4 conv
        wb_off, sb_off, ab_off = [], [], []
        for bi in range(4):
            if bi < 3:
                pieces = ["q", "k", "v", "g", "o", "eg", "eu", "ed", "sg", "su", "sd"]
            else:
                pieces = ["q", "kva", "kvb", "o", "eg", "eu", "ed", "sg", "su", "sd"]
            for p in pieces:
                wb_off.append(lay.wb[f"b{bi}.{p}"])
                sb_off.append(lay.sb[f"b{bi}.{p}"])
            if bi == 3:
                wb_off.append(0)
                sb_off.append(0)
        for bi in range(4):
            ab_off.extend([lay.ab[f"b{bi}.an"], lay.ab[f"b{bi}.mn"], lay.ab[f"b{bi}.rt"]])
            if bi < 3:
                ab_off.extend([lay.ab[f"b{bi}.beta"], lay.ab[f"b{bi}.conv"]])
            else:
                ab_off.extend([0, 0])

        self._wb, self._sb, self._ab, self._sc, self._bar = wb, sb, ab, sc, bar
        self._ws_c, self._ws_k = ws_c, ws_k
        self._x_out = x_out
        self._cosc, self._sinc = cosc, sinc
        self._ext = extension()
        self._ext.setup(wb, sb, ab, sc, bar, ws_c, ws_k, x_out, cosc, sinc,
                        wb_off, sb_off, ab_off)
        self._prepared = True

    # ------------------------------------------------------------------ #
    def step(self, hidden, state):
        if not self._prepared:
            self.prepare()
        return self._step_mega(hidden, state)

    def _step_mega(self, hidden, state):
        cfg = self.cfg
        ext = self._ext
        mla_idx = 3
        st = state
        m = st[mla_idx]
        L = m["c_kv"].shape[0]
        fresh = 0 if m.get("_ws") is not None and m["_ws"][0] is self._ws_c else 1
        if fresh:
            ckv_src = m["c_kv"]
            kr_src = m["k_rope"]
        else:
            ckv_src = m["c_kv"]
            kr_src = m["k_rope"]
        # attention chunking: aim ~64-128 items for CTA coverage, min 8-subchunk items
        tgt = max(16, min(NCHUNK_MAX, (L + 1 + 170) // 171))
        nchunk = tgt
        ch = (L + 1 + nchunk - 1) // nchunk
        if fresh or self._spec is None:
            self._spec = [
                hidden.data_ptr(),
                st[0]["S"].data_ptr(), st[0]["cq"].data_ptr(), st[0]["ck"].data_ptr(), st[0]["cv"].data_ptr(),
                st[1]["S"].data_ptr(), st[1]["cq"].data_ptr(), st[1]["ck"].data_ptr(), st[1]["cv"].data_ptr(),
                st[2]["S"].data_ptr(), st[2]["cq"].data_ptr(), st[2]["ck"].data_ptr(), st[2]["cv"].data_ptr(),
                ckv_src.data_ptr(), kr_src.data_ptr(),
            ]
        else:
            self._spec[0] = hidden.data_ptr()
            self._spec[13] = ckv_src.data_ptr()
            self._spec[14] = kr_src.data_ptr()
        spec = self._spec[:15] + [int(L), int(fresh), int(nchunk), int(ch), int(self._gen), int(-1)]
        ext.mstep(spec)
        m["c_kv"] = self._ws_c[: L + 1]
        m["k_rope"] = self._ws_k[: L + 1]
        m["_ws"] = (self._ws_c, self._ws_k)
        self._gen += 1
        return self._x_out, state

    def _step_eager(self, hidden, state):
        for i, blk in enumerate(self.blocks):
            hidden = blk.step(hidden, state[i])
        return hidden, state


# ==================================================================
# ===== sidecar: mega_impl.py (52755 bytes, loaded by solution.py) =====
# ==================================================================

"""Build + cache the CUDA megakernel extension for solution.py.

Single __global__ kernel: whole 4-block decode step (KDA x3, MLA x1, each + MoE).
Cooperative launch, custom 2-word grid barriers, dynamic atomic work counters.

Weight layout (flat blobs, produced by solution.Model.prepare):
  WB (uint8): per piece, packed int4 rows (in//2, out) row-major.
  SB (bf16) : per piece, [scales(G,N) flat][zeros(G,N) flat], G = in//128.
              expert pieces are ex-major: scales (E,G,N) flat then zeros (E,G,N).
  AB (bf16) : norms, router (64,2304), beta (32,2304), conv (3,4096,4).

Offset-vector piece order (Dev.wb/sb index = bi*11 + p):
  KDA: 0 q,1 k,2 v,3 g,4 o,5 eg,6 eu,7 ed,8 sg,9 su,10 sd
  MLA: 0 q,1 kva,2 kvb,3 o,4 eg,5 eu,6 ed,7 sg,8 su,9 sd,(10 unused)
Dev.ab index = bi*5 + a:  0 attn_norm,1 moe_norm,2 router,3 beta,4 conv
"""
from __future__ import annotations

import torch
from torch.utils.cpp_extension import load_inline

_EXT = None

CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <ATen/cuda/CUDAContext.h>

using bf16 = __nv_bfloat16;
using ull = unsigned long long;

#define THR 256
#define NWARP 8
#define HID 2304
#define CC 4096
#define KVL 512
#define QKR 64
#define QKN 128
#define QKD 192
#define NMH 32
#define VH 128
#define DK 128
#define EXP 64
#define MINT 1024
#define RSCALE 2.446f
#define DKSCALE 0.08838834764831845f   // 128^-0.5
#define MLASCALE 0.07216878364870323f  // 192^-0.5
#define W_EPS 1e-9f

// scratch offsets (must match solution.py)
#define SC_QKVG 0
#define SC_MLAQ 16384
#define SC_KV 22528
#define SC_QR 23168
#define SC_QABS 25216
#define SC_CTX 41536
#define SC_O 57920
#define SC_MOEH 62016
#define SC_HACC 71232
#define SC_MOACC 73536
#define SC_LOGIT 75840
#define SC_W8 75904
#define SC_IDS 75912
#define SC_XSUM 75920
#define SC_BETA (SC_XSUM+96)
#define SC_PART 76096
#define SC_PROF (76096 + 96*32*514)
#define SC_MOEHG (SC_PROF + 384)
#define SC_MOEHU (SC_MOEHG + 9216)
#define SC_MOEHG_IDX(j, k) (SC_MOEHG + (size_t)(j) * MINT + (k))
#define SC_MOEHU_IDX(j, k) (SC_MOEHU + (size_t)(j) * MINT + (k))
// BAR (ull) offsets
#define BAR_ARR 0
#define BAR_REL 32
#define BAR_RFLAG 64
#define BAR_RDONE 72
#define BAR_WORK 80

struct Dev {
  const uint8_t* WB;
  const bf16* SB;
  const bf16* AB;
  float* SC;
  ull* BAR;
  bf16* WSC;
  bf16* WSK;
  bf16* XOUT;
  const float* COSC;
  const float* SINC;
  long long wb[44];
  long long sb[44];
  long long ab[20];
};

struct Call {
  const bf16* x_in;
  float* S0; bf16* cq0; bf16* ck0; bf16* cv0;
  float* S1; bf16* cq1; bf16* ck1; bf16* cv1;
  float* S2; bf16* cq2; bf16* ck2; bf16* cv2;
  const bf16* ckv; const bf16* krc;
  int L; int fresh; int nchunk; int ch; int dbg_stop;
  long long gen;
};

struct SmGemv {
  float xn[4096];
  float xsu[128];
  float red[2 * NWARP][128];
  uint8_t wroll[NWARP][2][4096];
};
struct SmAttn {
  bf16 qt[592][40];     // 46KB: transposed q (576 feat rows, 32 head cols + pad)
  bf16 cpos[32][520];   // 33KB: pos-major c rows (576 used + pad)
  bf16 psm[32 * 40];    // 2.5KB: P tiles (32 h x 40 pos)
  float ctab[64];       // per-warp col max/sum temps
  float msum[32];       // l (per head)
  float mrow[32];       // m (per head)
  float alp[32];        // exp(m_old - m_new) per head
};
struct SmS {
  float qc[128];
  float kc[128];
  float vc[128];
  float gc[128];
  float red[NWARP][32];
  float aux[32];
};
union SmU {
  SmGemv g;
  SmAttn a;
  SmS s;
};

__device__ __forceinline__ float b2f(bf16 v) { return __bfloat162float(v); }
__device__ __forceinline__ bf16 f2b(float v) { return __float2bfloat16(v); }
__device__ __forceinline__ uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }
// lane q fills 16B segment (q%8)*16 of tile-local row (q/8): batch of 4 tile rows starting at (src_row0) of the matrix
__device__ __forceinline__ void cp_unit_tile(uint8_t* dst_row0, const uint8_t* wp, int N, int src_row0, int c0, int lane) {
  int r = lane >> 3;
  int s16 = lane & 7;
  const uint8_t* srcp = wp + (size_t)(src_row0 + r) * (long long)N + c0 + s16 * 16;
  uint8_t* dstp = dst_row0 + r * 128 + s16 * 16;
  asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(smem_u32(dstp)), "l"(srcp));
}
__device__ __forceinline__ void cp_commit() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void cp_wait1() { asm volatile("cp.async.wait_group 1;"); }
__device__ __forceinline__ void cp_wait0() { asm volatile("cp.async.wait_group 0;"); }

// A-frag: lane l covers rows l/4 (+8), cols 2(l%4)(+8) of a 16x16 tile at (r0, k0)
__device__ __forceinline__ void ldmatrix_A(uint32_t a[4], const void* base, int row_stride, int lane) {
  // base: bf16 smem ptr to tile origin (r0, k0); row_stride in elems
  const bf16* p = (const bf16*)base + (lane % 16) * row_stride + (lane / 16) * 8;
  asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
               : "=r"(a[0]), "=r"(a[1]), "=r"(a[2]), "=r"(a[3])
               : "r"(smem_u32(p)));
}
// B-frag (trans): lane l covers cols l/4, k-rows 2(l%4)(+8) of a 16x8 (k x n) tile at (k0, n0)
__device__ __forceinline__ void ldmatrix_Btrans(uint32_t b[2], const void* base, int row_stride, int lane) {
  // base: bf16 smem ptr to tile origin (k0, n0); row_stride in elems (row = k)
  const bf16* p = (const bf16*)base + (lane % 16) * row_stride;
  asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
               : "=r"(b[0]), "=r"(b[1])
               : "r"(smem_u32(p)));
}
__device__ __forceinline__ void mma_bf16(float d[4], const uint32_t a[4], const uint32_t b[2]) {
  asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};"
               : "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3])
               : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]));
}
__device__ __forceinline__ float magic(uint32_t nib) {
  return __uint_as_float(0x4B000000u | nib) - 8388608.0f;
}

__device__ __forceinline__ void gsync(const Dev& D, int slot, ull tgt, int G) {
  __syncthreads();
  __threadfence();
  if (threadIdx.x == 0) {
    ull old = atomicAdd(D.BAR + BAR_ARR + slot, 1ULL);
    if (old == tgt - 1) atomicExch(D.BAR + BAR_REL + slot, tgt);
    volatile ull* r = D.BAR + BAR_REL + slot;
    ull v = *r;
    while (v < tgt) { v = *r; }
  }
  __threadfence();
  __syncthreads();
}

__device__ __forceinline__ int next_work(const Dev& D, int cid, int* sm_t) {
  if (threadIdx.x == 0) *sm_t = (int)atomicAdd(D.BAR + BAR_WORK + cid, 1ULL);
  __syncthreads();
  return *sm_t;
}

// block reduce helper: returns total sum of per-thread value
__device__ __forceinline__ float block_sum(float v, float* sm_red) {
  #pragma unroll
  for (int o = 16; o > 0; o >>= 1) v += __shfl_down_sync(0xffffffffu, v, o);
  if ((threadIdx.x & 31) == 0) sm_red[threadIdx.x >> 5] = v;
  __syncthreads();
  float s = 0.f;
  if (threadIdx.x < NWARP) s = sm_red[threadIdx.x];
  #pragma unroll
  for (int o = 4; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
  if (threadIdx.x == 0) sm_red[0] = s;
  __syncthreads();
  return sm_red[0];
}

// stage normed x into smg.xn with exact bf16 rounding; returns rstd. Also fills xsu.
__device__ float stage_xn(SmGemv& smg, const bf16* xb, const float* xf, const bf16* nrm,
                          float* sm_red) {
  float part = 0.f;
  for (int k = threadIdx.x; k < HID; k += THR) {
    float xv = xb ? b2f(xb[k]) : b2f(f2b(xf[k]));
    smg.xn[k] = xv;
    part += xv * xv;
  }
  float ssq = block_sum(part, sm_red);
  float rstd = 1.0f / sqrtf(ssq / (float)HID + 1e-6f);
  for (int k = threadIdx.x; k < HID; k += THR)
    smg.xn[k] = b2f(f2b(smg.xn[k] * rstd * b2f(nrm[k])));
  __syncthreads();
  return rstd;
}

// stage normed x for blocks > 0: x_k = bf16(HACC_k + bf16(DELTA_k)) exactly like torch
__device__ float stage_xn_blk(SmGemv& smg, const float* hacc, const float* delta, const bf16* nrm,
                              float* sm_red) {
  float part = 0.f;
  for (int k = threadIdx.x; k < HID; k += THR) {
    float xb = b2f(f2b(hacc[k] + b2f(f2b(delta[k]))));
    smg.xn[k] = xb;
    part += xb * xb;
  }
  float ssq = block_sum(part, sm_red);
  float rstd = 1.0f / sqrtf(ssq / (float)HID + 1e-6f);
  for (int k = threadIdx.x; k < HID; k += THR)
    smg.xn[k] = b2f(f2b(smg.xn[k] * rstd * b2f(nrm[k])));
  __syncthreads();
  return rstd;
}

// fill xsu for len/64 units from staged xn
__device__ void compute_xsu(SmGemv& smg, int units) {
  int wid = threadIdx.x >> 5, lane = threadIdx.x & 31;
  for (int u = wid; u < units; u += NWARP) {
    float s = smg.xn[u * 64 + lane * 2] + smg.xn[u * 64 + lane * 2 + 1];
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
    if (lane == 0) smg.xsu[u] = s;
  }
  __syncthreads();
}

// stage raw fp32 x of given length (+ xsu)
__device__ void stage_raw(SmGemv& smg, const float* src, int len) {
  for (int k = threadIdx.x; k < len; k += THR) smg.xn[k] = src[k];
  __syncthreads();
  compute_xsu(smg, len >> 6);
}

// fused dequant GEMV tile: 128 cols x full-K, k-split across warps.
// MODE 0: store fp32 (y = rstd*v) | 1: fp32 residual (y = res + v) | 2: atomic wgt add
// MODE 3: bf16 residual (y = b2f(res_bf16) + v)
template <int MODE>
__device__ void gemv_tile(const Dev& D, SmGemv& smg, long long wb, long long sb, int N, int K,
                          int col0, float rstd, float* y, const void* res, float wgt,
                          long long zoff, const void* res2 = nullptr, int out0 = -1, int KARG0 = 0, int KARG1 = 0) {
  const int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
  const uint8_t* wp = D.WB + wb;
  long long zbase = sb + zoff;
  if (out0 < 0) out0 = col0;
  int c = col0 + lane * 4;
  float acc[4] = {0.f, 0.f, 0.f, 0.f};
  int units = (K >> 1) >> 5;
  if (KARG1 > 0) units = min(units, (KARG1 >> 1) >> 5);
  int ustart = (KARG0 >> 1) >> 5;
  int nu = 0;
  for (int u = ustart + wid; u < units; u += NWARP) nu++;
  if (nu > 0) {
    int u = ustart + wid;
    int bufp = 0;
    // prologue: load first tile
    uint8_t* dst0 = smg.wroll[wid][0];
    #pragma unroll
    for (int it = 0; it < 8; it++) cp_unit_tile(dst0 + it * 512, wp, N, u * 32 + it * 4, col0, lane);
    cp_commit();
    for (int itn = 0; itn < nu; itn++, u += NWARP) {
      int nxt = u + NWARP;
      if (KARG0 > 0) {}
      int r0 = u * 32;
      if (nxt < units) {
        uint8_t* dstd = smg.wroll[wid][bufp ^ 1];
        #pragma unroll
        for (int it = 0; it < 8; it++) cp_unit_tile(dstd + it * 512, wp, N, nxt * 32 + it * 4, col0, lane);
        cp_commit();
        cp_wait1();
      } else {
        cp_wait0();
      }
      const uint8_t* tile = smg.wroll[wid][bufp];
      bufp ^= 1;
      float dot[4] = {0.f, 0.f, 0.f, 0.f};
      #pragma unroll
      for (int ii = 0; ii < 32; ii += 4) {
        float xl0 = smg.xn[2 * (r0 + ii)], xh0 = smg.xn[2 * (r0 + ii) + 1];
        float xl1 = smg.xn[2 * (r0 + ii) + 2], xh1 = smg.xn[2 * (r0 + ii) + 3];
        float xl2 = smg.xn[2 * (r0 + ii) + 4], xh2 = smg.xn[2 * (r0 + ii) + 5];
        float xl3 = smg.xn[2 * (r0 + ii) + 6], xh3 = smg.xn[2 * (r0 + ii) + 7];
        #pragma unroll
        for (int j = 0; j < 4; j++) {
          uint32_t wv = ((const uint32_t*)tile)[(ii + j) * 32 + lane & (1024 - 1)];
          float xl = (j == 0) ? xl0 : (j == 1 ? xl1 : (j == 2 ? xl2 : xl3));
          float xh = (j == 0) ? xh0 : (j == 1 ? xh1 : (j == 2 ? xh2 : xh3));
          #pragma unroll
          for (int b = 0; b < 4; b++) {
            uint32_t by = (wv >> (8 * b)) & 0xFFu;
            dot[b] += magic(by & 0xFu) * xl + magic(by >> 4) * xh;
          }
        }
      }
      int g = u >> 1;
      uint2 s4 = *(const uint2*)(D.SB + sb + (size_t)g * N + c);
      uint2 z4 = *(const uint2*)(D.SB + zbase + (size_t)g * N + c);
      float xs = smg.xsu[u];
      #pragma unroll
      for (int b = 0; b < 4; b++) {
        acc[b] += b2f(((const bf16*)&s4)[b]) * (dot[b] - b2f(((const bf16*)&z4)[b]) * xs);
      }
    }
  }
  #pragma unroll
  for (int b = 0; b < 4; b++) smg.red[wid][lane * 4 + b] = acc[b];
  __syncthreads();
  int l = threadIdx.x & 31;
  if (l < 16) {
    int col = (threadIdx.x >> 5) * 16 + l;
    float v = 0.f;
    #pragma unroll
    for (int w = 0; w < NWARP; w++) v += smg.red[w][col];
    v *= rstd;
    int cg = out0 + col;
    if (MODE == 0) y[cg] = v;
    else if (MODE == 4) y[cg] = b2f(f2b(v));
    else if (MODE == 1) {
      // composite torch-exact: h = bf16( bf16(prevH + bf16(prevD)) + bf16(v) )
      float xb = b2f(f2b(((const float*)res)[cg] + b2f(f2b(((const float*)res2)[cg]))));
      y[cg] = b2f(f2b(xb + b2f(f2b(v))));
    } else if (MODE == 2) atomicAdd(y + cg, wgt * v);
    else y[cg] = b2f(f2b(b2f(((const bf16*)res)[cg]) + b2f(f2b(v))));
  }
  __syncthreads();
}

// gate+up sequential GEMV: one accumulator set per matrix (register-lean)
__device__ void gu_tile(const Dev& D, SmGemv& smg, long long wbg, long long wbu,
                        long long sbg, long long sbu, int N, int K, int col0, float rstd,
                        float* out, long long zoff) {
  const int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
  int c = col0 + lane * 4;
  int units = (K >> 1) >> 5;
  float accg[4] = {0.f, 0.f, 0.f, 0.f};
  float accu[4] = {0.f, 0.f, 0.f, 0.f};
  for (int pass = 0; pass < 2; pass++) {
    const uint8_t* wp = pass == 0 ? (D.WB + wbg) : (D.WB + wbu);
    float* accp = pass == 0 ? accg : accu;
    long long szb = pass == 0 ? sbg : sbu;
    for (int u = wid; u < units; u += NWARP) {
      int r0 = u * 32;
      float dot[4] = {0.f, 0.f, 0.f, 0.f};
      const uint8_t* rp = wp + (size_t)r0 * N + c;
      #pragma unroll
      for (int ii = 0; ii < 32; ii += 8) {
        uint32_t w[8];
        #pragma unroll
        for (int j = 0; j < 8; j++) w[j] = ((const uint32_t*)rp)[(ii + j) * (N >> 2)];
        #pragma unroll
        for (int j = 0; j < 8; j++) {
          uint32_t wv = w[j];
          float xl = smg.xn[2 * (r0 + ii + j)];
          float xh = smg.xn[2 * (r0 + ii + j) + 1];
          #pragma unroll
          for (int b = 0; b < 4; b++) {
            uint32_t by = (wv >> (8 * b)) & 0xFFu;
            dot[b] += magic(by & 0xFu) * xl + magic(by >> 4) * xh;
          }
        }
      }
      int g = u >> 1;
      uint2 s4 = *(const uint2*)(D.SB + szb + (size_t)g * N + c);
      uint2 z4 = *(const uint2*)(D.SB + szb + zoff + (size_t)g * N + c);
      float xs = smg.xsu[u];
      #pragma unroll
      for (int b = 0; b < 4; b++) accp[b] += b2f(((const bf16*)&s4)[b]) * (dot[b] - b2f(((const bf16*)&z4)[b]) * xs);
    }
  }
  #pragma unroll
  for (int b = 0; b < 4; b++) {
    smg.red[wid][lane * 4 + b] = accg[b];
    smg.red[wid + NWARP][lane * 4 + b] = accu[b];
  }
  __syncthreads();
  int l = threadIdx.x & 31, wid2 = threadIdx.x >> 5;
  float gv = 0.f, uv = 0.f;
  int col = wid2 * 16 + l;
  bool act = l < 16;
  if (act) {
    #pragma unroll
    for (int w = 0; w < NWARP; w++) {
      gv += smg.red[w][col];
      uv += smg.red[w + NWARP][col];
    }
    gv *= rstd;
    uv *= rstd;
    float hv = gv / (1.f + expf(-gv)) * uv;
    out[col0 + col] = hv;
  }
  __syncthreads();
}

// gate+up K-split: accumulate scaled partial sums, atomicAdd into MOEH_G / MOEH_U
__device__ void gu_tile_atomic(const Dev& D, SmGemv& smg, long long wbg, long long wbu,
                               long long sbg, long long sbu, int N, int K, int col0, float rstd,
                               float* outg, float* outu, long long zoff, int KARG0, int KARG1) {
  const int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
  const uint8_t* wgp = D.WB + wbg;
  const uint8_t* wup = D.WB + wbu;
  int c = col0 + lane * 4;
  float accg[4] = {0.f, 0.f, 0.f, 0.f};
  float accu[4] = {0.f, 0.f, 0.f, 0.f};
  int units = (K >> 1) >> 5;
  if (KARG1 > 0) units = min(units, (KARG1 >> 1) >> 5);
  int ustart = (KARG0 >> 1) >> 5;
  for (int u = ustart + wid; u < units; u += NWARP) {
    int r0 = u * 32;
    float dg[4] = {0.f, 0.f, 0.f, 0.f};
    float du[4] = {0.f, 0.f, 0.f, 0.f};
    const uint8_t* rg = wgp + (size_t)r0 * N + c;
    const uint8_t* ru = wup + (size_t)r0 * N + c;
    #pragma unroll
    for (int ii = 0; ii < 32; ii += 8) {
      uint32_t wg[8], wu[8];
      #pragma unroll
      for (int j = 0; j < 8; j++) {
        wg[j] = ((const uint32_t*)rg)[(ii + j) * (N >> 2)];
        wu[j] = ((const uint32_t*)ru)[(ii + j) * (N >> 2)];
      }
      #pragma unroll
      for (int j = 0; j < 8; j++) {
        uint32_t wgv = wg[j], wuv = wu[j];
        float xl = smg.xn[2 * (r0 + ii + j)];
        float xh = smg.xn[2 * (r0 + ii + j) + 1];
        #pragma unroll
        for (int b = 0; b < 4; b++) {
          uint32_t bg = (wgv >> (8 * b)) & 0xFFu;
          dg[b] += magic(bg & 0xFu) * xl + magic(bg >> 4) * xh;
          uint32_t bu = (wuv >> (8 * b)) & 0xFFu;
          du[b] += magic(bu & 0xFu) * xl + magic(bu >> 4) * xh;
        }
      }
    }
    int g = u >> 1;
    uint2 sg4 = *(const uint2*)(D.SB + sbg + (size_t)g * N + c);
    uint2 zg4 = *(const uint2*)(D.SB + sbg + zoff + (size_t)g * N + c);
    uint2 su4 = *(const uint2*)(D.SB + sbu + (size_t)g * N + c);
    uint2 zu4 = *(const uint2*)(D.SB + sbu + zoff + (size_t)g * N + c);
    float xs = smg.xsu[u];
    #pragma unroll
    for (int b = 0; b < 4; b++) {
      accg[b] += b2f(((const bf16*)&sg4)[b]) * (dg[b] - b2f(((const bf16*)&zg4)[b]) * xs);
      accu[b] += b2f(((const bf16*)&su4)[b]) * (du[b] - b2f(((const bf16*)&zu4)[b]) * xs);
    }
  }
  float v0 = accg[0] * rstd, v1 = accg[1] * rstd, v2 = accg[2] * rstd, v3 = accg[3] * rstd;
  float w0 = accu[0] * rstd, w1 = accu[1] * rstd, w2 = accu[2] * rstd, w3 = accu[3] * rstd;
  atomicAdd(outg + c + 0, v0);
  atomicAdd(outg + c + 1, v1);
  atomicAdd(outg + c + 2, v2);
  atomicAdd(outg + c + 3, v3);
  atomicAdd(outu + c + 0, w0);
  atomicAdd(outu + c + 1, w1);
  atomicAdd(outu + c + 2, w2);
  atomicAdd(outu + c + 3, w3);
}

// router + moe gate/up + down for one block (identical for KDA/MLA blocks).
// po: MoE piece offset within the block's 11 pieces (KDA 5, MLA 4)
#define GSYNC_M() do { if (cl.dbg_stop == 9999) { __threadfence(); __syncthreads(); unsigned long long c; asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(c)); if (threadIdx.x == 0) { atomicMax((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1) + 1), c); atomicMin((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1) + 48), c); } } gsync(D, slot, tgt, G); if (cl.dbg_stop == 9999) { unsigned long long c2; asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(c2)); if (threadIdx.x == 0) atomicMax((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1)), c2); } slot++; if (cl.dbg_stop >= 0 && slot > cl.dbg_stop) return; } while (0)
__device__ void moe_phase(const Dev& D, SmU& sm, int bi, const Call& cl, int cid_gu, int cid_dn,
                          const bf16* mnrm, float* sm_red, int* sm_t, int& slot, ull tgt, int G, int po) {
  // stage normed h
  float rstd3 = stage_xn(sm.g, nullptr, D.SC + SC_HACC, mnrm, sm_red);
  compute_xsu(sm.g, HID >> 6);
  float* SC = D.SC;
  const int wbbase = bi * 11;
  const int abbase = bi * 5;

  int NT3 = 4 + 16 + 2 + 128;
  while (true) {
    int t = next_work(D, cid_gu, sm_t);
    if (t >= NT3) break;
    if (t < 4) {
      // router k-chunk (576 of 2304) over 64 cols
      const bf16* rt = D.AB + D.ab[abbase + 2];
      int k0 = t * 576;
      int wid = threadIdx.x >> 5, lane = threadIdx.x & 31;
      for (int cc2 = wid * 8; cc2 < wid * 8 + 8; cc2++) {
        float s = 0.f;
        for (int k = k0 + lane * 18; k < k0 + lane * 18 + 18; k++)
          s += sm.g.xn[k] * b2f(rt[(size_t)cc2 * HID + k]);
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
        if (lane == 0) atomicAdd(SC + SC_LOGIT + cc2, s);
      }
      __syncthreads();
      if (threadIdx.x == 0) {
        int old = (int)atomicAdd(D.BAR + BAR_RDONE + bi, 1ULL);
        *sm_t = (old == 3) ? -1 : t;
      }
      __syncthreads();
      if (*sm_t == -1 && threadIdx.x < 32) {
        // last router chunk done: softmax + top8 (warp 0)
        float p2[2];
        p2[0] = b2f(f2b(SC[SC_LOGIT + threadIdx.x * 2]));
        p2[1] = b2f(f2b(SC[SC_LOGIT + threadIdx.x * 2 + 1]));
        float mx = fmaxf(p2[0], p2[1]);
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) mx = fmaxf(mx, __shfl_down_sync(0xffffffffu, mx, o));
        mx = __shfl_sync(0xffffffffu, mx, 0);
        p2[0] = expf(p2[0] - mx);
        p2[1] = expf(p2[1] - mx);
        float ss = p2[0] + p2[1];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) ss += __shfl_down_sync(0xffffffffu, ss, o);
        ss = __shfl_sync(0xffffffffu, ss, 0);
        p2[0] /= ss;
        p2[1] /= ss;
        for (int j = 0; j < 8; j++) {
          float best = -1.f;
          int bi2 = -1;
          if (p2[0] > best) { best = p2[0]; bi2 = threadIdx.x * 2; }
          if (p2[1] > best) { best = p2[1]; bi2 = threadIdx.x * 2 + 1; }
          #pragma unroll
          for (int o = 16; o > 0; o >>= 1) {
            float ob = __shfl_down_sync(0xffffffffu, best, o);
            int oi = __shfl_down_sync(0xffffffffu, bi2, o);
            if (ob > best || (ob == best && oi >= 0 && oi < bi2)) { best = ob; bi2 = oi; }
          }
          best = __shfl_sync(0xffffffffu, best, 0);
          bi2 = __shfl_sync(0xffffffffu, bi2, 0);
          if (threadIdx.x == 0) {
            ((int*)(SC + SC_IDS))[j] = bi2;
            SC[SC_W8 + j] = best;
          }
          __syncwarp();
          if (threadIdx.x * 2 == bi2) p2[0] = -1.f;
          if (threadIdx.x * 2 + 1 == bi2) p2[1] = -1.f;
          __syncwarp();
        }
        if (threadIdx.x == 0) {
          float ws = 0.f;
          for (int j = 0; j < 8; j++) ws += SC[SC_W8 + j];
          for (int j = 0; j < 8; j++) SC[SC_W8 + j] = SC[SC_W8 + j] / (ws + W_EPS) * RSCALE;
          __threadfence();
          atomicExch(D.BAR + BAR_RFLAG + bi, (ull)cl.gen + 1);
        }
      }
    } else if (t < 20) {
      int td = t - 4;
      int j2 = td >> 1;
      int kh = td & 1;
      int col = j2 * 128;
      gu_tile_atomic(D, sm.g, D.wb[wbbase + po + 3], D.wb[wbbase + po + 4], D.sb[wbbase + po + 3], D.sb[wbbase + po + 4],
              MINT, HID, col, 1.f, SC + SC_MOEHG + (size_t)8 * MINT, SC + SC_MOEHU + (size_t)8 * MINT,
              (HID >> 7) * (long long)MINT, kh * (HID >> 1), (kh + 1) * (HID >> 1));
    } else if (t < 22) {
      int c0 = (t - 20) * 1152;
      for (int k = c0 + threadIdx.x; k < c0 + 1152; k += THR) SC[SC_MOACC + k] = 0.f;
    } else {
      if (threadIdx.x == 0) {
        volatile ull* f = D.BAR + BAR_RFLAG + bi;
        ull v = *f;
        while (v < (ull)cl.gen + 1) { v = *f; }
      }
      __threadfence();
      __syncthreads();
      int td = t - 22;
      int j = td >> 4;
      int col = ((td >> 1) & 7) * 128;
      int kh = td & 1;
      int e = ((int*)(SC + SC_IDS))[j];
      if (e < 0 || e >= EXP) { if (threadIdx.x == 0) D.SC[7] = 4000.f + bi * 100.f + j; e = 0; }
      long long wbe = D.wb[wbbase + po + 0] + (size_t)e * ((HID / 2) * MINT);
      long long wue = D.wb[wbbase + po + 1] + (size_t)e * ((HID / 2) * MINT);
      long long sbe = D.sb[wbbase + po + 0] + (size_t)e * ((HID >> 7) * MINT);
      long long sue = D.sb[wbbase + po + 1] + (size_t)e * (long long)((HID >> 7) * MINT);
      gu_tile_atomic(D, sm.g, wbe, wue, sbe, sue, MINT, HID, col, 1.f,
              SC + SC_MOEHG + (size_t)j * MINT, SC + SC_MOEHU + (size_t)j * MINT,
              (HID >> 7) * (long long)MINT * EXP, kh * (HID >> 1), (kh + 1) * (HID >> 1));
    }
  }
  GSYNC_M();
  // ---- down ----
  int NT4 = 162;
  int* doneflag = sm_t;  // reuse
  while (true) {
    int t = next_work(D, cid_dn, doneflag);
    if (t >= NT4) break;
    int j = t / 18;
    int col = (t % 18) * 128;
    int kh = 0;
    // stage FFN hidden with silu: h = silu(g)*u (fp32)
    for (int k = threadIdx.x; k < MINT; k += THR) {
      float g = SC[SC_MOEHG_IDX(j, k)];
      sm.g.xn[k] = g / (1.f + expf(-g)) * SC[SC_MOEHU_IDX(j, k)];
    }
    __syncthreads();
    compute_xsu(sm.g, MINT >> 6);
    if (j < 8) {
      int e = ((int*)(SC + SC_IDS))[j];
      if (e < 0 || e >= EXP) { if (threadIdx.x == 0) D.SC[7] = 6000.f + bi * 100.f + j; e = 0; }
      float wgt = SC[SC_W8 + j];
      gemv_tile<2>(D, sm.g, D.wb[wbbase + po + 2] + (size_t)e * ((MINT / 2) * HID),
                   D.sb[wbbase + po + 2] + (size_t)e * ((MINT >> 7) * HID), HID, MINT, col, 1.f,
                   SC + SC_MOACC, nullptr, wgt, (MINT >> 7) * (long long)HID * EXP);
    } else {
      gemv_tile<2>(D, sm.g, D.wb[wbbase + po + 5], D.sb[wbbase + po + 5], HID, MINT, col, 1.f,
                   SC + SC_MOACC, nullptr, 1.f, (MINT >> 7) * (long long)HID);
    }
  }
}

__global__ void mega(const Dev D, const Call cl) {
  extern __shared__ char smem_raw[];
  SmU& sm = *(SmU*)smem_raw;
  __shared__ int sm_t;
  __shared__ float sm_red[32];
  const int G = gridDim.x;
  const ull tgt = (ull)(cl.gen + 1) * (ull)G;
  int slot = 0;
#define GSYNC() do { if (cl.dbg_stop == 9999) { __threadfence(); __syncthreads(); unsigned long long c; asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(c)); if (threadIdx.x == 0) { atomicMax((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1) + 1), c); atomicMin((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1) + 48), c); } } gsync(D, slot, tgt, G); if (cl.dbg_stop == 9999) { unsigned long long c2; asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(c2)); if (threadIdx.x == 0) atomicMax((unsigned long long*)((unsigned long long*)(D.SC + SC_PROF) + (slot << 1)), c2); } slot++; if (cl.dbg_stop >= 0 && slot > cl.dbg_stop) return; } while (0)
  float* SC = D.SC;

  float* Sptr[3] = {cl.S0, cl.S1, cl.S2};
  bf16* cq[3] = {(bf16*)cl.cq0, (bf16*)cl.cq1, (bf16*)cl.cq2};
  bf16* ck[3] = {(bf16*)cl.ck0, (bf16*)cl.ck1, (bf16*)cl.ck2};
  bf16* cv[3] = {(bf16*)cl.cv0, (bf16*)cl.cv1, (bf16*)cl.cv2};

  for (int bi = 0; bi < 4; bi++) {
    const int wbbase = bi * 11;
    const int abbase = bi * 5;
    const bf16* xb = (bi == 0) ? cl.x_in : nullptr;
    const float* xf = (bi == 0) ? nullptr : (SC + SC_MOACC);
    const bf16* anrm = D.AB + D.ab[abbase + 0];
    const float* xf_blk = (bi == 0) ? nullptr : (SC + SC_HACC);
    const float* xf_blk2 = (bi == 0) ? nullptr : (SC + SC_MOACC);
    const bf16* mnrm = D.AB + D.ab[abbase + 1];
    const int c0 = bi * 5;   // KDA counter base; MLA uses 15..

    if (bi < 3) {
      // =================== KDA block =================== //
      float rstd = (bi == 0) ? stage_xn(sm.g, xb, xf, anrm, sm_red)
                             : stage_xn_blk(sm.g, xf_blk, xf_blk2, anrm, sm_red);
      compute_xsu(sm.g, HID >> 6);

      int NT = 129;
      while (true) {
        int t = next_work(D, c0 + 0, &sm_t);
        if (t >= NT) break;
        if (t < 128) {
          int mat = t >> 5;
          int col = (t & 31) * 128;
          gemv_tile<0>(D, sm.g, D.wb[wbbase + mat], D.sb[wbbase + mat], CC, HID, col, 1.f,
                       SC + SC_QKVG + mat * CC, nullptr, 0.f, (HID >> 7) * (long long)CC);
        } else {
          // side: zero logits + rdone; beta
          for (int k = threadIdx.x; k < 64; k += THR) SC[SC_LOGIT + k] = 0.f;
          __syncthreads();
          if (threadIdx.x == 0) D.BAR[BAR_RDONE + bi] = 0;
          const bf16* bw = D.AB + D.ab[abbase + 3];
          int h = threadIdx.x >> 5, lane = threadIdx.x & 31;
          if (h < 32) {
            float s = 0.f;
            for (int k = lane * 72; k < lane * 72 + 72; k++) s += sm.g.xn[k] * b2f(bw[(size_t)h * HID + k]);
            #pragma unroll
            for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
            if (lane == 0) SC[SC_BETA + h] = b2f(f2b(s));
          }
          __syncthreads();
        }
      }
      GSYNC();

      // ---- P1: conv + S update ----
      {
        int NT1 = 128;
        while (true) {
          int t = next_work(D, c0 + 1, &sm_t);
          if (t >= NT1) break;
          int h = t >> 2;
          int dv0 = (t & 3) * 32;
          if (threadIdx.x < 128) {
            int c = h * 128 + threadIdx.x;
            #pragma unroll
            for (int qv = 0; qv < 3; qv++) {
              const bf16* wnd = qv == 0 ? cq[bi] : (qv == 1 ? ck[bi] : cv[bi]);
              const bf16* cw = D.AB + D.ab[abbase + 4] + (size_t)qv * CC * 4;
              float val = b2f(f2b(SC[SC_QKVG + qv * CC + c]));
              float o = 0.f;
              o += b2f(wnd[0 * CC + c]) * b2f(cw[c * 4 + 0]);
              o += b2f(wnd[1 * CC + c]) * b2f(cw[c * 4 + 1]);
              o += b2f(wnd[2 * CC + c]) * b2f(cw[c * 4 + 2]);
              o += val * b2f(cw[c * 4 + 3]);
              float sv = o / (1.f + expf(-o));
              sv = b2f(f2b(sv));
              if (qv == 0) sm.s.qc[threadIdx.x] = sv * DKSCALE;
              else if (qv == 1) sm.s.kc[threadIdx.x] = sv;
              else sm.s.vc[threadIdx.x] = sv;
            }
            float gg = b2f(f2b(SC[SC_QKVG + 3 * CC + c]));
            float sp = logf(1.f + expf(-fabsf(gg))) + fmaxf(gg, 0.f);
            sm.s.gc[threadIdx.x] = expf(-sp);
          }
          __syncthreads();
          float beta = 1.f / (1.f + expf(-SC[SC_BETA + h]));
          int wid = threadIdx.x >> 5, lane = threadIdx.x & 31;
          float* Sh = Sptr[bi] + (size_t)h * (DK * DK);
          float o = 0.f;
          for (int pass = 0; pass < 2; pass++) {
            int dk0 = wid * 8 + pass * 64;
            float sv[8];
            #pragma unroll
            for (int i = 0; i < 8; i++) {
              int dk = dk0 + i;
              sv[i] = Sh[(size_t)dk * DK + dv0 + lane] * sm.s.gc[dk];
            }
            float pr = 0.f;
            #pragma unroll
            for (int i = 0; i < 8; i++) pr += sv[i] * sm.s.kc[dk0 + i];
            sm.s.red[wid][lane] = pr;
            __syncthreads();
            #pragma unroll
            for (int w = 0; w < NWARP; w++) pr += sm.s.red[w][lane];
            float pred = pr;
            float vp = sm.s.vc[dv0 + lane] - pred;
            #pragma unroll
            for (int i = 0; i < 8; i++) {
              int dk = dk0 + i;
              float snew = sv[i] + beta * sm.s.kc[dk] * vp;
              Sh[(size_t)dk * DK + dv0 + lane] = snew;
              o += snew * sm.s.qc[dk];
            }
            __syncthreads();
          }
          sm.s.red[wid][lane] = o;
          __syncthreads();
          if (wid == 0) {
            float oo = 0.f;
            #pragma unroll
            for (int w = 0; w < NWARP; w++) oo += sm.s.red[w][lane];
            SC[SC_O + h * 128 + dv0 + lane] = b2f(f2b(oo));
          }
          __syncthreads();
        }
      }
      GSYNC();

      // ---- P2: o_proj (+residual) + conv-window shift ----
      {
        stage_raw(sm.g, SC + SC_O, CC);
        int NT2 = 30 + 3;
        while (true) {
          int t = next_work(D, c0 + 2, &sm_t);
          if (t >= NT2) break;
          if (t < 18) {
            int col = t * 128;
            if (bi == 0)
              gemv_tile<3>(D, sm.g, D.wb[wbbase + 4], D.sb[wbbase + 4], HID, CC, col, 1.f,
                           SC + SC_HACC, cl.x_in, 0.f, (CC >> 7) * (long long)HID);
            else
              gemv_tile<1>(D, sm.g, D.wb[wbbase + 4], D.sb[wbbase + 4], HID, CC, col, 1.f,
                           SC + SC_HACC, SC + SC_HACC, 0.f, (CC >> 7) * (long long)HID, SC + SC_MOACC);
          } else if (t < 21) {
            // zero MOEHG/U for this block's MoE
            int c0 = (t - 18) * 6144;
            for (int k = c0 + threadIdx.x; k < c0 + 6144; k += THR) {
              SC[SC_MOEHG + k] = 0.f;
              SC[SC_MOEHU + k] = 0.f;
            }
          } else {
            int tt = t - 21;
            int qv = tt / 4;
            int cb = (tt % 4) * 1024;
            bf16* wnd = qv == 0 ? cq[bi] : (qv == 1 ? ck[bi] : cv[bi]);
            for (int cc2 = cb + threadIdx.x; cc2 < cb + 1024; cc2 += THR) {
              bf16 r1 = wnd[1 * CC + cc2];
              bf16 r2 = wnd[2 * CC + cc2];
              bf16 val = f2b(SC[SC_QKVG + qv * CC + cc2]);
              wnd[0 * CC + cc2] = r1;
              wnd[1 * CC + cc2] = r2;
              wnd[2 * CC + cc2] = val;
            }
          }
        }
      }
      GSYNC();
      // ---- P3/P4: MoE ----
      moe_phase(D, sm, bi, cl, c0 + 3, c0 + 4, mnrm, sm_red, &sm_t, slot, tgt, G, 5);
      GSYNC();
    } else {
      // =================== MLA block =================== //
      float rstd = (bi == 0) ? stage_xn(sm.g, xb, xf, anrm, sm_red)
                             : stage_xn_blk(sm.g, xf_blk, xf_blk2, anrm, sm_red);
      compute_xsu(sm.g, HID >> 6);
      {
        int NT = 54;
        while (true) {
          int t = next_work(D, 15, &sm_t);
          if (t >= NT) break;
          if (t < 48) {
            int col = t * 128;
            gemv_tile<0>(D, sm.g, D.wb[wbbase + 0], D.sb[wbbase + 0], NMH * QKD, HID, col, 1.f,
                         SC + SC_MLAQ, nullptr, 0.f, (HID >> 7) * (long long)(NMH * QKD));
          } else if (t < 53) {
            int col = (t - 48) * 128;
            gemv_tile<0>(D, sm.g, D.wb[wbbase + 1], D.sb[wbbase + 1], KVL + QKR, HID, col, 1.f,
                         SC + SC_KV, nullptr, 0.f, (HID >> 7) * (long long)(KVL + QKR));
          } else {
            for (int k = threadIdx.x; k < 64; k += THR) SC[SC_LOGIT + k] = 0.f;
            __syncthreads();
            if (threadIdx.x == 0) D.BAR[BAR_RDONE + bi] = 0;
          }
        }
      }
      GSYNC();

      // ---- P1: q_abs absorb + rope + cache append (+copy if fresh) ----
      {
        int NT = 261 + (cl.fresh ? cl.L : 0);
        __shared__ float qs[128];
        __shared__ float zqv[2];
        while (true) {
          int t = next_work(D, 16, &sm_t);
          if (t >= NT) break;
          if (t < 256) {
            int h = t >> 3;
            int jc = t & 7;
            int g = jc >> 1;
            if (threadIdx.x < 128) {
              int d = threadIdx.x;
              float qv = b2f(f2b(SC[SC_MLAQ + h * QKD + d]));
              float s = b2f(D.SB[D.sb[wbbase + 2] + (size_t)g * (NMH * 256) + h * 256 + d]);
              float z = b2f(D.SB[D.sb[wbbase + 2] + (size_t)4 * (NMH * 256) + (size_t)g * (NMH * 256) + h * 256 + d]);
              float qsc = qv * s;
              qs[d] = qsc;
              float zp = z * qsc;
              #pragma unroll
              for (int o = 16; o > 0; o >>= 1) zp += __shfl_down_sync(0xffffffffu, zp, o);
              if ((threadIdx.x & 31) == 0) sm_red[threadIdx.x >> 5] = zp;
            }
            __syncthreads();
            if (threadIdx.x == 0) zqv[0] = sm_red[0] + sm_red[1] + sm_red[2] + sm_red[3];
            __syncthreads();
            if (threadIdx.x < 32) {
              int lane = threadIdx.x;
              int j2 = jc * 32 + lane;
              float zq = zqv[0];
              const uint8_t* rp = D.WB + D.wb[wbbase + 2] + (size_t)j2 * (NMH * 256) + h * 256;
              float dlo = 0.f, dhi = 0.f;
              #pragma unroll
              for (int i = 0; i < 32; i++) {
                uint32_t wv = ((const uint32_t*)rp)[i];
                #pragma unroll
                for (int b = 0; b < 4; b++) {
                  uint32_t by = (wv >> (8 * b)) & 0xFFu;
                  float qsv = qs[i * 4 + b];
                  dlo += magic(by & 0xFu) * qsv;
                  dhi += magic(by >> 4) * qsv;
                }
              }
              SC[SC_QABS + h * KVL + 2 * j2] = dlo - zq;
              SC[SC_QABS + h * KVL + 2 * j2 + 1] = dhi - zq;
            }
            __syncthreads();
          } else if (t < 260) {
            int e0 = (t - 256) * 512;
            for (int e = e0 + threadIdx.x; e < e0 + 512; e += THR) {
              int h = e / QKR;
              int i = e % QKR;
              int ii = i >> 1;
              float cs = D.COSC[(size_t)cl.L * (QKR / 2) + ii];
              float sn = D.SINC[(size_t)cl.L * (QKR / 2) + ii];
              int ie = i & ~1;
              float xe = b2f(f2b(SC[SC_MLAQ + h * QKD + QKN + ie]));
              float xo = b2f(f2b(SC[SC_MLAQ + h * QKD + QKN + ie + 1]));
              float outv = (i & 1) ? (xo * cs + xe * sn) : (xe * cs - xo * sn);
              SC[SC_QR + h * QKR + i] = b2f(f2b(outv));
            }
          } else if (t == 260) {
            for (int e = threadIdx.x; e < KVL; e += THR) D.WSC[(size_t)cl.L * KVL + e] = f2b(SC[SC_KV + e]);
            for (int e = threadIdx.x; e < QKR; e += THR) {
              int ii = e >> 1;
              float cs = D.COSC[(size_t)cl.L * (QKR / 2) + ii];
              float sn = D.SINC[(size_t)cl.L * (QKR / 2) + ii];
              int ie = e & ~1;
              float xe = b2f(f2b(SC[SC_KV + KVL + ie]));
              float xo = b2f(f2b(SC[SC_KV + KVL + ie + 1]));
              float outv = (e & 1) ? (xo * cs + xe * sn) : (xe * cs - xo * sn);
              D.WSK[(size_t)cl.L * QKR + e] = f2b(outv);
            }
          } else {
            int r = t - 261;
            for (int e = threadIdx.x; e < KVL; e += THR) D.WSC[(size_t)r * KVL + e] = cl.ckv[(size_t)r * KVL + e];
            for (int e = threadIdx.x; e < QKR; e += THR) D.WSK[(size_t)r * QKR + e] = cl.krc[(size_t)r * QKR + e];
          }
        }
      }
      GSYNC();

      // ---- P2: attention (flash-mma engine) ----
      {
        int NT = cl.nchunk;
        while (true) {
          int t = next_work(D, 17, &sm_t);
          if (t >= NT) break;
          int p0 = t * cl.ch;
          int p1 = min(p0 + cl.ch, cl.L + 1);
          int lane = threadIdx.x & 31;
          int wid = threadIdx.x >> 5;
          int ht = wid & 1;   // acc h-tile
          int jo0 = (wid >> 1) & 3;   // base j-octet
          // stage qt (transposed q) once: qt[fe][h], 576 rows, 40 cols
          for (int e = threadIdx.x; e < 576 * 32; e += THR) {
            int fe = e >> 5, h = e & 31;
            bf16 v = f2b(0.f);
            if (fe < KVL) v = f2b(SC[SC_QABS + h * KVL + fe]);
            else v = f2b(SC[SC_QR + h * QKR + (fe - KVL)]);
            sm.a.qt[fe][h] = v;
          }
          if (threadIdx.x < 32) {
            sm.a.msum[threadIdx.x] = 0.f;
            sm.a.mrow[threadIdx.x] = -1e30f;
            sm.a.alp[threadIdx.x] = 1.f;
          }
          __syncthreads();
          float accf[4][4];
          #pragma unroll
          for (int b = 0; b < 4; b++)
            #pragma unroll
            for (int i = 0; i < 4; i++) accf[b][i] = 0.f;
          for (int ps = p0; ps < p1; ps += 32) {
            int np = min(32, p1 - ps);
            int npr = (np + 15) >> 4;   // pos 16-tiles to compute (1 or 2)
            for (int e = threadIdx.x; e < np * 64; e += THR) {
              int r = e >> 6;
              ((uint4*)&sm.a.cpos[r][0])[e & 63] = ((const uint4*)(D.WSC + (size_t)(ps + r) * KVL))[e & 63];
            }
            for (int e = threadIdx.x; e < np * 8; e += THR) {
              int r = e >> 3;
              ((uint4*)&sm.a.cpos[r][KVL])[e & 7] = ((const uint4*)(D.WSK + (size_t)(ps + r) * QKR))[e & 7];
            }
            if (np < 32) {
              for (int e = threadIdx.x; e < (32 - np) * 72; e += THR) {
                int r = np + (e / 72), q2 = e % 72;
                ((uint4*)&sm.a.cpos[r][0])[q2] = make_uint4(0, 0, 0, 0);
              }
            }
            __syncthreads();
            // ---- score pass: S^T tiles (2 pos x 4 h-oct) ----
            float d[4] = {0.f, 0.f, 0.f, 0.f};
            int spot = wid >> 2;      // 0..3? wid<8 -> 0,1
            int socth = wid & 3;
            if (wid < 8 && spot < npr) {
              for (int k = 0; k < 576; k += 16) {
                uint32_t a[4], b[2];
                ldmatrix_A(a, &sm.a.cpos[spot * 16][k], 520, lane);
                ldmatrix_Btrans(b, &sm.a.qt[k][socth * 8], 40, lane);
                mma_bf16(d, a, b);
              }
              // mask positions beyond np; scale
              #pragma unroll
              for (int i = 0; i < 4; i++) {
                int locpos = spot * 16 + (lane >> 2) + ((i & 2) ? 8 : 0);
                if (ps + locpos >= p1) d[i] = -1e30f;
                else d[i] *= MLASCALE;
              }
              // head-col max: even col va, odd col vb, reduce over xor 4,8,16
              #pragma unroll
              for (int part = 0; part < 2; part++) {
                float va = fmaxf(part ? d[2] : d[0], part ? d[3] : d[1]);
                va = fmaxf(va, __shfl_xor_sync(0xffffffffu, va, 4));
                va = fmaxf(va, __shfl_xor_sync(0xffffffffu, va, 8));
                va = fmaxf(va, __shfl_xor_sync(0xffffffffu, va, 16));
                int lc = (lane & 3) * 2 + part;
                sm.a.ctab[wid * 8 + lc] = va;
              }
            }
            __syncthreads();
            // merge over pos-tiles (warps {socth, socth+4}), new m, alpha
            if (threadIdx.x < 32) {
              int h = threadIdx.x;
              int so = (h >> 3) & 3;   // h-octet of this head
              float mold = sm.a.mrow[h];
              float mnew = mold;
              int lc = h & 7;
              mnew = fmaxf(mnew, fmaxf(sm.a.ctab[so * 8 + lc], sm.a.ctab[(so + 4) * 8 + lc]));
              sm.a.mrow[h] = mnew;
              sm.a.alp[h] = expf(mold - mnew);
            }
            __syncthreads();
            // P store (transposed scalar) + sum per head
            if (wid < 8 && spot < npr) {
              #pragma unroll
              for (int i = 0; i < 4; i++) {
                int lc = (lane & 3) * 2 + (i & 1);
                int h = socth * 8 + lc;
                int locpos = spot * 16 + (lane >> 2) + ((i & 2) ? 8 : 0);
                float pw = (d[i] <= -1e29f) ? 0.f : expf(d[i] - sm.a.mrow[h]);
                sm.a.psm[h * 40 + locpos] = f2b(pw);
                d[i] = pw;
              }
              #pragma unroll
              for (int part = 0; part < 2; part++) {
                float va = (part ? d[2] : d[0]) + (part ? d[3] : d[1]);
                va += __shfl_xor_sync(0xffffffffu, va, 4);
                va += __shfl_xor_sync(0xffffffffu, va, 8);
                va += __shfl_xor_sync(0xffffffffu, va, 16);
                int lc = (lane & 3) * 2 + part;
                sm.a.ctab[wid * 8 + lc] = va;
              }
            }
            __syncthreads();
            // l update
            if (threadIdx.x < 32) {
              int h = threadIdx.x;
              int so = (h >> 3) & 3;
              int lc = h & 7;
              float ssum = sm.a.ctab[so * 8 + lc] + sm.a.ctab[(so + 4) * 8 + lc];
              sm.a.msum[h] = sm.a.msum[h] * sm.a.alp[h] + ssum;
            }
            // rescale acc rows (4 batches of 4 tiles)
            #pragma unroll
            for (int p = 0; p < 4; p++)
              #pragma unroll
              for (int i = 0; i < 4; i++) {
                int hh = ht * 16 + (lane >> 2) + ((i & 2) ? 8 : 0);
                accf[p][i] *= sm.a.alp[hh];
              }
            __syncthreads();
            // ---- acc mma ----
            for (int kc = 0; kc < np; kc += 16) {
              uint32_t a[4], b[2];
              #pragma unroll
              for (int bb = 0; bb < 4; bb++) {
                ldmatrix_A(a, &sm.a.psm[(ht * 16) * 40 + kc], 40, lane);
                #pragma unroll
                for (int q2 = 0; q2 < 4; q2++) {
                  int joct = jo0 + 2 * bb + q2 * 16;
                  ldmatrix_Btrans(b, &sm.a.cpos[kc][joct * 8], 520, lane);
                  mma_bf16(accf[q2], a, b);
                }
              }
            }
          }
          // ---- write partial ----
          {
            float* base = SC + SC_PART + (size_t)t * NMH * 514;
            if (threadIdx.x < 32) {
              int h = threadIdx.x;
              base[h * 514 + 0] = sm.a.mrow[h];
              base[h * 514 + 1] = sm.a.msum[h];
            }
            int c0 = (lane & 3) * 2;
            #pragma unroll
            for (int q2 = 0; q2 < 16; q2++) {
              int hrow = ht * 16 + (lane >> 2);
              int bb = q2 >> 2;
              int joct = jo0 + 2 * bb + (q2 & 3) * 16;
              int col0 = joct * 8 + c0;
              base[hrow * 514 + 2 + col0 + 0] = accf[q2 & 3][0];
              base[hrow * 514 + 2 + col0 + 1] = accf[q2 & 3][1];
              base[(hrow + 8) * 514 + 2 + col0 + 0] = accf[q2 & 3][2];
              base[(hrow + 8) * 514 + 2 + col0 + 1] = accf[q2 & 3][3];
            }
          }
          __syncthreads();
        }
      }
      GSYNC();

      // ---- P3: combine partials ----
      {
        int NT = 128;
        while (true) {
          int t = next_work(D, 18, &sm_t);
          if (t >= NT) break;
          int h = t >> 2;
          int jq = (t & 3) * 128;
          if (threadIdx.x == 0) {
            float mstar = -1e30f;
            for (int c = 0; c < cl.nchunk; c++)
              mstar = fmaxf(mstar, SC[SC_PART + (size_t)c * NMH * 514 + h * 514]);
            float lstar = 0.f;
            for (int c = 0; c < cl.nchunk; c++)
              lstar += SC[SC_PART + (size_t)c * NMH * 514 + h * 514 + 1] *
                       expf(SC[SC_PART + (size_t)c * NMH * 514 + h * 514] - mstar);
            sm_red[0] = mstar;
            sm_red[1] = lstar;
          }
          __syncthreads();
          float mstar = sm_red[0], lstar = sm_red[1];
          for (int j = jq + threadIdx.x; j < jq + 128; j += THR) {
            float s = 0.f;
            for (int c = 0; c < cl.nchunk; c++) {
              float mc = SC[SC_PART + (size_t)c * NMH * 514 + h * 514];
              s += expf(mc - mstar) * SC[SC_PART + (size_t)c * NMH * 514 + h * 514 + 2 + j];
            }
            SC[SC_CTX + h * KVL + j] = s / lstar;
          }
          __syncthreads();
        }
      }
      GSYNC();

      // ---- P4: o-absorb GEMV (per head 128 outs, K=512) ----
      {
        int NT = 32;
        while (true) {
          int t = next_work(D, 19, &sm_t);
          if (t >= NT) break;
          int h = t;
          stage_raw(sm.g, SC + SC_CTX + (size_t)h * KVL, KVL);
          gemv_tile<4>(D, sm.g, D.wb[wbbase + 2], D.sb[wbbase + 2], NMH * 256, KVL,
                       h * 256 + 128, 1.f, SC + SC_O, nullptr, 0.f, 4 * (long long)(NMH * 256), nullptr,
                       h * 128);
        }
      }
      GSYNC();

      // ---- P5: o_proj + residual + MOEHG/U zero ----
      {
        stage_raw(sm.g, SC + SC_O, CC);
        int NT = 18 + 3;
        while (true) {
          int t = next_work(D, 20, &sm_t);
          if (t >= NT) break;
          if (t < 18) {
            int col = t * 128;
            gemv_tile<1>(D, sm.g, D.wb[wbbase + 3], D.sb[wbbase + 3], HID, CC, col, 1.f,
                         SC + SC_HACC, SC + SC_HACC, 0.f, (CC >> 7) * (long long)HID, SC + SC_MOACC);
          } else {
            int c0 = (t - 18) * 6144;
            for (int k = c0 + threadIdx.x; k < c0 + 6144; k += THR) {
              SC[SC_MOEHG + k] = 0.f;
              SC[SC_MOEHU + k] = 0.f;
            }
          }
        }
      }
      GSYNC();
      // ---- P6/P7: MoE ----
      moe_phase(D, sm, bi, cl, 21, 22, mnrm, sm_red, &sm_t, slot, tgt, G, 4);
      GSYNC();
    }
  }

  // ---- P8: write hidden out ----
  {
    int NT = 1;
    while (true) {
      int t = next_work(D, 23, &sm_t);
      if (t >= NT) break;
      for (int k = threadIdx.x; k < HID; k += THR) {
        float xb = b2f(f2b(SC[SC_HACC + k] + b2f(f2b(SC[SC_MOACC + k]))));
        D.XOUT[k] = f2b(xb);
      }
    }
  }
  GSYNC();
  // zero work counters for the next launch (all CTAs redundantly, no atomics)
  for (int k = threadIdx.x; k < 64; k += THR) D.BAR[BAR_WORK + k] = 0ULL;
}
"""

CPP_DECL = r"""
void mstep(const std::vector<int64_t>& a);
void setup(torch::Tensor wb, torch::Tensor sb, torch::Tensor ab, torch::Tensor sc, torch::Tensor bar,
           torch::Tensor wsc, torch::Tensor wsk, torch::Tensor xout, torch::Tensor cosc, torch::Tensor sinc,
           std::vector<int64_t> wb_off, std::vector<int64_t> sb_off, std::vector<int64_t> ab_off);
void step(torch::Tensor x_in, torch::Tensor S0, torch::Tensor cq0, torch::Tensor ck0, torch::Tensor cv0,
          torch::Tensor S1, torch::Tensor cq1, torch::Tensor ck1, torch::Tensor cv1,
          torch::Tensor S2, torch::Tensor cq2, torch::Tensor ck2, torch::Tensor cv2,
          torch::Tensor ckv, torch::Tensor krc,
          int64_t L, int64_t fresh, int64_t nchunk, int64_t ch, int64_t gen, int64_t dbg_stop);
"""

HOST_SRC = r"""
static Dev g_dev;
static int g_G = 0;
static int g_threads = THR;

void setup(torch::Tensor wb, torch::Tensor sb, torch::Tensor ab, torch::Tensor sc, torch::Tensor bar,
           torch::Tensor wsc, torch::Tensor wsk, torch::Tensor xout, torch::Tensor cosc, torch::Tensor sinc,
           std::vector<int64_t> wb_off, std::vector<int64_t> sb_off, std::vector<int64_t> ab_off) {
  TORCH_CHECK(wb_off.size() == 44 && sb_off.size() == 44 && ab_off.size() == 20, "bad offsets");
  Dev d;
  d.WB = (const uint8_t*)wb.data_ptr();
  d.SB = (const bf16*)sb.data_ptr();
  d.AB = (const bf16*)ab.data_ptr();
  d.SC = sc.data_ptr<float>();
  d.BAR = (ull*)bar.data_ptr();
  d.WSC = (bf16*)wsc.data_ptr();
  d.WSK = (bf16*)wsk.data_ptr();
  d.XOUT = (bf16*)xout.data_ptr();
  d.COSC = cosc.data_ptr<float>();
  d.SINC = sinc.data_ptr<float>();
  for (int i = 0; i < 44; i++) { d.wb[i] = wb_off[i]; d.sb[i] = sb_off[i]; }
  for (int i = 0; i < 20; i++) d.ab[i] = ab_off[i];
  g_dev = d;
  int occ = 0;
  cudaError_t e1 = cudaFuncSetAttribute(mega, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)sizeof(SmU) + 1024);
  cudaError_t e2 = cudaOccupancyMaxActiveBlocksPerMultiprocessor(&occ, (const void*)mega, g_threads, sizeof(SmU) + 1024);
  TORCH_CHECK(occ >= 1, "mega kernel not resident");
  cudaDeviceProp prop;
  cudaGetDeviceProperties(&prop, 0);
  g_G = prop.multiProcessorCount * occ;
}

void mstep(const std::vector<int64_t>& a) {
  Call cl;
  cl.x_in = (const bf16*)a[0];
  cl.S0 = (float*)a[1]; cl.cq0 = (bf16*)a[2]; cl.ck0 = (bf16*)a[3]; cl.cv0 = (bf16*)a[4];
  cl.S1 = (float*)a[5]; cl.cq1 = (bf16*)a[6]; cl.ck1 = (bf16*)a[7]; cl.cv1 = (bf16*)a[8];
  cl.S2 = (float*)a[9]; cl.cq2 = (bf16*)a[10]; cl.ck2 = (bf16*)a[11]; cl.cv2 = (bf16*)a[12];
  cl.ckv = (const bf16*)a[13]; cl.krc = (const bf16*)a[14];
  cl.L = (int)a[15]; cl.fresh = (int)a[16]; cl.nchunk = (int)a[17]; cl.ch = (int)a[18];
  cl.gen = a[19]; cl.dbg_stop = (int)a[20];
  Dev d = g_dev;
  void* args[] = {(void*)&d, (void*)&cl};
  cudaStream_t stream = at::cuda::getCurrentCUDAStream();
  cudaError_t err = cudaLaunchCooperativeKernel((const void*)mega, dim3(g_G), dim3(g_threads), args,
                                                sizeof(SmU) + 1024, stream);
  TORCH_CHECK(err == cudaSuccess, "mega launch failed: ", cudaGetErrorString(err));
}

void step(torch::Tensor x_in, torch::Tensor S0, torch::Tensor cq0, torch::Tensor ck0, torch::Tensor cv0,
          torch::Tensor S1, torch::Tensor cq1, torch::Tensor ck1, torch::Tensor cv1,
          torch::Tensor S2, torch::Tensor cq2, torch::Tensor ck2, torch::Tensor cv2,
          torch::Tensor ckv, torch::Tensor krc,
          int64_t L, int64_t fresh, int64_t nchunk, int64_t ch, int64_t gen, int64_t dbg_stop) {
  Call cl;
  cl.x_in = (const bf16*)x_in.data_ptr();
  cl.S0 = S0.data_ptr<float>(); cl.cq0 = (bf16*)cq0.data_ptr(); cl.ck0 = (bf16*)ck0.data_ptr(); cl.cv0 = (bf16*)cv0.data_ptr();
  cl.S1 = S1.data_ptr<float>(); cl.cq1 = (bf16*)cq1.data_ptr(); cl.ck1 = (bf16*)ck1.data_ptr(); cl.cv1 = (bf16*)cv1.data_ptr();
  cl.S2 = S2.data_ptr<float>(); cl.cq2 = (bf16*)cq2.data_ptr(); cl.ck2 = (bf16*)ck2.data_ptr(); cl.cv2 = (bf16*)cv2.data_ptr();
  cl.ckv = (const bf16*)ckv.data_ptr(); cl.krc = (const bf16*)krc.data_ptr();
  cl.L = (int)L; cl.fresh = (int)fresh; cl.nchunk = (int)nchunk; cl.ch = (int)ch; cl.gen = gen;
  cl.dbg_stop = (int)dbg_stop;
  Dev d = g_dev;
  void* args[] = {(void*)&d, (void*)&cl};
  cudaStream_t stream = at::cuda::getCurrentCUDAStream();
  cudaError_t err = cudaLaunchCooperativeKernel((const void*)mega, dim3(g_G), dim3(g_threads), args,
                                                sizeof(SmU) + 1024, stream);
  TORCH_CHECK(err == cudaSuccess, "mega launch failed: ", cudaGetErrorString(err));
}
"""


def build_cuda_source() -> str:
    return CUDA_SRC + "\n" + HOST_SRC


def extension():
    global _EXT
    if _EXT is None:
        _EXT = load_inline(
            name="kimi_mega",
            cpp_sources=CPP_DECL,
            cuda_sources=[build_cuda_source()],
            functions=["setup", "step", "mstep"],
            extra_cuda_cflags=["-O3", "-lineinfo"],
            verbose=False,
        )
    return _EXT

20260716_091712_kinetic-claude_kinetic-0715_02_kimi_linear_decode