"""Kimi-Linear W4A16 hybrid decode — single fused megakernel. One @triton.jit launch per step() fuses int4 dequant-GEMV, KDA, MLA, MoE, RMSNorm and residuals. Multi-CTA with software grid barriers between phases. """ from __future__ import annotations from dataclasses import dataclass, field import torch import torch.nn as nn import triton import triton.language as tl OP_TYPE = "kimi_linear_w4a16_decode" HARDWARE_REQUIRED = ["RTX_PRO_6000"] EPS = 1.0e-6 GROUP_SIZE = 128 # Workspace layout (bf16 elements) # See _ws_layout() for sizes. # f32 workspace: scores[L*H] + topk[2*NA] @dataclass(frozen=True) class Config: hidden: int = 2304 kda_heads: int = 32 kda_head_dim: int = 128 short_conv: int = 4 mla_heads: int = 32 kv_lora: int = 512 qk_nope: int = 128 qk_rope: int = 64 v_head: int = 128 rope_theta: float = 10000.0 n_experts: int = 64 n_active: int = 8 n_shared: int = 1 moe_inter: int = 1024 routed_scaling: float = 2.446 group: int = 128 pattern: tuple = ("K", "K", "K", "M") dtype: torch.dtype = field(default=torch.bfloat16) def build_config(shape: dict) -> Config: return Config(n_experts=int(shape.get("n_experts", 64))) def _pack_int4(w_q: torch.Tensor) -> torch.Tensor: lo = w_q[0::2] & 0xF hi = w_q[1::2] & 0xF return (lo | (hi << 4)).contiguous() def quantize(w_io: torch.Tensor, group: int = GROUP_SIZE): K, N = w_io.shape ng = K // group wg = w_io.view(ng, group, N).float() wmin = wg.min(dim=1, keepdim=True).values wmax = wg.max(dim=1, keepdim=True).values scales = (wmax - wmin).clamp_min(1e-8) / 15.0 zeros = (-wmin / scales).round().clamp(0, 15) w_q = ((wg / scales) + zeros).round().clamp(0, 15).to(torch.uint8).view(K, N) return _pack_int4(w_q), scales.squeeze(1).to(torch.bfloat16), zeros.squeeze(1).to(torch.bfloat16) class QuantLinear(nn.Module): 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) 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) 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) # --------------------------------------------------------------------------- # # Triton device helpers (no nested defs) # --------------------------------------------------------------------------- # @triton.jit def _barrier(bar_ptr, phase, n_blocks): tl.atomic_add(bar_ptr + phase, 1) while tl.load(bar_ptr + phase, volatile=True) < n_blocks: pass @triton.jit def _gemv_int4( x_ptr, wq_ptr, scales_ptr, zeros_ptr, y_ptr, K, N, pid, n_pids, BLOCK_N: tl.constexpr, GROUP: tl.constexpr, ): n_tiles = tl.cdiv(N, BLOCK_N) for t in range(pid, n_tiles, n_pids): n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < N acc = tl.zeros((BLOCK_N,), dtype=tl.float32) num_g = K // GROUP # preload x once per group loop — stream weights for g in range(num_g): k_base = g * GROUP s = tl.load(scales_ptr + g * N + offs_n, mask=mask_n, other=0.0).to(tl.float32) z = tl.load(zeros_ptr + g * N + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) packed = tl.load( wq_ptr + offs_pk[:, None] * N + offs_n[None, :], mask=mask_n[None, :], other=0, ) w0 = ((packed & 0xF).to(tl.float32) - z[None, :]) * s[None, :] w1 = (((packed >> 4) & 0xF).to(tl.float32) - z[None, :]) * s[None, :] x0 = tl.load(x_ptr + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(x_ptr + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) acc += tl.sum(x0[:, None] * w0, axis=0) acc += tl.sum(x1[:, None] * w1, axis=0) tl.store(y_ptr + offs_n, acc.to(tl.bfloat16), mask=mask_n) @triton.jit def _gemv_bf16(x_ptr, w_ptr, y_ptr, K, N, pid, n_pids, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr): n_tiles = tl.cdiv(N, BLOCK_N) for t in range(pid, n_tiles, n_pids): n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < N acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for k0 in range(0, K, BLOCK_K): offs_k = k0 + tl.arange(0, BLOCK_K) mask_k = offs_k < K x = tl.load(x_ptr + offs_k, mask=mask_k, other=0.0).to(tl.float32) w = tl.load( w_ptr + offs_k[:, None] * N + offs_n[None, :], mask=mask_k[:, None] & mask_n[None, :], other=0.0, ).to(tl.float32) acc += tl.sum(x[:, None] * w, axis=0) tl.store(y_ptr + offs_n, acc.to(tl.bfloat16), mask=mask_n) @triton.jit def _rmsnorm(x_ptr, w_ptr, y_ptr, D: tl.constexpr, pid): if pid == 0: # D=2304 is not a power of 2; accumulate in tiles of 256 acc = 0.0 for t in range(0, D, 256): offs = t + tl.arange(0, 256) mask = offs < D x = tl.load(x_ptr + offs, mask=mask, other=0.0).to(tl.float32) acc += tl.sum(x * x, axis=0) rstd = 1.0 / tl.sqrt(acc / float(D) + 1e-6) for t in range(0, D, 256): offs = t + tl.arange(0, 256) mask = offs < D x = tl.load(x_ptr + offs, mask=mask, other=0.0).to(tl.float32) ww = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32) tl.store(y_ptr + offs, (x * rstd * ww).to(tl.bfloat16), mask=mask) @triton.jit def _residual_add(x_ptr, delta_ptr, D, pid, n_pids, BLOCK_N: tl.constexpr): n_tiles = tl.cdiv(D, BLOCK_N) for t in range(pid, n_tiles, n_pids): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < D a = tl.load(x_ptr + offs, mask=mask, other=0.0).to(tl.float32) b = tl.load(delta_ptr + offs, mask=mask, other=0.0).to(tl.float32) tl.store(x_ptr + offs, (a + b).to(tl.bfloat16), mask=mask) @triton.jit def _zero(y_ptr, N, pid, n_pids, BLOCK_N: tl.constexpr): n_tiles = tl.cdiv(N, BLOCK_N) for t in range(pid, n_tiles, n_pids): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < N tl.store(y_ptr + offs, tl.zeros((BLOCK_N,), dtype=tl.bfloat16), mask=mask) @triton.jit def _short_conv(val_ptr, prev_ptr, conv_w_ptr, ch_idx, out_ptr, C, SC, pid, n_pids, BLOCK_N: tl.constexpr): n_tiles = tl.cdiv(C, BLOCK_N) for t in range(pid, n_tiles, n_pids): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < C acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for k in range(SC - 1): p = tl.load(prev_ptr + k * C + offs, mask=mask, other=0.0).to(tl.float32) w = tl.load(conv_w_ptr + ch_idx * C * SC + offs * SC + k, mask=mask, other=0.0).to(tl.float32) acc += p * w val = tl.load(val_ptr + offs, mask=mask, other=0.0).to(tl.float32) wlast = tl.load(conv_w_ptr + ch_idx * C * SC + offs * SC + (SC - 1), mask=mask, other=0.0).to(tl.float32) acc += val * wlast out = acc * tl.sigmoid(acc) tl.store(out_ptr + offs, out.to(tl.bfloat16), mask=mask) for k in range(SC - 2): nxt = tl.load(prev_ptr + (k + 1) * C + offs, mask=mask, other=0.0) tl.store(prev_ptr + k * C + offs, nxt, mask=mask) tl.store(prev_ptr + (SC - 2) * C + offs, val.to(tl.bfloat16), mask=mask) @triton.jit def _kda_delta(q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, S_ptr, out_ptr, H, Dk, kda_scale, pid, n_pids): for h in range(pid, H, n_pids): offs = tl.arange(0, Dk) g = tl.load(g_ptr + h * Dk + offs).to(tl.float32) # softplus sp = tl.where(g > 20.0, g, tl.log(1.0 + tl.exp(tl.where(g < -20.0, -20.0, g)))) decay = tl.exp(-sp) k = tl.load(k_ptr + h * Dk + offs).to(tl.float32) v = tl.load(v_ptr + h * Dk + offs).to(tl.float32) q = tl.load(q_ptr + h * Dk + offs).to(tl.float32) * kda_scale beta = tl.load(beta_ptr + h).to(tl.float32) beta = 1.0 / (1.0 + tl.exp(-beta)) pred = tl.zeros((Dk,), dtype=tl.float32) for i in range(Dk): di = tl.load(g_ptr + h * Dk + i).to(tl.float32) spi = tl.where(di > 20.0, di, tl.log(1.0 + tl.exp(tl.where(di < -20.0, -20.0, di)))) di_dec = tl.exp(-spi) ki = tl.load(k_ptr + h * Dk + i).to(tl.float32) row = tl.load(S_ptr + h * Dk * Dk + i * Dk + offs) row = row * di_dec tl.store(S_ptr + h * Dk * Dk + i * Dk + offs, row) pred += row * ki delta = v - pred o = tl.zeros((Dk,), dtype=tl.float32) for i in range(Dk): ki = tl.load(k_ptr + h * Dk + i).to(tl.float32) qi = tl.load(q_ptr + h * Dk + i).to(tl.float32) * kda_scale row = tl.load(S_ptr + h * Dk * Dk + i * Dk + offs) row = row + beta * ki * delta tl.store(S_ptr + h * Dk * Dk + i * Dk + offs, row) o += row * qi tl.store(out_ptr + h * Dk + offs, o.to(tl.bfloat16)) @triton.jit def _silu_mul(g_ptr, u_ptr, out_ptr, N, pid, n_pids, BLOCK_N: tl.constexpr): n_tiles = tl.cdiv(N, BLOCK_N) for t in range(pid, n_tiles, n_pids): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < N g = tl.load(g_ptr + offs, mask=mask, other=0.0).to(tl.float32) u = tl.load(u_ptr + offs, mask=mask, other=0.0).to(tl.float32) h = g * tl.sigmoid(g) * u tl.store(out_ptr + offs, h.to(tl.bfloat16), mask=mask) @triton.jit def _axpy(y_ptr, x_ptr, scale, N, pid, n_pids, BLOCK_N: tl.constexpr): n_tiles = tl.cdiv(N, BLOCK_N) for t in range(pid, n_tiles, n_pids): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < N a = tl.load(y_ptr + offs, mask=mask, other=0.0).to(tl.float32) b = tl.load(x_ptr + offs, mask=mask, other=0.0).to(tl.float32) tl.store(y_ptr + offs, (a + scale * b).to(tl.bfloat16), mask=mask) @triton.jit def _router_topk(x_ptr, router_w, fws_ptr, D, E, NA, routed_scaling, pid): """logits = x @ W (W is D x E), softmax, top-NA. Writes w[0:NA], idx[NA:2*NA] as f32 in fws.""" if pid == 0: # compute logits offs_e = tl.arange(0, 64) # E=64 acc = tl.zeros((64,), dtype=tl.float32) for k0 in range(0, D, 64): offs_k = k0 + tl.arange(0, 64) mask_k = offs_k < D x = tl.load(x_ptr + offs_k, mask=mask_k, other=0.0).to(tl.float32) w = tl.load( router_w + offs_k[:, None] * E + offs_e[None, :], mask=mask_k[:, None], other=0.0, ).to(tl.float32) acc += tl.sum(x[:, None] * w, axis=0) mx = tl.max(acc, axis=0) ex = tl.exp(acc - mx) sm = ex / tl.sum(ex, axis=0) mask = tl.full((64,), 1.0, tl.float32) for j in range(NA): vals = tl.where(mask > 0.5, sm, float("-inf")) best_v = tl.max(vals, axis=0) idxs = tl.arange(0, 64).to(tl.float32) is_best = vals >= (best_v - 1e-12) cand = tl.where(is_best, idxs, 999.0) best_i = tl.min(cand, axis=0) tl.store(fws_ptr + j, best_v) tl.store(fws_ptr + NA + j, best_i) mask = tl.where(idxs == best_i, 0.0, mask) wsum = 0.0 for j in range(NA): wsum += tl.load(fws_ptr + j) for j in range(NA): w = tl.load(fws_ptr + j) / (wsum + 1e-9) * routed_scaling tl.store(fws_ptr + j, w) # --------------------------------------------------------------------------- # # Megakernel — fully inlined phase sequence, no nested defs # --------------------------------------------------------------------------- # @triton.jit def mega_kernel( hidden_ptr, fws_ptr, bws_ptr, bar_ptr, n_blocks, mla_pos, rope_theta, kda_scale, mla_scale, routed_scaling, # layer0 KDA weights q0_wq, q0_s, q0_z, k0_wq, k0_s, k0_z, v0_wq, v0_s, v0_z, g0_wq, g0_s, g0_z, o0_wq, o0_s, o0_z, beta0_w, conv0_w, an0, mn0, r0_w, gate0_wq, gate0_s, gate0_z, up0_wq, up0_s, up0_z, down0_wq, down0_s, down0_z, sgate0_wq, sgate0_s, sgate0_z, sup0_wq, sup0_s, sup0_z, sdown0_wq, sdown0_s, sdown0_z, S0, cq0, ck0, cv0, # layer1 q1_wq, q1_s, q1_z, k1_wq, k1_s, k1_z, v1_wq, v1_s, v1_z, g1_wq, g1_s, g1_z, o1_wq, o1_s, o1_z, beta1_w, conv1_w, an1, mn1, r1_w, gate1_wq, gate1_s, gate1_z, up1_wq, up1_s, up1_z, down1_wq, down1_s, down1_z, sgate1_wq, sgate1_s, sgate1_z, sup1_wq, sup1_s, sup1_z, sdown1_wq, sdown1_s, sdown1_z, S1, cq1, ck1, cv1, # layer2 q2_wq, q2_s, q2_z, k2_wq, k2_s, k2_z, v2_wq, v2_s, v2_z, g2_wq, g2_s, g2_z, o2_wq, o2_s, o2_z, beta2_w, conv2_w, an2, mn2, r2_w, gate2_wq, gate2_s, gate2_z, up2_wq, up2_s, up2_z, down2_wq, down2_s, down2_z, sgate2_wq, sgate2_s, sgate2_z, sup2_wq, sup2_s, sup2_z, sdown2_wq, sdown2_s, sdown2_z, S2, cq2, ck2, cv2, # layer3 MLA qm_wq, qm_s, qm_z, kva_wq, kva_s, kva_z, kvb_wq, kvb_s, kvb_z, om_wq, om_s, om_z, an3, mn3, r3_w, gate3_wq, gate3_s, gate3_z, up3_wq, up3_s, up3_z, down3_wq, down3_s, down3_z, sgate3_wq, sgate3_s, sgate3_z, sup3_wq, sup3_s, sup3_z, sdown3_wq, sdown3_s, sdown3_z, ckv_src, krope_src, ckv_dst, krope_dst, D: tl.constexpr, H: tl.constexpr, Dk: tl.constexpr, C: tl.constexpr, SC: tl.constexpr, KV: tl.constexpr, QN: tl.constexpr, QR: tl.constexpr, VH: tl.constexpr, E: tl.constexpr, NA: tl.constexpr, NS: tl.constexpr, MI: tl.constexpr, GROUP: tl.constexpr, BLOCK_N: tl.constexpr, ): pid = tl.program_id(0) ph = 0 # bf16 workspace offsets OFF_XN = 0 OFF_Q = D OFF_K = D + C OFF_V = D + 2 * C OFF_G = D + 3 * C OFF_BETA = D + 4 * C OFF_ATTN = D + 4 * C + H # MoE workspace: NA experts of gate/up/down intermediates OFF_MG = OFF_ATTN + C # [NA, MI] gate or silu*up OFF_MU = OFF_MG + NA * MI # [NA, MI] up OFF_MD = OFF_MU + NA * MI # [NA, D] down outs OFF_MO = OFF_MD + NA * D # [D] moe accum OFF_TMP = OFF_MO + D OFF_QFULL = OFF_TMP + D OFF_KV = OFF_QFULL + H * (QN + QR) OFF_QABS = OFF_KV + KV + QR # ---- KDA layer ---- _rmsnorm(hidden_ptr, an0, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # independent projections (no inter-barrier) _gemv_int4(bws_ptr + OFF_XN, q0_wq, q0_s, q0_z, bws_ptr + OFF_Q, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, k0_wq, k0_s, k0_z, bws_ptr + OFF_K, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, v0_wq, v0_s, v0_z, bws_ptr + OFF_V, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, g0_wq, g0_s, g0_z, bws_ptr + OFF_G, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_bf16(bws_ptr + OFF_XN, beta0_w, bws_ptr + OFF_BETA, D, H, pid, n_blocks, BLOCK_N, 64) _barrier(bar_ptr, ph, n_blocks); ph += 1 _short_conv(bws_ptr + OFF_Q, cq0, conv0_w, 0, bws_ptr + OFF_Q, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_K, ck0, conv0_w, 1, bws_ptr + OFF_K, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_V, cv0, conv0_w, 2, bws_ptr + OFF_V, C, SC, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _kda_delta(bws_ptr + OFF_Q, bws_ptr + OFF_K, bws_ptr + OFF_V, bws_ptr + OFF_G, bws_ptr + OFF_BETA, S0, bws_ptr + OFF_ATTN, H, Dk, kda_scale, pid, n_blocks) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_ATTN, o0_wq, o0_s, o0_z, bws_ptr + OFF_TMP, C, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_TMP, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _rmsnorm(hidden_ptr, mn0, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 _router_topk(bws_ptr + OFF_XN, r0_w, fws_ptr, D, E, NA, routed_scaling, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel gate+up for all active experts: unit = j * n_tiles + t n_tiles_mi = tl.cdiv(MI, BLOCK_N) for unit in range(pid, NA * n_tiles_mi, n_blocks): j = unit // n_tiles_mi t = unit % n_tiles_mi e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < MI # gate acc_g = tl.zeros((BLOCK_N,), dtype=tl.float32) acc_u = tl.zeros((BLOCK_N,), dtype=tl.float32) wqg = gate0_wq + e_id * (D // 2) * MI sg = gate0_s + e_id * (D // GROUP) * MI zg = gate0_z + e_id * (D // GROUP) * MI wqu = up0_wq + e_id * (D // 2) * MI su = up0_s + e_id * (D // GROUP) * MI zu = up0_z + e_id * (D // GROUP) * MI for gg in range(D // GROUP): k_base = gg * GROUP s_g = tl.load(sg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_g = tl.load(zg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) s_u = tl.load(su + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_u = tl.load(zu + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pg = tl.load(wqg + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) pu = tl.load(wqu + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0g = ((pg & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w1g = (((pg >> 4) & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w0u = ((pu & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] w1u = (((pu >> 4) & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] acc_g += tl.sum(x0[:, None] * w0g, axis=0) + tl.sum(x1[:, None] * w1g, axis=0) acc_u += tl.sum(x0[:, None] * w0u, axis=0) + tl.sum(x1[:, None] * w1u, axis=0) # silu(gate)*up hh = acc_g * tl.sigmoid(acc_g) * acc_u tl.store(bws_ptr + OFF_MG + j * MI + offs_n, hh.to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel down for all experts n_tiles_d = tl.cdiv(D, BLOCK_N) for unit in range(pid, NA * n_tiles_d, n_blocks): j = unit // n_tiles_d t = unit % n_tiles_d e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) wqd = down0_wq + e_id * (MI // 2) * D sd = down0_s + e_id * (MI // GROUP) * D zd = down0_z + e_id * (MI // GROUP) * D for gg in range(MI // GROUP): k_base = gg * GROUP scv = tl.load(sd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) zv = tl.load(zd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pk = tl.load(wqd + offs_pk[:, None] * D + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0 = ((pk & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] w1 = (((pk >> 4) & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] acc += tl.sum(x0[:, None] * w0, axis=0) + tl.sum(x1[:, None] * w1, axis=0) wj = tl.load(fws_ptr + j) tl.store(bws_ptr + OFF_MD + j * D + offs_n, (acc * wj).to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # sum over experts into OFF_MO (single parallel pass) n_tiles_sum = tl.cdiv(D, BLOCK_N) for t in range(pid, n_tiles_sum, n_blocks): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for j in range(NA): acc += tl.load(bws_ptr + OFF_MD + j * D + offs, mask=mask, other=0.0).to(tl.float32) tl.store(bws_ptr + OFF_MO + offs, acc.to(tl.bfloat16), mask=mask) _barrier(bar_ptr, ph, n_blocks); ph += 1 # shared experts (NS=1) for ss in range(NS): _gemv_int4(bws_ptr + OFF_XN, sgate0_wq + ss * (D // 2) * MI, sgate0_s + ss * (D // GROUP) * MI, sgate0_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MG, D, MI, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, sup0_wq + ss * (D // 2) * MI, sup0_s + ss * (D // GROUP) * MI, sup0_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MU, D, MI, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _silu_mul(bws_ptr + OFF_MG, bws_ptr + OFF_MU, bws_ptr + OFF_MG, MI, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_MG, sdown0_wq + ss * (MI // 2) * D, sdown0_s + ss * (MI // GROUP) * D, sdown0_z + ss * (MI // GROUP) * D, bws_ptr + OFF_TMP, MI, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _axpy(bws_ptr + OFF_MO, bws_ptr + OFF_TMP, 1.0, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_MO, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 # ---- KDA layer ---- _rmsnorm(hidden_ptr, an1, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # independent projections (no inter-barrier) _gemv_int4(bws_ptr + OFF_XN, q1_wq, q1_s, q1_z, bws_ptr + OFF_Q, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, k1_wq, k1_s, k1_z, bws_ptr + OFF_K, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, v1_wq, v1_s, v1_z, bws_ptr + OFF_V, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, g1_wq, g1_s, g1_z, bws_ptr + OFF_G, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_bf16(bws_ptr + OFF_XN, beta1_w, bws_ptr + OFF_BETA, D, H, pid, n_blocks, BLOCK_N, 64) _barrier(bar_ptr, ph, n_blocks); ph += 1 _short_conv(bws_ptr + OFF_Q, cq1, conv1_w, 0, bws_ptr + OFF_Q, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_K, ck1, conv1_w, 1, bws_ptr + OFF_K, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_V, cv1, conv1_w, 2, bws_ptr + OFF_V, C, SC, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _kda_delta(bws_ptr + OFF_Q, bws_ptr + OFF_K, bws_ptr + OFF_V, bws_ptr + OFF_G, bws_ptr + OFF_BETA, S1, bws_ptr + OFF_ATTN, H, Dk, kda_scale, pid, n_blocks) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_ATTN, o1_wq, o1_s, o1_z, bws_ptr + OFF_TMP, C, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_TMP, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _rmsnorm(hidden_ptr, mn1, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 _router_topk(bws_ptr + OFF_XN, r1_w, fws_ptr, D, E, NA, routed_scaling, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel gate+up for all active experts: unit = j * n_tiles + t n_tiles_mi = tl.cdiv(MI, BLOCK_N) for unit in range(pid, NA * n_tiles_mi, n_blocks): j = unit // n_tiles_mi t = unit % n_tiles_mi e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < MI # gate acc_g = tl.zeros((BLOCK_N,), dtype=tl.float32) acc_u = tl.zeros((BLOCK_N,), dtype=tl.float32) wqg = gate1_wq + e_id * (D // 2) * MI sg = gate1_s + e_id * (D // GROUP) * MI zg = gate1_z + e_id * (D // GROUP) * MI wqu = up1_wq + e_id * (D // 2) * MI su = up1_s + e_id * (D // GROUP) * MI zu = up1_z + e_id * (D // GROUP) * MI for gg in range(D // GROUP): k_base = gg * GROUP s_g = tl.load(sg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_g = tl.load(zg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) s_u = tl.load(su + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_u = tl.load(zu + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pg = tl.load(wqg + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) pu = tl.load(wqu + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0g = ((pg & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w1g = (((pg >> 4) & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w0u = ((pu & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] w1u = (((pu >> 4) & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] acc_g += tl.sum(x0[:, None] * w0g, axis=0) + tl.sum(x1[:, None] * w1g, axis=0) acc_u += tl.sum(x0[:, None] * w0u, axis=0) + tl.sum(x1[:, None] * w1u, axis=0) # silu(gate)*up hh = acc_g * tl.sigmoid(acc_g) * acc_u tl.store(bws_ptr + OFF_MG + j * MI + offs_n, hh.to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel down for all experts n_tiles_d = tl.cdiv(D, BLOCK_N) for unit in range(pid, NA * n_tiles_d, n_blocks): j = unit // n_tiles_d t = unit % n_tiles_d e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) wqd = down1_wq + e_id * (MI // 2) * D sd = down1_s + e_id * (MI // GROUP) * D zd = down1_z + e_id * (MI // GROUP) * D for gg in range(MI // GROUP): k_base = gg * GROUP scv = tl.load(sd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) zv = tl.load(zd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pk = tl.load(wqd + offs_pk[:, None] * D + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0 = ((pk & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] w1 = (((pk >> 4) & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] acc += tl.sum(x0[:, None] * w0, axis=0) + tl.sum(x1[:, None] * w1, axis=0) wj = tl.load(fws_ptr + j) tl.store(bws_ptr + OFF_MD + j * D + offs_n, (acc * wj).to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # sum over experts into OFF_MO (single parallel pass) n_tiles_sum = tl.cdiv(D, BLOCK_N) for t in range(pid, n_tiles_sum, n_blocks): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for j in range(NA): acc += tl.load(bws_ptr + OFF_MD + j * D + offs, mask=mask, other=0.0).to(tl.float32) tl.store(bws_ptr + OFF_MO + offs, acc.to(tl.bfloat16), mask=mask) _barrier(bar_ptr, ph, n_blocks); ph += 1 # shared experts (NS=1) for ss in range(NS): _gemv_int4(bws_ptr + OFF_XN, sgate1_wq + ss * (D // 2) * MI, sgate1_s + ss * (D // GROUP) * MI, sgate1_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MG, D, MI, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, sup1_wq + ss * (D // 2) * MI, sup1_s + ss * (D // GROUP) * MI, sup1_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MU, D, MI, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _silu_mul(bws_ptr + OFF_MG, bws_ptr + OFF_MU, bws_ptr + OFF_MG, MI, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_MG, sdown1_wq + ss * (MI // 2) * D, sdown1_s + ss * (MI // GROUP) * D, sdown1_z + ss * (MI // GROUP) * D, bws_ptr + OFF_TMP, MI, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _axpy(bws_ptr + OFF_MO, bws_ptr + OFF_TMP, 1.0, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_MO, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 # ---- KDA layer ---- _rmsnorm(hidden_ptr, an2, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # independent projections (no inter-barrier) _gemv_int4(bws_ptr + OFF_XN, q2_wq, q2_s, q2_z, bws_ptr + OFF_Q, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, k2_wq, k2_s, k2_z, bws_ptr + OFF_K, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, v2_wq, v2_s, v2_z, bws_ptr + OFF_V, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, g2_wq, g2_s, g2_z, bws_ptr + OFF_G, D, C, pid, n_blocks, BLOCK_N, GROUP) _gemv_bf16(bws_ptr + OFF_XN, beta2_w, bws_ptr + OFF_BETA, D, H, pid, n_blocks, BLOCK_N, 64) _barrier(bar_ptr, ph, n_blocks); ph += 1 _short_conv(bws_ptr + OFF_Q, cq2, conv2_w, 0, bws_ptr + OFF_Q, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_K, ck2, conv2_w, 1, bws_ptr + OFF_K, C, SC, pid, n_blocks, BLOCK_N) _short_conv(bws_ptr + OFF_V, cv2, conv2_w, 2, bws_ptr + OFF_V, C, SC, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _kda_delta(bws_ptr + OFF_Q, bws_ptr + OFF_K, bws_ptr + OFF_V, bws_ptr + OFF_G, bws_ptr + OFF_BETA, S2, bws_ptr + OFF_ATTN, H, Dk, kda_scale, pid, n_blocks) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_ATTN, o2_wq, o2_s, o2_z, bws_ptr + OFF_TMP, C, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_TMP, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _rmsnorm(hidden_ptr, mn2, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 _router_topk(bws_ptr + OFF_XN, r2_w, fws_ptr, D, E, NA, routed_scaling, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel gate+up for all active experts: unit = j * n_tiles + t n_tiles_mi = tl.cdiv(MI, BLOCK_N) for unit in range(pid, NA * n_tiles_mi, n_blocks): j = unit // n_tiles_mi t = unit % n_tiles_mi e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < MI # gate acc_g = tl.zeros((BLOCK_N,), dtype=tl.float32) acc_u = tl.zeros((BLOCK_N,), dtype=tl.float32) wqg = gate2_wq + e_id * (D // 2) * MI sg = gate2_s + e_id * (D // GROUP) * MI zg = gate2_z + e_id * (D // GROUP) * MI wqu = up2_wq + e_id * (D // 2) * MI su = up2_s + e_id * (D // GROUP) * MI zu = up2_z + e_id * (D // GROUP) * MI for gg in range(D // GROUP): k_base = gg * GROUP s_g = tl.load(sg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_g = tl.load(zg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) s_u = tl.load(su + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_u = tl.load(zu + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pg = tl.load(wqg + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) pu = tl.load(wqu + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0g = ((pg & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w1g = (((pg >> 4) & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w0u = ((pu & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] w1u = (((pu >> 4) & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] acc_g += tl.sum(x0[:, None] * w0g, axis=0) + tl.sum(x1[:, None] * w1g, axis=0) acc_u += tl.sum(x0[:, None] * w0u, axis=0) + tl.sum(x1[:, None] * w1u, axis=0) # silu(gate)*up hh = acc_g * tl.sigmoid(acc_g) * acc_u tl.store(bws_ptr + OFF_MG + j * MI + offs_n, hh.to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel down for all experts n_tiles_d = tl.cdiv(D, BLOCK_N) for unit in range(pid, NA * n_tiles_d, n_blocks): j = unit // n_tiles_d t = unit % n_tiles_d e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) wqd = down2_wq + e_id * (MI // 2) * D sd = down2_s + e_id * (MI // GROUP) * D zd = down2_z + e_id * (MI // GROUP) * D for gg in range(MI // GROUP): k_base = gg * GROUP scv = tl.load(sd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) zv = tl.load(zd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pk = tl.load(wqd + offs_pk[:, None] * D + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0 = ((pk & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] w1 = (((pk >> 4) & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] acc += tl.sum(x0[:, None] * w0, axis=0) + tl.sum(x1[:, None] * w1, axis=0) wj = tl.load(fws_ptr + j) tl.store(bws_ptr + OFF_MD + j * D + offs_n, (acc * wj).to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # sum over experts into OFF_MO (single parallel pass) n_tiles_sum = tl.cdiv(D, BLOCK_N) for t in range(pid, n_tiles_sum, n_blocks): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for j in range(NA): acc += tl.load(bws_ptr + OFF_MD + j * D + offs, mask=mask, other=0.0).to(tl.float32) tl.store(bws_ptr + OFF_MO + offs, acc.to(tl.bfloat16), mask=mask) _barrier(bar_ptr, ph, n_blocks); ph += 1 # shared experts (NS=1) for ss in range(NS): _gemv_int4(bws_ptr + OFF_XN, sgate2_wq + ss * (D // 2) * MI, sgate2_s + ss * (D // GROUP) * MI, sgate2_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MG, D, MI, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, sup2_wq + ss * (D // 2) * MI, sup2_s + ss * (D // GROUP) * MI, sup2_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MU, D, MI, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _silu_mul(bws_ptr + OFF_MG, bws_ptr + OFF_MU, bws_ptr + OFF_MG, MI, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_MG, sdown2_wq + ss * (MI // 2) * D, sdown2_s + ss * (MI // GROUP) * D, sdown2_z + ss * (MI // GROUP) * D, bws_ptr + OFF_TMP, MI, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _axpy(bws_ptr + OFF_MO, bws_ptr + OFF_TMP, 1.0, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_MO, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 # ---- MLA layer ---- pos = mla_pos _rmsnorm(hidden_ptr, an3, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 qdim = H * (QN + QR) _gemv_int4(bws_ptr + OFF_XN, qm_wq, qm_s, qm_z, bws_ptr + OFF_QFULL, D, qdim, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, kva_wq, kva_s, kva_z, bws_ptr + OFF_KV, D, KV + QR, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 # RoPE if pid == 0: for hh in range(H): base = bws_ptr + OFF_QFULL + hh * (QN + QR) + QN for ii in range(QR // 2): inv = 1.0 / tl.exp((2.0 * ii / float(QR)) * tl.log(rope_theta)) ang = pos * inv cosv = tl.cos(ang) sinv = tl.sin(ang) even = tl.load(base + 2 * ii).to(tl.float32) odd = tl.load(base + 2 * ii + 1).to(tl.float32) tl.store(base + 2 * ii, (even * cosv - odd * sinv).to(tl.bfloat16)) tl.store(base + 2 * ii + 1, (odd * cosv + even * sinv).to(tl.bfloat16)) basek = bws_ptr + OFF_KV + KV for ii in range(QR // 2): inv = 1.0 / tl.exp((2.0 * ii / float(QR)) * tl.log(rope_theta)) ang = pos * inv cosv = tl.cos(ang) sinv = tl.sin(ang) even = tl.load(basek + 2 * ii).to(tl.float32) odd = tl.load(basek + 2 * ii + 1).to(tl.float32) tl.store(basek + 2 * ii, (even * cosv - odd * sinv).to(tl.bfloat16)) tl.store(basek + 2 * ii + 1, (odd * cosv + even * sinv).to(tl.bfloat16)) _barrier(bar_ptr, ph, n_blocks); ph += 1 # copy cache + append for ll in range(pid, pos, n_blocks): for r0 in range(0, KV, 128): offs = r0 + tl.arange(0, 128) mask = offs < KV tl.store(ckv_dst + ll * KV + offs, tl.load(ckv_src + ll * KV + offs, mask=mask, other=0.0), mask=mask) for r0 in range(0, QR, 64): offs = r0 + tl.arange(0, 64) mask = offs < QR tl.store(krope_dst + ll * QR + offs, tl.load(krope_src + ll * QR + offs, mask=mask, other=0.0), mask=mask) if pid == 0: for ii in range(KV): tl.store(ckv_dst + pos * KV + ii, tl.load(bws_ptr + OFF_KV + ii)) for ii in range(QR): tl.store(krope_dst + pos * QR + ii, tl.load(bws_ptr + OFF_KV + KV + ii)) _barrier(bar_ptr, ph, n_blocks); ph += 1 L = pos + 1 OUT_COLS = H * (QN + VH) # q_abs[h,r] = sum_d W_k[r,h,d] * q_nope[h,d] for h in range(pid, H, n_blocks): for r0 in range(0, KV, 64): offs_r = r0 + tl.arange(0, 64) mask_r = offs_r < KV acc = tl.zeros((64,), dtype=tl.float32) for d in range(QN): col = h * (QN + VH) + d qd = tl.load(bws_ptr + OFF_QFULL + h * (QN + QR) + d).to(tl.float32) packed = tl.load(kvb_wq + (offs_r // 2) * OUT_COLS + col, mask=mask_r, other=0) lo = (packed & 0xF).to(tl.float32) hi = ((packed >> 4) & 0xF).to(tl.float32) is_even = (offs_r % 2) == 0 wq = tl.where(is_even, lo, hi) gg = offs_r // GROUP scv = tl.load(kvb_s + gg * OUT_COLS + col, mask=mask_r, other=0.0).to(tl.float32) zv = tl.load(kvb_z + gg * OUT_COLS + col, mask=mask_r, other=0.0).to(tl.float32) acc += (wq - zv) * scv * qd tl.store(bws_ptr + OFF_QABS + h * KV + offs_r, acc.to(tl.bfloat16), mask=mask_r) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Scores: partition (head, L-tile) across all CTAs n_lt = tl.cdiv(L, 16) n_units = H * n_lt for unit in range(pid, n_units, n_blocks): h = unit // n_lt l0 = (unit % n_lt) * 16 offs_l = l0 + tl.arange(0, 16) mask_l = offs_l < L acc = tl.zeros((16,), dtype=tl.float32) for r0 in range(0, KV, 64): offs_r = r0 + tl.arange(0, 64) mask_r = offs_r < KV qa = tl.load(bws_ptr + OFF_QABS + h * KV + offs_r, mask=mask_r, other=0.0).to(tl.float32) c = tl.load( ckv_dst + offs_l[:, None] * KV + offs_r[None, :], mask=mask_l[:, None] & mask_r[None, :], other=0.0, ).to(tl.float32) acc += tl.sum(c * qa[None, :], axis=1) qr = tl.load(bws_ptr + OFF_QFULL + h * (QN + QR) + QN + tl.arange(0, QR)).to(tl.float32) kr = tl.load( krope_dst + offs_l[:, None] * QR + tl.arange(0, QR)[None, :], mask=mask_l[:, None], other=0.0, ).to(tl.float32) acc += tl.sum(kr * qr[None, :], axis=1) tl.store(fws_ptr + offs_l * H + h, acc * mla_scale, mask=mask_l) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Softmax stats per head: store mx at fws[L*H + h], se at fws[L*H + H + h] stats = fws_ptr + L * H for h in range(pid, H, n_blocks): mx = float("-inf") for l0 in range(0, L, 32): offs_l = l0 + tl.arange(0, 32) mask_l = offs_l < L sc = tl.load(fws_ptr + offs_l * H + h, mask=mask_l, other=float("-inf")) sc_m = tl.where(mask_l, sc, float("-inf")) mx = tl.maximum(mx, tl.max(sc_m, axis=0)) se = 0.0 for l0 in range(0, L, 32): offs_l = l0 + tl.arange(0, 32) mask_l = offs_l < L sc = tl.load(fws_ptr + offs_l * H + h, mask=mask_l, other=0.0) se += tl.sum(tl.where(mask_l, tl.exp(sc - mx), 0.0), axis=0) tl.store(stats + h, mx) tl.store(stats + H + h, se) _barrier(bar_ptr, ph, n_blocks); ph += 1 # attn_c: partition (head, r-tile) across all CTAs n_rt = KV // 64 for unit in range(pid, H * n_rt, n_blocks): h = unit // n_rt r0 = (unit % n_rt) * 64 mx = tl.load(stats + h) se = tl.load(stats + H + h) offs_r = r0 + tl.arange(0, 64) mask_r = offs_r < KV acc = tl.zeros((64,), dtype=tl.float32) for l0 in range(0, L, 16): offs_l = l0 + tl.arange(0, 16) mask_l = offs_l < L sc = tl.load(fws_ptr + offs_l * H + h, mask=mask_l, other=0.0) p = tl.where(mask_l, tl.exp(sc - mx) / se, 0.0) c = tl.load( ckv_dst + offs_l[:, None] * KV + offs_r[None, :], mask=mask_l[:, None] & mask_r[None, :], other=0.0, ).to(tl.float32) acc += tl.sum(p[:, None] * c, axis=0) tl.store(bws_ptr + OFF_QABS + h * KV + offs_r, acc.to(tl.bfloat16), mask=mask_r) _barrier(bar_ptr, ph, n_blocks); ph += 1 # o[h,d] = sum_r W_v[r,h,d] * attn_c[h,r] for h in range(pid, H, n_blocks): for d in range(VH): col = h * (QN + VH) + QN + d accv = 0.0 for r0 in range(0, KV, 64): offs_r = r0 + tl.arange(0, 64) mask_r = offs_r < KV packed = tl.load(kvb_wq + (offs_r // 2) * OUT_COLS + col, mask=mask_r, other=0) lo = (packed & 0xF).to(tl.float32) hi = ((packed >> 4) & 0xF).to(tl.float32) is_even = (offs_r % 2) == 0 wq = tl.where(is_even, lo, hi) gg = offs_r // GROUP scv = tl.load(kvb_s + gg * OUT_COLS + col, mask=mask_r, other=0.0).to(tl.float32) zv = tl.load(kvb_z + gg * OUT_COLS + col, mask=mask_r, other=0.0).to(tl.float32) w = (wq - zv) * scv a = tl.load(bws_ptr + OFF_QABS + h * KV + offs_r, mask=mask_r, other=0.0).to(tl.float32) accv += tl.sum(w * a, axis=0) tl.store(bws_ptr + OFF_ATTN + h * VH + d, accv.to(tl.bfloat16)) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_ATTN, om_wq, om_s, om_z, bws_ptr + OFF_TMP, H * VH, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_TMP, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _rmsnorm(hidden_ptr, mn3, bws_ptr + OFF_XN, D, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 _router_topk(bws_ptr + OFF_XN, r3_w, fws_ptr, D, E, NA, routed_scaling, pid) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel gate+up for all active experts: unit = j * n_tiles + t n_tiles_mi = tl.cdiv(MI, BLOCK_N) for unit in range(pid, NA * n_tiles_mi, n_blocks): j = unit // n_tiles_mi t = unit % n_tiles_mi e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < MI # gate acc_g = tl.zeros((BLOCK_N,), dtype=tl.float32) acc_u = tl.zeros((BLOCK_N,), dtype=tl.float32) wqg = gate3_wq + e_id * (D // 2) * MI sg = gate3_s + e_id * (D // GROUP) * MI zg = gate3_z + e_id * (D // GROUP) * MI wqu = up3_wq + e_id * (D // 2) * MI su = up3_s + e_id * (D // GROUP) * MI zu = up3_z + e_id * (D // GROUP) * MI for gg in range(D // GROUP): k_base = gg * GROUP s_g = tl.load(sg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_g = tl.load(zg + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) s_u = tl.load(su + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) z_u = tl.load(zu + gg * MI + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pg = tl.load(wqg + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) pu = tl.load(wqu + offs_pk[:, None] * MI + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_XN + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0g = ((pg & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w1g = (((pg >> 4) & 0xF).to(tl.float32) - z_g[None, :]) * s_g[None, :] w0u = ((pu & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] w1u = (((pu >> 4) & 0xF).to(tl.float32) - z_u[None, :]) * s_u[None, :] acc_g += tl.sum(x0[:, None] * w0g, axis=0) + tl.sum(x1[:, None] * w1g, axis=0) acc_u += tl.sum(x0[:, None] * w0u, axis=0) + tl.sum(x1[:, None] * w1u, axis=0) # silu(gate)*up hh = acc_g * tl.sigmoid(acc_g) * acc_u tl.store(bws_ptr + OFF_MG + j * MI + offs_n, hh.to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # Parallel down for all experts n_tiles_d = tl.cdiv(D, BLOCK_N) for unit in range(pid, NA * n_tiles_d, n_blocks): j = unit // n_tiles_d t = unit % n_tiles_d e_id = tl.load(fws_ptr + NA + j).to(tl.int32) n0 = t * BLOCK_N offs_n = n0 + tl.arange(0, BLOCK_N) mask_n = offs_n < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) wqd = down3_wq + e_id * (MI // 2) * D sd = down3_s + e_id * (MI // GROUP) * D zd = down3_z + e_id * (MI // GROUP) * D for gg in range(MI // GROUP): k_base = gg * GROUP scv = tl.load(sd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) zv = tl.load(zd + gg * D + offs_n, mask=mask_n, other=0.0).to(tl.float32) offs_pk = (k_base // 2) + tl.arange(0, GROUP // 2) pk = tl.load(wqd + offs_pk[:, None] * D + offs_n[None, :], mask=mask_n[None, :], other=0) x0 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2)).to(tl.float32) x1 = tl.load(bws_ptr + OFF_MG + j * MI + k_base + 2 * tl.arange(0, GROUP // 2) + 1).to(tl.float32) w0 = ((pk & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] w1 = (((pk >> 4) & 0xF).to(tl.float32) - zv[None, :]) * scv[None, :] acc += tl.sum(x0[:, None] * w0, axis=0) + tl.sum(x1[:, None] * w1, axis=0) wj = tl.load(fws_ptr + j) tl.store(bws_ptr + OFF_MD + j * D + offs_n, (acc * wj).to(tl.bfloat16), mask=mask_n) _barrier(bar_ptr, ph, n_blocks); ph += 1 # sum over experts into OFF_MO (single parallel pass) n_tiles_sum = tl.cdiv(D, BLOCK_N) for t in range(pid, n_tiles_sum, n_blocks): offs = t * BLOCK_N + tl.arange(0, BLOCK_N) mask = offs < D acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for j in range(NA): acc += tl.load(bws_ptr + OFF_MD + j * D + offs, mask=mask, other=0.0).to(tl.float32) tl.store(bws_ptr + OFF_MO + offs, acc.to(tl.bfloat16), mask=mask) _barrier(bar_ptr, ph, n_blocks); ph += 1 # shared experts (NS=1) for ss in range(NS): _gemv_int4(bws_ptr + OFF_XN, sgate3_wq + ss * (D // 2) * MI, sgate3_s + ss * (D // GROUP) * MI, sgate3_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MG, D, MI, pid, n_blocks, BLOCK_N, GROUP) _gemv_int4(bws_ptr + OFF_XN, sup3_wq + ss * (D // 2) * MI, sup3_s + ss * (D // GROUP) * MI, sup3_z + ss * (D // GROUP) * MI, bws_ptr + OFF_MU, D, MI, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _silu_mul(bws_ptr + OFF_MG, bws_ptr + OFF_MU, bws_ptr + OFF_MG, MI, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _gemv_int4(bws_ptr + OFF_MG, sdown3_wq + ss * (MI // 2) * D, sdown3_s + ss * (MI // GROUP) * D, sdown3_z + ss * (MI // GROUP) * D, bws_ptr + OFF_TMP, MI, D, pid, n_blocks, BLOCK_N, GROUP) _barrier(bar_ptr, ph, n_blocks); ph += 1 _axpy(bws_ptr + OFF_MO, bws_ptr + OFF_TMP, 1.0, D, pid, n_blocks, BLOCK_N) _barrier(bar_ptr, ph, n_blocks); ph += 1 _residual_add(hidden_ptr, bws_ptr + OFF_MO, D, pid, n_blocks, BLOCK_N) # --------------------------------------------------------------------------- # # Host Model # --------------------------------------------------------------------------- # class Model(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern) self.reset_parameters() self._ws_b = None self._ws_f = None self._bar = None self._ckv_buf = None self._krope_buf = None self._n_blocks = 114 self._prepared = False def reset_parameters(self): g = torch.Generator(device="cpu").manual_seed(1234) for mod in self.modules(): if isinstance(mod, (QuantLinear, QuantExperts)): mod.init_random(g) elif isinstance(mod, nn.Linear): nn.init.normal_(mod.weight, 0.0, 0.02, generator=g) elif isinstance(mod, KDA): nn.init.normal_(mod.conv_w, 0.0, 0.1, generator=g) def _prepare(self, device, max_ctx: int): cfg = self.cfg D, H, C = cfg.hidden, cfg.kda_heads, cfg.kda_heads * cfg.kda_head_dim MI = cfg.moe_inter NA = cfg.n_active MI = cfg.moe_inter bsz = ( D + 4 * C + H + C + NA * MI + NA * MI + NA * D # gate/up/down expert bufs + D + D # moe out + tmp + H * (cfg.qk_nope + cfg.qk_rope) + cfg.kv_lora + cfg.qk_rope + H * cfg.kv_lora + 2048 ) self._ws_b = torch.empty(bsz, device=device, dtype=torch.bfloat16) fsz = (max_ctx + 64) * H + 256 self._ws_f = torch.empty(fsz, device=device, dtype=torch.float32) # many barriers: ~200 phases self._bar = torch.zeros(2048, device=device, dtype=torch.int32) self._ckv_buf = torch.empty(max_ctx + 64, cfg.kv_lora, device=device, dtype=cfg.dtype) self._krope_buf = torch.empty(max_ctx + 64, cfg.qk_rope, device=device, dtype=cfg.dtype) self._max_ctx = max_ctx self._prepared = True def step(self, hidden, state): cfg = self.cfg device = hidden.device mla_i = cfg.pattern.index("M") pos = state[mla_i]["c_kv"].shape[0] need = pos + 16 if not self._prepared or self._ckv_buf is None or self._ckv_buf.shape[0] < need: self._prepare(device, max(need, pos + 256)) # Cache transposed tiny bf16 weights (router / beta) across steps if getattr(self, "_weight_id", None) is not id(self.blocks[0].moe.router.weight): self._router_T = [blk.moe.router.weight.detach().T.contiguous() for blk in self.blocks] self._beta_T = [ blk.attn.beta_proj.weight.detach().T.contiguous() if isinstance(blk.attn, KDA) else None for blk in self.blocks ] self._weight_id = id(self.blocks[0].moe.router.weight) router_T = self._router_T beta_T = self._beta_T for i, kind in enumerate(cfg.pattern): if kind == "K" and state[i]["S"].dtype != torch.float32: state[i]["S"] = state[i]["S"].float() self._bar.zero_() D = cfg.hidden H = cfg.kda_heads Dk = cfg.kda_head_dim C = H * Dk n_blocks = self._n_blocks BLOCK_N = 64 b0, b1, b2, b3 = self.blocks a0, a1, a2, am = b0.attn, b1.attn, b2.attn, b3.attn m0, m1, m2, m3 = b0.moe, b1.moe, b2.moe, b3.moe st0, st1, st2 = state[0], state[1], state[2] ckv_src = state[mla_i]["c_kv"].contiguous() krope_src = state[mla_i]["k_rope"].contiguous() def qp(ql): return ql.w_q, ql.scales, ql.zeros def ep(qe): return qe.w_q, qe.scales, qe.zeros h = hidden.contiguous() mega_kernel[(n_blocks,)]( h, self._ws_f, self._ws_b, self._bar, n_blocks, pos, float(cfg.rope_theta), float(Dk ** -0.5), float((cfg.qk_nope + cfg.qk_rope) ** -0.5), float(cfg.routed_scaling), *qp(a0.q_proj), *qp(a0.k_proj), *qp(a0.v_proj), *qp(a0.g_proj), *qp(a0.o_proj), beta_T[0], a0.conv_w, b0.attn_norm, b0.moe_norm, router_T[0], *ep(m0.gate), *ep(m0.up), *ep(m0.down), *ep(m0.s_gate), *ep(m0.s_up), *ep(m0.s_down), st0["S"], st0["cq"], st0["ck"], st0["cv"], *qp(a1.q_proj), *qp(a1.k_proj), *qp(a1.v_proj), *qp(a1.g_proj), *qp(a1.o_proj), beta_T[1], a1.conv_w, b1.attn_norm, b1.moe_norm, router_T[1], *ep(m1.gate), *ep(m1.up), *ep(m1.down), *ep(m1.s_gate), *ep(m1.s_up), *ep(m1.s_down), st1["S"], st1["cq"], st1["ck"], st1["cv"], *qp(a2.q_proj), *qp(a2.k_proj), *qp(a2.v_proj), *qp(a2.g_proj), *qp(a2.o_proj), beta_T[2], a2.conv_w, b2.attn_norm, b2.moe_norm, router_T[2], *ep(m2.gate), *ep(m2.up), *ep(m2.down), *ep(m2.s_gate), *ep(m2.s_up), *ep(m2.s_down), st2["S"], st2["cq"], st2["ck"], st2["cv"], *qp(am.q_proj), *qp(am.kv_a), *qp(am.kv_b), *qp(am.o_proj), b3.attn_norm, b3.moe_norm, router_T[3], *ep(m3.gate), *ep(m3.up), *ep(m3.down), *ep(m3.s_gate), *ep(m3.s_up), *ep(m3.s_down), ckv_src, krope_src, self._ckv_buf, self._krope_buf, D=D, H=H, Dk=Dk, C=C, SC=cfg.short_conv, KV=cfg.kv_lora, QN=cfg.qk_nope, QR=cfg.qk_rope, VH=cfg.v_head, E=cfg.n_experts, NA=cfg.n_active, NS=cfg.n_shared, MI=cfg.moe_inter, GROUP=GROUP_SIZE, BLOCK_N=BLOCK_N, ) state[mla_i]["c_kv"] = self._ckv_buf[: pos + 1] state[mla_i]["k_rope"] = self._krope_buf[: pos + 1] return h, state 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