"""Kimi-Linear W4A16 hybrid decode — single fused cooperative megakernel. Exposes the same interface as reference.py: Model(cfg); step(hidden, state). The entire per-token forward (3 KDA + 1 MLA layers, each + MoE, int4 dequant-GEMV, RMSNorm, residuals, both state updates) runs in ONE CUDA cooperative kernel launch. The int4 weights are streamed once — dequant is fused into the GEMV. """ from __future__ import annotations import os from dataclasses import dataclass, field import torch import torch.nn as nn EPS = 1.0e-6 GROUP_SIZE = 128 @dataclass(frozen=True) class Config: hidden: int = 2304 kda_heads: int = 32 kda_head_dim: int = 128 short_conv: int = 4 mla_heads: int = 32 kv_lora: int = 512 qk_nope: int = 128 qk_rope: int = 64 v_head: int = 128 rope_theta: float = 10000.0 n_experts: int = 64 n_active: int = 8 n_shared: int = 1 moe_inter: int = 1024 routed_scaling: float = 2.446 group: int = 128 pattern: tuple = ("K", "K", "K", "M") dtype: torch.dtype = field(default=torch.bfloat16) def build_config(shape: dict) -> Config: return Config(n_experts=int(shape.get("n_experts", 64))) # --------------------------------------------------------------------------- # # module definitions (identical buffer/parameter names to reference.py) # --------------------------------------------------------------------------- # 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) class QuantLinear(nn.Module): def __init__(self, in_f: int, out_f: int, group: int = GROUP_SIZE): super().__init__() assert in_f % group == 0 and in_f % 2 == 0 self.in_f, self.out_f, self.group = in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen: torch.Generator, std: float = 0.02) -> None: w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q.copy_(wq) self.scales.copy_(s) self.zeros.copy_(z) def weight_bf(self) -> torch.Tensor: return dequant(self.w_q, self.scales, self.zeros, self.in_f, self.group) def forward(self, x: torch.Tensor) -> torch.Tensor: return (x.float() @ self.weight_bf().float()).to(torch.bfloat16) class QuantExperts(nn.Module): def __init__(self, n: int, in_f: int, out_f: int, group: int = GROUP_SIZE): super().__init__() self.n, self.in_f, self.out_f, self.group = n, in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(n, in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen: torch.Generator, std: float = 0.02) -> None: for e in range(self.n): w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q[e].copy_(wq) self.scales[e].copy_(s) self.zeros[e].copy_(z) 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 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 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 class MoE(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg d, m, E = cfg.hidden, cfg.moe_inter, cfg.n_experts self.router = nn.Linear(d, E, bias=False, dtype=cfg.dtype) self.gate = QuantExperts(E, d, m, cfg.group) self.up = QuantExperts(E, d, m, cfg.group) self.down = QuantExperts(E, m, d, cfg.group) self.s_gate = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_up = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_down = QuantExperts(cfg.n_shared, m, d, cfg.group) class Block(nn.Module): def __init__(self, cfg: Config, kind: str): super().__init__() self.kind = kind self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.attn = KDA(cfg) if kind == "K" else MLA(cfg) self.moe = MoE(cfg) 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._ext = None self._ptrs = None self._sig_ptr = None self._first = 0 self._ws = None # ------------------------------------------------------------------ # # workspace / pointer-table construction # ------------------------------------------------------------------ # def _prepare(self, state): from dev import kernel_build cfg = self.cfg L = state[3]["c_kv"].shape[0] capacity = L + 512 # headroom for autoregressive growth dev = torch.device("cuda:0") ws = {} ws["h_fp32"] = torch.empty(cfg.hidden, device=dev, dtype=torch.float32) C = cfg.kda_heads * cfg.kda_head_dim ws["q"] = torch.empty(C, device=dev, dtype=torch.float32) ws["k"] = torch.empty(C, device=dev, dtype=torch.float32) ws["v"] = torch.empty(C, device=dev, dtype=torch.float32) ws["g"] = torch.empty(C, device=dev, dtype=torch.float32) ws["beta"] = torch.empty(cfg.kda_heads, device=dev, dtype=torch.float32) ws["mla_q"] = torch.empty(cfg.mla_heads * (cfg.qk_nope + cfg.qk_rope), device=dev, dtype=torch.float32) ws["kv"] = torch.empty(cfg.kv_lora + cfg.qk_rope, device=dev, dtype=torch.float32) ws["qabs"] = torch.empty(cfg.mla_heads * cfg.kv_lora, device=dev, dtype=torch.float32) ws["qrope"] = torch.empty(cfg.mla_heads * cfg.qk_rope, device=dev, dtype=torch.float32) ws["scores"] = torch.empty(capacity * cfg.mla_heads, device=dev, dtype=torch.float32) ws["ocomp"] = torch.empty(cfg.mla_heads * cfg.kv_lora, device=dev, dtype=torch.float32) ws["norm2"] = torch.empty(cfg.hidden, device=dev, dtype=torch.float32) ws["ridx"] = torch.empty(cfg.n_active, device=dev, dtype=torch.int32) ws["rw"] = torch.empty(cfg.n_active, device=dev, dtype=torch.float32) nchunks = (capacity + 63) // 64 ws["chmax"] = torch.empty(nchunks * cfg.mla_heads, device=dev, dtype=torch.float32) ws["chsum"] = torch.empty(nchunks * cfg.mla_heads, device=dev, dtype=torch.float32) ws["new_cq"] = torch.empty(3 * C, device=dev, dtype=torch.bfloat16) ws["new_ck"] = torch.empty(3 * C, device=dev, dtype=torch.bfloat16) ws["new_cv"] = torch.empty(3 * C, device=dev, dtype=torch.bfloat16) ws["qc_scr"] = torch.empty(3 * C, device=dev, dtype=torch.float32) ws["kc_scr"] = torch.empty(3 * C, device=dev, dtype=torch.float32) ws["vc_scr"] = torch.empty(3 * C, device=dev, dtype=torch.float32) ws["pred_scr"] = torch.empty(C, device=dev, dtype=torch.float32) ws["g_scr"] = torch.empty(9 * cfg.moe_inter, device=dev, dtype=torch.float32) ws["u_scr"] = torch.empty(9 * cfg.moe_inter, device=dev, dtype=torch.float32) ws["gmax"] = torch.empty(cfg.mla_heads, device=dev, dtype=torch.float32) ws["gsum"] = torch.empty(cfg.mla_heads, device=dev, dtype=torch.float32) ws["h_scr"] = torch.empty(9 * cfg.moe_inter, device=dev, dtype=torch.float32) ws["c_kv_cap"] = torch.empty(capacity * cfg.kv_lora, device=dev, dtype=torch.bfloat16) ws["k_rope_cap"] = torch.empty(capacity * cfg.qk_rope, device=dev, dtype=torch.bfloat16) self._ws = ws P = kernel_build # norms ptrs = [] for b in range(4): ptrs.append(self.blocks[b].attn_norm.data_ptr()) for b in range(4): ptrs.append(self.blocks[b].moe_norm.data_ptr()) # KDA blocks 0..2 for b in range(3): a = self.blocks[b].attn st = state[b] ptrs += [ a.q_proj.w_q.data_ptr(), a.q_proj.scales.data_ptr(), a.q_proj.zeros.data_ptr(), a.k_proj.w_q.data_ptr(), a.k_proj.scales.data_ptr(), a.k_proj.zeros.data_ptr(), a.v_proj.w_q.data_ptr(), a.v_proj.scales.data_ptr(), a.v_proj.zeros.data_ptr(), a.g_proj.w_q.data_ptr(), a.g_proj.scales.data_ptr(), a.g_proj.zeros.data_ptr(), a.beta_proj.weight.data_ptr(), a.conv_w.data_ptr(), a.o_proj.w_q.data_ptr(), a.o_proj.scales.data_ptr(), a.o_proj.zeros.data_ptr(), st["S"].data_ptr(), st["cq"].data_ptr(), st["ck"].data_ptr(), st["cv"].data_ptr(), ] # MLA block 3 b = 3 a = self.blocks[b].attn st = state[3] ptrs += [ a.q_proj.w_q.data_ptr(), a.q_proj.scales.data_ptr(), a.q_proj.zeros.data_ptr(), a.kv_a.w_q.data_ptr(), a.kv_a.scales.data_ptr(), a.kv_a.zeros.data_ptr(), a.kv_b.w_q.data_ptr(), a.kv_b.scales.data_ptr(), a.kv_b.zeros.data_ptr(), a.o_proj.w_q.data_ptr(), a.o_proj.scales.data_ptr(), a.o_proj.zeros.data_ptr(), ws["c_kv_cap"].data_ptr(), ws["k_rope_cap"].data_ptr(), st["c_kv"].data_ptr(), st["k_rope"].data_ptr(), ] # MoE blocks for b in range(4): moe = self.blocks[b].moe ptrs += [ moe.router.weight.data_ptr(), moe.gate.w_q.data_ptr(), moe.gate.scales.data_ptr(), moe.gate.zeros.data_ptr(), moe.up.w_q.data_ptr(), moe.up.scales.data_ptr(), moe.up.zeros.data_ptr(), moe.down.w_q.data_ptr(), moe.down.scales.data_ptr(), moe.down.zeros.data_ptr(), moe.s_gate.w_q.data_ptr(), moe.s_gate.scales.data_ptr(), moe.s_gate.zeros.data_ptr(), moe.s_up.w_q.data_ptr(), moe.s_up.scales.data_ptr(), moe.s_up.zeros.data_ptr(), moe.s_down.w_q.data_ptr(), moe.s_down.scales.data_ptr(), moe.s_down.zeros.data_ptr(), ] # workspace ptrs += [ ws["h_fp32"].data_ptr(), ws["q"].data_ptr(), ws["k"].data_ptr(), ws["v"].data_ptr(), ws["g"].data_ptr(), ws["beta"].data_ptr(), ws["mla_q"].data_ptr(), ws["kv"].data_ptr(), ws["qabs"].data_ptr(), ws["qrope"].data_ptr(), ws["scores"].data_ptr(), ws["ocomp"].data_ptr(), ws["norm2"].data_ptr(), ws["ridx"].data_ptr(), ws["rw"].data_ptr(), ws["chmax"].data_ptr(), ws["chsum"].data_ptr(), ws["new_cq"].data_ptr(), ws["new_ck"].data_ptr(), ws["new_cv"].data_ptr(), ws["qc_scr"].data_ptr(), ws["kc_scr"].data_ptr(), ws["vc_scr"].data_ptr(), ws["pred_scr"].data_ptr(), ws["g_scr"].data_ptr(), ws["u_scr"].data_ptr(), ws["gmax"].data_ptr(), ws["gsum"].data_ptr(), ws["h_scr"].data_ptr(), ] self._ptrs = torch.tensor(ptrs, dtype=torch.int64, device=dev) self._sig_ptr = ws["c_kv_cap"].data_ptr() self._first = 1 # ------------------------------------------------------------------ # # step -- the whole per-token forward is ONE cooperative CUDA kernel, # compiled via torch.utils.cpp_extension.load_inline (see dev/kernel_build.py). # ------------------------------------------------------------------ # def step(self, hidden, state): if self._ext is None: from dev import kernel_build self._ext = kernel_build.build_kernel() if self._ptrs is None or state[3]["c_kv"].data_ptr() != self._sig_ptr: self._prepare(state) L = state[3]["c_kv"].shape[0] self._ext.launch_kimi(self._ptrs, L, self._first, hidden) self._first = 0 cap = self._ws["c_kv_cap"] krc = self._ws["k_rope_cap"] state[3]["c_kv"] = cap[: (L + 1) * self.cfg.kv_lora].view(L + 1, self.cfg.kv_lora) state[3]["k_rope"] = krc[: (L + 1) * self.cfg.qk_rope].view(L + 1, self.cfg.qk_rope) return hidden, state # --------------------------------------------------------------------------- # # state / inputs (same as reference) # --------------------------------------------------------------------------- # def init_state(cfg: Config, context_len: int, seed: int) -> list: dev = torch.device("cuda:0") g = torch.Generator(device=dev).manual_seed(seed) H, Dk = cfg.kda_heads, cfg.kda_head_dim C = H * Dk state = [] for kind in cfg.pattern: if kind == "K": state.append({ "S": torch.randn(H, Dk, Dk, device=dev, generator=g) * 0.05, "cq": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "ck": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "cv": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, }) else: state.append({ "c_kv": torch.randn(context_len, cfg.kv_lora, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "k_rope": torch.randn(context_len, cfg.qk_rope, device=dev, generator=g, dtype=cfg.dtype) * 0.1, }) return state def init_token(cfg: Config, seed: int) -> torch.Tensor: dev = torch.device("cuda:0") g = torch.Generator(device=dev).manual_seed(seed + 1) return torch.randn(cfg.hidden, device=dev, generator=g, dtype=cfg.dtype) * 0.25