"""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 #include #include #include #include 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 __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& 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 wb_off, std::vector sb_off, std::vector 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 wb_off, std::vector sb_off, std::vector 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(); 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(); d.SINC = sinc.data_ptr(); 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& 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(); cl.cq0 = (bf16*)cq0.data_ptr(); cl.ck0 = (bf16*)ck0.data_ptr(); cl.cv0 = (bf16*)cv0.data_ptr(); cl.S1 = S1.data_ptr(); cl.cq1 = (bf16*)cq1.data_ptr(); cl.ck1 = (bf16*)ck1.data_ptr(); cl.cv1 = (bf16*)cv1.data_ptr(); cl.S2 = S2.data_ptr(); 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