KernelBench mega · H100
Kimi-Linear Decode Kimi K3 (256k)
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.
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