KernelBench mega · H100
Kimi-Linear Decode Kimi K3 (1M)
manually audited: clean
H100 kinetic-0715[1m] cell (4.0063x geomean speedup versus the eager reference over ctx 2048/8192/16384). Clean, genuine single-launch raw CUDA megakernel: unlike the RTX PRO 6000 sidecar runs, this solution.py is fully self-contained (1347 lines, no local module import) with the CUDA source embedded and compiled via load_inline. The timed step() path calls _ext.mega_step exactly once, which issues one megakernel<<<nCTA, 256, 40960B smem>>> launch; all 23 stages (three KDA blocks: fused int4 GEMVs, short conv, gated-delta S update, o_proj+residual, router/top-8 MoE gate/up/down; one MLA block: q/kv_a projections with in-kernel RoPE, latent cache append plus in-kernel copy-forward of the old cache rows, flash-decoding chunked online-softmax attention with a combine phase, kv_b value projection, o_proj, MoE) are separated by an atomics-based grid-wide barrier inside the same launch. No CUDA graph, no torch.compile, no per-token multi-kernel loop, no forbidden library, no cached or constant output: every launch reads the live hidden/state pointers and writes hout, all KDA S/conv windows, and the appended MLA cache into freshly allocated output tensors. The debug stage-stop parameter is hardcoded to 999 (run all stages) in the timed path, and KIMI_FORCE_EAGER only selects the slow eager oracle, which the official grading did not use (benchmark shows 2.79-3.39 ms/tok versus 11.5-13.3 ms baseline, and latency grows with context length as real live-cache attention should).
Kernel source (redacted)
"""Kimi-Linear W4A16 hybrid decode unit -- single-launch fused megakernel.
One decode step = ONE persistent CUDA kernel launch. All int4 weights are
unpacked + dequantized in-register inside fused GEMVs (never materialized);
KDA gated-delta recurrence, MLA absorbed latent attention (flash-decoding
partials + combine), 64-expert int4 MoE with router/top-8, short conv,
RMSNorms and residuals are all fused into stage pipelines separated by
grid-wide barriers inside the same launch. An eager PyTorch path is kept for
debugging (KIMI_FORCE_EAGER=1) but the timed step() path is one kernel launch.
"""
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
@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)))
# --------------------------------------------------------------------------- #
# W4A16 quantization (identical layout 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)
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 weight_bf(self, e: int) -> torch.Tensor:
return dequant(self.w_q[e], self.scales[e], self.zeros[e], self.in_f, self.group)
# --------------------------------------------------------------------------- #
# helpers
# --------------------------------------------------------------------------- #
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)
# --------------------------------------------------------------------------- #
# layers (eager fallback / debug oracle)
# --------------------------------------------------------------------------- #
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 step(self, x, st):
H, Dk = self.cfg.kda_heads, self.cfg.kda_head_dim
q, k, v = self.q_proj(x), self.k_proj(x), 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.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.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 step(self, x, st):
cfg = self.cfg
H = cfg.mla_heads
pos = st["c_kv"].shape[0]
q = 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.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.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.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))
# --------------------------------------------------------------------------- #
# Megakernel CUDA source (compiled once at import).
# --------------------------------------------------------------------------- #
_CUDA_SRC = r"""// Kimi-Linear W4A16 decode megakernel: one launch per token.
// Stages separated by grid-wide sync inside a single persistent kernel.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#include <cstdint>
using bf16 = __nv_bfloat16;
#define NTHR 256
#define NWARP 8
#define EPS 1e-6f
#define ROUTED_SCALE 2.446f
#define QSCALE_KDA 0.08838834764831845f // 128^-0.5
#define QSCALE_MLA 0.07216878364870323f // 192^-0.5
#define HID 2304
#define C4096 4096
#define NHEADS 32
#define DK 128
#define QDIM 6144
#define KVA 576
#define LORA 512
#define ROPE 64
#define QHB 192
#define NOPE 128
#define VH 128
#define KVBOUT 8192
#define NEXP 64
#define NACT 8
#define INTER 1024
#define LCCHUNK 256
// dynamic smem map
#define SM_XS 0 // float2[2048] = 16384
#define SM_RED 16384 // float[8][128] = 4096
#define SM_SPRE 20480 // bf16[32][128] = 8192
#define SM_ZPRE 28672 // bf16[32][128] = 8192
#define SM_YOUT 36864 // float[128] = 512
#define SM_YOUT2 37376 // float[128] = 512
#define SM_ROUTE 37888 // float rw[9]+pad, int ridx[9]
#define SM_LSM 38400 // float[64]
#define SMEM_SIZE 40960
struct QMat {
const uint8_t* w; // (R, N) packed int4
const bf16* s; // (G, N)
const bf16* z; // (G, N)
};
struct LayerW {
const bf16* anorm; // (2304,)
const bf16* mnorm; // (2304,)
QMat q, k, v, g, o; // KDA
const bf16* betaw; // (2304, 32) transposed
const bf16* convw; // (3, 4096, 4)
QMat mq, mka, mkb, mo; // MLA
QMat gate, up, down, sgate, sup, sdown; // experts: (E, R, N)
const float* wnr; // (2304, 64) fp32, moe_norm folded router
};
struct DynState {
const float* S_in[3]; float* S_out[3];
const bf16 *cq_in[3], *ck_in[3], *cv_in[3];
bf16 *cq_out[3], *ck_out[3], *cv_out[3];
const bf16 *ckv_in, *kr_in;
bf16 *ckv_out, *kr_out;
};
struct Params {
LayerW lyr[4];
bf16 *qs, *ks, *vs, *gs, *beta_raw; // qs sized 6144 (MLA reuses)
bf16* o_att;
bf16* h1;
float* acc0;
float* acc1;
float* sumsq;
float* logits;
float* qlat;
float* olat;
float* part;
const float* invf;
float* m_scratch;
unsigned* bar;
const bf16* hin;
bf16* hout;
DynState st;
int L;
int nchunks;
int nCTA;
int stop; // debug: return from kernel after stage with this id (999 = run all)
};
// ---------------------------------------------------------------------------
__device__ __forceinline__ float b2f(bf16 v) { return __bfloat162float(v); }
__device__ __forceinline__ bf16 f2b(float v) { return __float2bfloat16(v); }
__device__ __forceinline__ float silu(float x) { return x / (1.f + __expf(-x)); }
__device__ __forceinline__ float sigmoidf_(float x) { return 1.f / (1.f + __expf(-x)); }
__device__ __forceinline__ float nsoftplus(float x) { return x > 20.f ? x : log1pf(__expf(x)); }
__device__ void gbar(unsigned* bar) {
__syncthreads();
if (threadIdx.x == 0) {
unsigned* cnt = bar;
unsigned* gen = bar + 1;
__threadfence();
unsigned g = *((volatile unsigned*)gen);
unsigned arrived = atomicAdd(cnt, 1u) + 1u;
if (arrived == gridDim.x) {
atomicExch(cnt, 0u);
__threadfence();
atomicExch(gen, g + 1u);
} else {
while (*((volatile unsigned*)gen) == g) {
__nanosleep(64);
}
}
__threadfence();
}
__syncthreads();
}
// residual stream element j at input of layer l
__device__ __forceinline__ float resid_elem(const Params& P, int l, int j) {
if (l == 0) return b2f(P.hin[j]);
const float* accp = ((l - 1) & 1) ? P.acc1 : P.acc0;
return b2f(f2b(b2f(P.h1[j]) + b2f(f2b(accp[j]))));
}
// build rmsnorm'd xs in smem fp32; which: 0=attn norm over residual, 2=moe norm over h1
__device__ void build_xs_norm(const Params& P, int l, int which, float2* xs2, float* red) {
int t = threadIdx.x;
float ss = 0.f;
if (which == 0) {
for (int j = t; j < HID; j += NTHR) {
float rj = resid_elem(P, l, j);
ss += rj * rj;
}
#pragma unroll
for (int off = 16; off > 0; off >>= 1) ss += __shfl_xor_sync(0xffffffffu, ss, off);
__syncthreads();
if ((t & 31) == 0) red[t >> 5] = ss;
__syncthreads();
if (t == 0) {
float v = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; ++w) v += red[w];
red[8] = v;
}
__syncthreads();
}
float r = (which == 0) ? rsqrtf(red[8] / HID + EPS) : rsqrtf(*P.sumsq / HID + EPS);
const bf16* nw = (which == 2) ? P.lyr[l].mnorm : P.lyr[l].anorm;
for (int j = t; j < HID; j += NTHR) {
float x = (which == 2) ? b2f(P.h1[j]) : resid_elem(P, l, j);
float val = r * x * b2f(nw[j]);
((float*)xs2)[j] = b2f(f2b(val));
}
__syncthreads();
}
__device__ void build_xs_from_bf16(const bf16* src, int K, float2* xs2) {
int t = threadIdx.x;
for (int j = t; j < K; j += NTHR) ((float*)xs2)[j] = b2f(src[j]);
__syncthreads();
}
// ---------------------------------------------------------------------------
// Quant GEMV over 128 columns; result y[128] fp32 in yout (smem).
__device__ void qgemv128(const QMat m, int R, int N, int G, int c0,
const float2* xs2, bf16* spre, bf16* zpre,
float* red, float* yout) {
int t = threadIdx.x;
int lane = t & 31, warp = t >> 5;
for (int i = t; i < G * 128; i += NTHR) {
int gg = i >> 7, cc = i & 127;
int ncol = min(c0 + cc, N - 1);
spre[i] = m.s[(size_t)gg * N + ncol];
zpre[i] = m.z[(size_t)gg * N + ncol];
}
__syncthreads();
int rw = (R + NWARP - 1) / NWARP;
int r0 = warp * rw;
int r1 = min(R, r0 + rw);
float acc[4] = {0.f, 0.f, 0.f, 0.f};
if (r0 < r1) {
int colbase = c0 + 4 * lane;
bool valid[4];
#pragma unroll
for (int c = 0; c < 4; ++c) valid[c] = (colbase + c) < N;
bool allv = valid[0] && valid[3];
float adot[4] = {0.f, 0.f, 0.f, 0.f};
float ax = 0.f;
const uint8_t* wp = m.w + (size_t)r0 * N + colbase;
for (int r = r0; r < r1; ++r, wp += N) {
float2 xp = xs2[r];
uint32_t pack;
if (allv) {
pack = __ldg((const uint32_t*)wp);
} else {
pack = 0;
#pragma unroll
for (int c = 0; c < 4; ++c)
if (valid[c]) pack |= ((uint32_t)__ldg(wp + c)) << (8 * c);
}
float xsu = xp.x + xp.y;
ax += xsu;
uint32_t b0 = pack & 0xFF, b1 = (pack >> 8) & 0xFF, b2 = (pack >> 16) & 0xFF, b3 = pack >> 24;
adot[0] = fmaf(xp.x, (float)(b0 & 0xF), adot[0]);
adot[0] = fmaf(xp.y, (float)(b0 >> 4), adot[0]);
adot[1] = fmaf(xp.x, (float)(b1 & 0xF), adot[1]);
adot[1] = fmaf(xp.y, (float)(b1 >> 4), adot[1]);
adot[2] = fmaf(xp.x, (float)(b2 & 0xF), adot[2]);
adot[2] = fmaf(xp.y, (float)(b2 >> 4), adot[2]);
adot[3] = fmaf(xp.x, (float)(b3 & 0xF), adot[3]);
adot[3] = fmaf(xp.y, (float)(b3 >> 4), adot[3]);
if (((r & 63) == 63) || (r == r1 - 1)) {
int gidx = r >> 6;
#pragma unroll
for (int c = 0; c < 4; ++c) {
float sv = b2f(spre[gidx * 128 + 4 * lane + c]);
float zv = b2f(zpre[gidx * 128 + 4 * lane + c]);
acc[c] = fmaf(sv, adot[c] - zv * ax, acc[c]);
adot[c] = 0.f;
}
ax = 0.f;
}
}
}
float* rrow = red + warp * 128;
#pragma unroll
for (int c = 0; c < 4; ++c) rrow[4 * lane + c] = acc[c];
__syncthreads();
if (t < 128) {
float y = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; ++w) y += red[w * 128 + t];
yout[t] = y;
}
__syncthreads();
}
// ---------------------------------------------------------------------------
__device__ void s1_kda(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS);
float* red = (float*)(smem + SM_RED);
bf16* spre = (bf16*)(smem + SM_SPRE);
bf16* zpre = (bf16*)(smem + SM_ZPRE);
float* yout = (float*)(smem + SM_YOUT);
build_xs_norm(P, l, 0, xs2, red);
const LayerW& W = P.lyr[l];
int t = threadIdx.x;
for (int it = blockIdx.x; it < 130; it += P.nCTA) {
if (it < 128) {
int mat = it >> 5;
int c0 = (it & 31) * 128;
QMat m = (mat == 0) ? W.q : (mat == 1) ? W.k : (mat == 2) ? W.v : W.g;
qgemv128(m, HID / 2, C4096, HID / 128, c0, xs2, spre, zpre, red, yout);
if (mat < 3) {
const bf16* prev = (mat == 0) ? P.st.cq_in[l] : (mat == 1) ? P.st.ck_in[l] : P.st.cv_in[l];
bf16* pout = (mat == 0) ? P.st.cq_out[l] : (mat == 1) ? P.st.ck_out[l] : P.st.cv_out[l];
bf16* store = (mat == 0) ? P.qs : (mat == 1) ? P.ks : P.vs;
const bf16* cw = W.convw + (size_t)mat * C4096 * 4;
for (int c = t; c < 128; c += NTHR) {
int ch = c0 + c;
bf16 raw = f2b(yout[c]);
float o = b2f(prev[0 * C4096 + ch]) * b2f(cw[ch * 4 + 0])
+ b2f(prev[1 * C4096 + ch]) * b2f(cw[ch * 4 + 1])
+ b2f(prev[2 * C4096 + ch]) * b2f(cw[ch * 4 + 2])
+ b2f(raw) * b2f(cw[ch * 4 + 3]);
store[ch] = f2b(silu(o));
pout[0 * C4096 + ch] = prev[1 * C4096 + ch];
pout[1 * C4096 + ch] = prev[2 * C4096 + ch];
pout[2 * C4096 + ch] = raw;
}
} else {
for (int c = t; c < 128; c += NTHR) P.gs[c0 + c] = f2b(yout[c]);
}
__syncthreads();
} else if (it == 128) {
int lane = t & 31, warp = t >> 5;
float acc = 0.f;
for (int j = warp * (HID / NWARP) + lane; j < (warp + 1) * (HID / NWARP); j += 32) {
acc = fmaf(((float*)xs2)[j], b2f(W.betaw[j * 32 + lane]), acc);
}
red[warp * 32 + lane] = acc;
__syncthreads();
if (t < 32) {
float y = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; ++w) y += red[w * 32 + t];
P.beta_raw[t] = f2b(y);
}
__syncthreads();
} else if (it == 129) {
if (t < 64) P.logits[t] = 0.f;
if (t == 0) *P.sumsq = 0.f;
__syncthreads();
}
}
}
__device__ void s2_kda(const Params& P, int l, char* smem) {
float* kv = (float*)smem; // 128
float* qv = kv + 128;
float* eg = qv + 128;
float* vv = eg + 128; // 32
float* pred = vv + 32; // 32
float* os = pred + 32; // 32
int t = threadIdx.x;
for (int it = blockIdx.x; it < 128; it += P.nCTA) {
int h = it >> 2;
int j0 = (it & 3) * 32;
for (int i = t; i < 128; i += NTHR) {
kv[i] = b2f(P.ks[h * DK + i]);
qv[i] = b2f(P.qs[h * DK + i]) * QSCALE_KDA;
eg[i] = __expf(-nsoftplus(b2f(P.gs[h * DK + i])));
}
for (int j = t; j < 32; j += NTHR) {
vv[j] = b2f(P.vs[h * DK + j0 + j]);
pred[j] = 0.f;
os[j] = 0.f;
}
__syncthreads();
float beta = sigmoidf_(b2f(P.beta_raw[h]));
int i = t >> 1;
int half = t & 1;
float ki = kv[i], egi = eg[i], qi = qv[i];
float sdec[16];
const float* srow = P.st.S_in[l] + (size_t)h * 16384 + i * DK + j0 + half * 16;
#pragma unroll
for (int jj = 0; jj < 16; jj += 4) {
float4 sv = *(const float4*)(srow + jj);
sdec[jj + 0] = sv.x * egi;
sdec[jj + 1] = sv.y * egi;
sdec[jj + 2] = sv.z * egi;
sdec[jj + 3] = sv.w * egi;
}
#pragma unroll
for (int jj = 0; jj < 16; ++jj) atomicAdd(&pred[half * 16 + jj], ki * sdec[jj]);
__syncthreads();
float* sout = P.st.S_out[l] + (size_t)h * 16384 + i * DK + j0 + half * 16;
float obk[16];
#pragma unroll
for (int jj = 0; jj < 16; ++jj) {
float d = vv[half * 16 + jj] - pred[half * 16 + jj];
obk[jj] = fmaf(beta * ki, d, sdec[jj]);
}
#pragma unroll
for (int jj = 0; jj < 16; jj += 4) {
*(float4*)(sout + jj) = make_float4(obk[jj], obk[jj + 1], obk[jj + 2], obk[jj + 3]);
}
#pragma unroll
for (int jj = 0; jj < 16; ++jj) atomicAdd(&os[half * 16 + jj], qi * obk[jj]);
__syncthreads();
for (int j = t; j < 32; j += NTHR) P.o_att[h * DK + j0 + j] = f2b(os[j]);
__syncthreads();
}
}
__device__ void s3_oproj(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS);
float* red = (float*)(smem + SM_RED);
bf16* spre = (bf16*)(smem + SM_SPRE);
bf16* zpre = (bf16*)(smem + SM_ZPRE);
float* yout = (float*)(smem + SM_YOUT);
float* routed = (float*)(smem + SM_LSM); // 64 scratch in smem
build_xs_from_bf16(P.o_att, C4096, xs2);
const LayerW& W = P.lyr[l];
QMat m = (l < 3) ? W.o : W.mo;
int t = threadIdx.x;
float* acc_cur = (l & 1) ? P.acc1 : P.acc0;
for (int it = blockIdx.x; it < 36; it += P.nCTA) {
if (it < 18) {
int c0 = it * 128;
qgemv128(m, C4096 / 2, HID, C4096 / 128, c0, xs2, spre, zpre, red, yout);
float ss_part = 0.f;
if (t < 64) routed[t] = 0.f;
__syncthreads();
for (int c = t; c < 128; c += NTHR) {
int j = c0 + c;
float tval = b2f(f2b(yout[c]));
float res = resid_elem(P, l, j);
float h = b2f(f2b(tval + res));
P.h1[j] = f2b(h);
yout[c] = h;
ss_part = fmaf(h, h, ss_part);
}
__syncthreads();
#pragma unroll
for (int off = 16; off > 0; off >>= 1) ss_part += __shfl_xor_sync(0xffffffffu, ss_part, off);
if ((t & 31) == 0) red[t >> 5] = ss_part;
__syncthreads();
if (t == 0) {
float v = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; ++w) v += red[w];
atomicAdd(P.sumsq, v);
}
{
int e = t & 63, jb = t >> 6;
float acc = 0.f;
for (int c = jb; c < 128; c += 4) {
int j = c0 + c;
acc = fmaf(yout[c], W.wnr[(size_t)j * 64 + e], acc);
}
atomicAdd(&routed[e], acc);
__syncthreads();
if (t < 64) atomicAdd(&P.logits[t], routed[t]);
}
__syncthreads();
} else {
int z0 = (it - 18) * 128;
for (int c = t; c < 128; c += NTHR) acc_cur[z0 + c] = 0.f;
__syncthreads();
}
}
}
__device__ void route_select(const Params& P, int l, float r, float* rw, int* ridx, float* lsm) {
int t = threadIdx.x;
if (t < 32) {
lsm[t] = P.logits[t] * r;
lsm[t + 32] = P.logits[t + 32] * r;
}
__syncthreads();
if (t == 0) {
float mx = -1e30f;
for (int e = 0; e < NEXP; ++e) mx = fmaxf(mx, lsm[e]);
float sum = 0.f;
for (int e = 0; e < NEXP; ++e) { lsm[e] = __expf(lsm[e] - mx); sum += lsm[e]; }
float inv = 1.f / sum;
for (int e = 0; e < NEXP; ++e) lsm[e] *= inv;
for (int j = 0; j < NACT; ++j) {
int best = -1; float bv = -1.f;
for (int e = 0; e < NEXP; ++e) {
if (lsm[e] > bv) { bv = lsm[e]; best = e; }
}
if (best < 0) best = 0;
ridx[j] = best;
rw[j] = bv;
lsm[best] = -1.f;
}
float wsum = 0.f;
for (int j = 0; j < NACT; ++j) wsum += rw[j];
for (int j = 0; j < NACT; ++j) rw[j] = rw[j] / (wsum + 1e-9f) * ROUTED_SCALE;
}
__syncthreads();
}
__device__ void s4_gateup(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS);
float* red = (float*)(smem + SM_RED);
bf16* spre = (bf16*)(smem + SM_SPRE);
bf16* zpre = (bf16*)(smem + SM_ZPRE);
float* yout = (float*)(smem + SM_YOUT);
float* yout2 = (float*)(smem + SM_YOUT2);
float* rw = (float*)(smem + SM_ROUTE);
int* ridx = (int*)(smem + SM_ROUTE + 64);
float* lsm = (float*)(smem + SM_LSM);
const LayerW& W = P.lyr[l];
build_xs_norm(P, l, 2, xs2, red);
float r = rsqrtf(*P.sumsq / HID + EPS);
route_select(P, l, r, rw, ridx, lsm);
int t = threadIdx.x;
for (int it = blockIdx.x; it < 72; it += P.nCTA) {
int slot = it >> 3;
int c0 = (it & 7) * 128;
int e = (slot < NACT) ? ridx[slot] : 0;
QMat g = (slot == NACT) ? W.sgate : W.gate;
QMat u = (slot == NACT) ? W.sup : W.up;
size_t offw = (size_t)e * (HID / 2) * INTER;
size_t offs = (size_t)e * (HID / 128) * INTER;
QMat mg = {g.w + offw, g.s + offs, g.z + offs};
QMat mu = {u.w + offw, u.s + offs, u.z + offs};
qgemv128(mg, HID / 2, INTER, HID / 128, c0, xs2, spre, zpre, red, yout);
qgemv128(mu, HID / 2, INTER, HID / 128, c0, xs2, spre, zpre, red, yout2);
for (int c = t; c < 128; c += NTHR) {
P.m_scratch[slot * INTER + c0 + c] = silu(yout[c]) * yout2[c];
}
__syncthreads();
}
}
__device__ void s5_down(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS);
float* red = (float*)(smem + SM_RED);
bf16* spre = (bf16*)(smem + SM_SPRE);
bf16* zpre = (bf16*)(smem + SM_ZPRE);
float* yout = (float*)(smem + SM_YOUT);
float* rw = (float*)(smem + SM_ROUTE);
int* ridx = (int*)(smem + SM_ROUTE + 64);
float* lsm = (float*)(smem + SM_LSM);
const LayerW& W = P.lyr[l];
float r = rsqrtf(*P.sumsq / HID + EPS);
route_select(P, l, r, rw, ridx, lsm);
int t = threadIdx.x;
float* acc_cur = (l & 1) ? P.acc1 : P.acc0;
for (int it = blockIdx.x; it < 162; it += P.nCTA) {
int slot = it / 18;
int c0 = (it % 18) * 128;
int e = (slot < NACT) ? ridx[slot] : 0;
float w_e = (slot < NACT) ? rw[slot] : 1.f;
QMat dm = (slot == NACT) ? W.sdown : W.down;
size_t offw = (size_t)e * (INTER / 2) * HID;
size_t offs = (size_t)e * (INTER / 128) * HID;
QMat m = {dm.w + offw, dm.s + offs, dm.z + offs};
for (int j = t; j < INTER; j += NTHR) ((float*)xs2)[j] = P.m_scratch[slot * INTER + j];
__syncthreads();
qgemv128(m, INTER / 2, HID, INTER / 128, c0, xs2, spre, zpre, red, yout);
for (int c = t; c < 128; c += NTHR) {
atomicAdd(&acc_cur[c0 + c], w_e * yout[c]);
}
__syncthreads();
}
}
// ---------------------------------------------------------------------------
__device__ void m1_mla(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS);
float* red = (float*)(smem + SM_RED);
bf16* spre = (bf16*)(smem + SM_SPRE);
bf16* zpre = (bf16*)(smem + SM_ZPRE);
float* yout = (float*)(smem + SM_YOUT);
build_xs_norm(P, l, 0, xs2, red);
const LayerW& W = P.lyr[l];
int t = threadIdx.x;
int pos = P.L;
for (int it = blockIdx.x; it < 54; it += P.nCTA) {
if (it < 48) {
int c0 = it * 128;
qgemv128(W.mq, HID / 2, QDIM, HID / 128, c0, xs2, spre, zpre, red, yout);
for (int c = t; c < 128; c += NTHR) {
int col = c0 + c;
int off = col % QHB;
float v = yout[c];
if (off >= NOPE) {
int p = col ^ 1;
int pf = p - c0;
float pv = yout[pf];
int idx = (off - NOPE) >> 1;
bool even = ((off - NOPE) & 1) == 0;
float ang = pos * P.invf[idx];
float cs, sn;
sincosf(ang, &sn, &cs);
float e = even ? v : pv;
float o = even ? pv : v;
P.qs[col] = f2b(even ? (e * cs - o * sn) : (o * cs + e * sn));
} else {
P.qs[col] = f2b(v);
}
}
__syncthreads();
} else if (it < 53) {
int c0 = (it - 48) * 128;
qgemv128(W.mka, HID / 2, KVA, HID / 128, c0, xs2, spre, zpre, red, yout);
for (int c = t; c < 128; c += NTHR) {
int col = c0 + c;
if (col >= KVA) continue;
float v = yout[c];
if (col < LORA) {
P.st.ckv_out[(size_t)pos * LORA + col] = f2b(v);
} else {
int pf = (col ^ 1) - c0;
float pv = yout[pf];
int idx = (col - LORA) >> 1;
bool even = ((col - LORA) & 1) == 0;
float ang = pos * P.invf[idx];
float cs, sn;
sincosf(ang, &sn, &cs);
float e = even ? v : pv;
float o = even ? pv : v;
P.st.kr_out[(size_t)pos * ROPE + (col - LORA)] = f2b(even ? (e * cs - o * sn) : (o * cs + e * sn));
}
}
__syncthreads();
} else if (it == 53) {
if (t < 64) P.logits[t] = 0.f;
if (t == 0) *P.sumsq = 0.f;
__syncthreads();
}
}
}
__device__ void m2_qlat(const Params& P, int l, char* smem) {
float* qno = (float*)smem; // 128 fp32
float* spre = (float*)(smem + 1024); // [4][128] fp32
float* zpre = (float*)(smem + 3072); // [4][128] fp32
int t = threadIdx.x;
const LayerW& W = P.lyr[l];
for (int it = blockIdx.x; it < NHEADS; it += P.nCTA) {
int h = it;
for (int i = t; i < 128; i += NTHR) qno[i] = b2f(P.qs[h * QHB + i]);
int nbase = h * 256;
for (int i = t; i < 4 * 128; i += NTHR) {
int gg = i >> 7, cc = i & 127;
spre[i] = b2f(W.mkb.s[(size_t)gg * KVBOUT + nbase + cc]);
zpre[i] = b2f(W.mkb.z[(size_t)gg * KVBOUT + nbase + cc]);
}
__syncthreads();
for (int jj = t; jj < LORA; jj += NTHR) {
const uint8_t* wrow = W.mkb.w + (size_t)(jj >> 1) * KVBOUT + nbase;
int g = jj >> 7;
bool odd = jj & 1;
float acc = 0.f;
#pragma unroll 4
for (int d4 = 0; d4 < 128; d4 += 4) {
uint32_t pack = __ldg((const uint32_t*)(wrow + d4));
#pragma unroll
for (int c = 0; c < 4; ++c) {
int nib = (pack >> (8 * c)) & 0xFF;
float w = (float)(odd ? (nib >> 4) : (nib & 0xF));
int d = d4 + c;
float coef = spre[g * 128 + d] * (w - zpre[g * 128 + d]);
acc = fmaf(coef, qno[d], acc);
}
}
P.qlat[h * LORA + jj] = acc;
}
__syncthreads();
}
}
__device__ void m3_attn(const Params& P, int l, char* smem) {
float* qlat = (float*)(smem + SM_XS); // [8][512] = 16KB
float* oacc = (float*)(smem + 16384); // [8][512] = 16KB
float* qrope = (float*)(smem + 32768); // [8][64] = 2KB
int t = threadIdx.x;
int lane = t & 31, warp = t >> 5;
int nc = P.nchunks;
for (int it = blockIdx.x; it < 4 * nc; it += P.nCTA) {
int hg = it / nc;
int c = it % nc;
int l0 = c * LCCHUNK;
int l1 = min(l0 + LCCHUNK, P.L + 1);
for (int i = t; i < 8 * 512; i += NTHR) {
int hh = i >> 9, j = i & 511;
qlat[hh * 512 + j] = P.qlat[(hg * 8 + hh) * LORA + j];
}
for (int i = t; i < 8 * 64; i += NTHR) {
int hh = i >> 6, j = i & 63;
qrope[hh * 64 + j] = b2f(P.qs[(hg * 8 + hh) * QHB + NOPE + j]);
}
for (int i = t; i < 8 * 512; i += NTHR) oacc[i] = 0.f;
__syncthreads();
int hloc = warp;
float m = -INFINITY, lsum = 0.f;
for (int cl = l0; cl < l1; ++cl) {
const bf16* crow = (cl < P.L) ? (P.st.ckv_in + (size_t)cl * LORA)
: (P.st.ckv_out + (size_t)P.L * LORA);
const bf16* krow = (cl < P.L) ? (P.st.kr_in + (size_t)cl * ROPE)
: (P.st.kr_out + (size_t)P.L * ROPE);
uint4 a0 = __ldg((const uint4*)(crow + lane * 16));
uint4 a1 = __ldg((const uint4*)(crow + lane * 16 + 8));
uint32_t krp = __ldg((const uint32_t*)(krow + lane * 2));
if (hg == 0 && hloc == 0 && cl < P.L) {
*(uint4*)(P.st.ckv_out + (size_t)cl * LORA + lane * 16) = a0;
*(uint4*)(P.st.ckv_out + (size_t)cl * LORA + lane * 16 + 8) = a1;
*(uint32_t*)(P.st.kr_out + (size_t)cl * ROPE + lane * 2) = krp;
}
float cv[16];
{
const bf16* pa = (const bf16*)&a0;
const bf16* pb = (const bf16*)&a1;
#pragma unroll
for (int i = 0; i < 8; ++i) cv[i] = b2f(pa[i]);
#pragma unroll
for (int i = 0; i < 8; ++i) cv[8 + i] = b2f(pb[i]);
}
float kr0 = b2f(((const bf16*)&krp)[0]);
float kr1 = b2f(((const bf16*)&krp)[1]);
float partial = 0.f;
#pragma unroll
for (int i = 0; i < 16; ++i)
partial = fmaf(cv[i], qlat[hloc * 512 + lane * 16 + i], partial);
partial = fmaf(kr0, qrope[hloc * 64 + lane * 2], partial);
partial = fmaf(kr1, qrope[hloc * 64 + lane * 2 + 1], partial);
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
partial += __shfl_xor_sync(0xffffffffu, partial, off);
float s = partial * QSCALE_MLA;
float mnew = fmaxf(m, s);
float f = __expf(m - mnew);
float pt = __expf(s - mnew);
lsum = lsum * f + pt;
m = mnew;
#pragma unroll
for (int i = 0; i < 16; ++i) {
float* o = &oacc[hloc * 512 + lane * 16 + i];
*o = fmaf(*o, f, pt * cv[i]);
}
}
float* dst = P.part + (((size_t)c * 4 + hg) * 8 + hloc) * 514;
#pragma unroll
for (int i = 0; i < 16; ++i) dst[lane * 16 + i] = oacc[hloc * 512 + lane * 16 + i];
if (lane == 0) { dst[512] = m; dst[513] = lsum; }
__syncthreads();
}
}
__device__ void m4_combine(const Params& P, int l, char* smem) {
int t = threadIdx.x;
int nc = P.nchunks;
for (int it = blockIdx.x; it < 128; it += P.nCTA) {
int h = it >> 2;
int dc = it & 3;
int j0 = dc * 128;
int hg = h >> 3, hloc = h & 7;
float M = -INFINITY;
for (int c = 0; c < nc; ++c) {
const float* src = P.part + (((size_t)c * 4 + hg) * 8 + hloc) * 514;
M = fmaxf(M, src[512]);
}
float den = 0.f;
float acc = 0.f;
int j = j0 + (t & 127);
bool active = t < 128;
for (int c = 0; c < nc; ++c) {
const float* src = P.part + (((size_t)c * 4 + hg) * 8 + hloc) * 514;
float mc = src[512], lc = src[513];
float f = __expf(mc - M);
den += f * lc;
if (active) acc = fmaf(f, src[j], acc);
}
if (active) P.olat[h * LORA + j] = acc / den;
__syncthreads();
}
}
__device__ void m5_wv(const Params& P, int l, char* smem) {
float2* xs2 = (float2*)(smem + SM_XS); // 256 pairs
float* spre = (float*)(smem + 2048); // [4][32] fp32
float* zpre = (float*)(smem + 2560); // [4][32] fp32
float* red = (float*)(smem + 3072); // [8][32]
int t = threadIdx.x;
const LayerW& W = P.lyr[l];
for (int it = blockIdx.x; it < 128; it += P.nCTA) {
int h = it >> 2;
int dc = it & 3;
int nbase = h * 256 + NOPE + dc * 32;
for (int i = t; i < LORA; i += NTHR) ((float*)xs2)[i] = P.olat[h * LORA + i];
for (int i = t; i < 4 * 32; i += NTHR) {
int gg = i >> 5, cc = i & 31;
spre[i] = b2f(W.mkb.s[(size_t)gg * KVBOUT + nbase + cc]);
zpre[i] = b2f(W.mkb.z[(size_t)gg * KVBOUT + nbase + cc]);
}
__syncthreads();
int lane = t & 31, warp = t >> 5;
int r0 = warp * 32, r1 = r0 + 32;
float acc = 0.f, ax = 0.f;
const uint8_t* wp = W.mkb.w + (size_t)r0 * KVBOUT + nbase + lane;
for (int r = r0; r < r1; ++r, wp += KVBOUT) {
float2 xp = xs2[r];
uint32_t b = __ldg(wp);
float wl = (float)(b & 0xF), wh = (float)(b >> 4);
ax += xp.x + xp.y;
acc = fmaf(xp.x, wl, acc);
acc = fmaf(xp.y, wh, acc);
}
{
int gidx = r0 >> 6;
float sv = spre[gidx * 32 + lane];
float zv = zpre[gidx * 32 + lane];
acc = sv * (acc - zv * ax);
}
red[warp * 32 + lane] = acc;
__syncthreads();
if (t < 32) {
float y = 0.f;
#pragma unroll
for (int w = 0; w < NWARP; ++w) y += red[w * 32 + t];
P.o_att[h * VH + dc * 32 + t] = f2b(y);
}
__syncthreads();
}
}
__device__ void stage_writeout(const Params& P, char* smem) {
int t = threadIdx.x;
for (int it = blockIdx.x; it < 18; it += P.nCTA) {
int j0 = it * 128;
for (int j = j0 + t; j < j0 + 128; j += NTHR) {
float h = b2f(P.h1[j]);
float a = b2f(f2b(P.acc1[j]));
P.hout[j] = f2b(b2f(f2b(h + a)));
}
}
}
// ---------------------------------------------------------------------------
#define SYNC_OR_STOP(sid_) do { gbar(P.bar); if (P.stop == (sid_)) return; } while(0)
__global__ void megakernel(Params P) {
extern __shared__ char smem[];
for (int l = 0; l < 4; ++l) {
if (l < 3) {
s1_kda(P, l, smem); SYNC_OR_STOP(l * 5 + 0);
s2_kda(P, l, smem); SYNC_OR_STOP(l * 5 + 1);
s3_oproj(P, l, smem); SYNC_OR_STOP(l * 5 + 2);
s4_gateup(P, l, smem); SYNC_OR_STOP(l * 5 + 3);
s5_down(P, l, smem); SYNC_OR_STOP(l * 5 + 4);
} else {
m1_mla(P, l, smem); SYNC_OR_STOP(15 + 0);
m2_qlat(P, l, smem); SYNC_OR_STOP(15 + 1);
m3_attn(P, l, smem); SYNC_OR_STOP(15 + 2);
m4_combine(P, l, smem); SYNC_OR_STOP(15 + 3);
m5_wv(P, l, smem); SYNC_OR_STOP(15 + 4);
s3_oproj(P, l, smem); SYNC_OR_STOP(15 + 5);
s4_gateup(P, l, smem); SYNC_OR_STOP(15 + 6);
s5_down(P, l, smem); SYNC_OR_STOP(15 + 7);
}
}
stage_writeout(P, smem);
}
// ---------------------------------------------------------------------------
static Params g_P;
static bool g_inited = false;
static bool g_attr_set = false;
static void set_qmat(QMat& m, const at::Tensor& w, const at::Tensor& s, const at::Tensor& z) {
m.w = (const uint8_t*)w.data_ptr();
m.s = (const bf16*)s.data_ptr();
m.z = (const bf16*)z.data_ptr();
}
void mega_init(std::vector<at::Tensor> wts, std::vector<at::Tensor> scratch) {
Params& P = g_P;
int idx = 0;
for (int l = 0; l < 4; ++l) {
LayerW& W = P.lyr[l];
W.anorm = (const bf16*)wts[idx++].data_ptr();
W.mnorm = (const bf16*)wts[idx++].data_ptr();
set_qmat(W.q, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.k, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.v, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.g, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.o, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
W.betaw = (const bf16*)wts[idx++].data_ptr();
W.convw = (const bf16*)wts[idx++].data_ptr();
set_qmat(W.mq, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.mka, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.mkb, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.mo, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
W.wnr = (const float*)wts[idx++].data_ptr();
set_qmat(W.gate, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.up, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.down, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.sgate, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.sup, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
set_qmat(W.sdown, wts[idx], wts[idx + 1], wts[idx + 2]); idx += 3;
}
int s = 0;
P.qs = (bf16*)scratch[s++].data_ptr();
P.ks = (bf16*)scratch[s++].data_ptr();
P.vs = (bf16*)scratch[s++].data_ptr();
P.gs = (bf16*)scratch[s++].data_ptr();
P.beta_raw = (bf16*)scratch[s++].data_ptr();
P.o_att = (bf16*)scratch[s++].data_ptr();
P.h1 = (bf16*)scratch[s++].data_ptr();
P.acc0 = (float*)scratch[s++].data_ptr();
P.acc1 = (float*)scratch[s++].data_ptr();
P.sumsq = (float*)scratch[s++].data_ptr();
P.logits = (float*)scratch[s++].data_ptr();
P.qlat = (float*)scratch[s++].data_ptr();
P.olat = (float*)scratch[s++].data_ptr();
P.part = (float*)scratch[s++].data_ptr();
P.invf = (const float*)scratch[s++].data_ptr();
P.m_scratch = (float*)scratch[s++].data_ptr();
P.bar = (unsigned*)scratch[s++].data_ptr();
g_inited = true;
}
at::Tensor mega_step(at::Tensor hin, at::Tensor hout,
std::vector<at::Tensor> kda_in, std::vector<at::Tensor> kda_out,
at::Tensor ckv_in, at::Tensor kr_in,
at::Tensor ckv_out, at::Tensor kr_out,
int64_t L, int64_t nCTA, int64_t stop) {
TORCH_CHECK(g_inited, "mega not inited");
Params P = g_P;
P.stop = (int)stop;
P.hin = (const bf16*)hin.data_ptr();
P.hout = (bf16*)hout.data_ptr();
for (int i = 0; i < 3; ++i) {
P.st.S_in[i] = (const float*)kda_in[i * 4 + 0].data_ptr();
P.st.cq_in[i] = (const bf16*)kda_in[i * 4 + 1].data_ptr();
P.st.ck_in[i] = (const bf16*)kda_in[i * 4 + 2].data_ptr();
P.st.cv_in[i] = (const bf16*)kda_in[i * 4 + 3].data_ptr();
P.st.S_out[i] = (float*)kda_out[i * 4 + 0].data_ptr();
P.st.cq_out[i] = (bf16*)kda_out[i * 4 + 1].data_ptr();
P.st.ck_out[i] = (bf16*)kda_out[i * 4 + 2].data_ptr();
P.st.cv_out[i] = (bf16*)kda_out[i * 4 + 3].data_ptr();
}
P.st.ckv_in = (const bf16*)ckv_in.data_ptr();
P.st.kr_in = (const bf16*)kr_in.data_ptr();
P.st.ckv_out = (bf16*)ckv_out.data_ptr();
P.st.kr_out = (bf16*)kr_out.data_ptr();
P.L = (int)L;
P.nchunks = (int)((L + 1 + LCCHUNK - 1) / LCCHUNK);
P.nCTA = (int)nCTA;
if (!g_attr_set) {
cudaFuncSetAttribute(megakernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE);
g_attr_set = true;
}
auto stream = at::cuda::getCurrentCUDAStream();
megakernel<<<(int)nCTA, NTHR, SMEM_SIZE, stream>>>(P);
return hout;
}
int64_t mega_occupancy() {
if (!g_attr_set) {
cudaFuncSetAttribute(megakernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE);
g_attr_set = true;
}
int dev = 0;
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
int nb = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&nb, megakernel, NTHR, SMEM_SIZE);
return (int64_t)prop.multiProcessorCount * nb;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("mega_init", &mega_init);
m.def("mega_step", &mega_step);
m.def("mega_occupancy", &mega_occupancy);
}
"""
_ext = None
_ext_err = None
if os.environ.get("KIMI_FORCE_EAGER", "0") != "1":
try:
from torch.utils.cpp_extension import load_inline
_cc = torch.cuda.get_device_capability(0)
_name = f"kimi_linear_megak_sm{_cc[0]}{_cc[1]}"
_ext = load_inline(
name=_name,
cpp_sources=[],
cuda_sources=[_CUDA_SRC],
extra_cuda_cflags=["-O3", "--expt-relaxed-constexpr"],
verbose=False,
)
except Exception as e: # pragma: no cover
_ext_err = e
class Model(nn.Module):
"""Drops the eager block-iteration step for the single-launch megakernel."""
def __init__(self, cfg: Config):
super().__init__()
self.cfg = cfg
self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern)
self._built = False
self._keep = None
self._nCTA = 0
def _build_mega(self):
dev = next(self.parameters()).device
cfg = self.cfg
wts = []
for l, blk in enumerate(self.blocks):
W = []
a = blk.attn
W.append(blk.attn_norm.detach().contiguous())
W.append(blk.moe_norm.detach().contiguous())
if l < 3:
for prj in (a.q_proj, a.k_proj, a.v_proj, a.g_proj, a.o_proj):
W += [prj.w_q.contiguous(), prj.scales.contiguous(), prj.zeros.contiguous()]
W.append(a.beta_proj.weight.detach().t().contiguous())
W.append(a.conv_w.detach().contiguous())
else:
for prj in (a.q_proj, a.q_proj, a.q_proj, a.q_proj, a.o_proj):
W += [prj.w_q.contiguous(), prj.scales.contiguous(), prj.zeros.contiguous()]
W.append(torch.zeros(32, cfg.hidden, dtype=cfg.dtype, device=dev).t().contiguous())
W.append(torch.zeros(3, cfg.kda_heads * cfg.kda_head_dim, cfg.short_conv,
dtype=cfg.dtype, device=dev))
for prj in (a.q_proj, getattr(a, "kv_a", a.q_proj), getattr(a, "kv_b", a.q_proj), a.o_proj):
W += [prj.w_q.contiguous(), prj.scales.contiguous(), prj.zeros.contiguous()]
wr = blk.moe.router.weight.detach()
W.append((wr.float() * blk.moe_norm.detach().float()[None, :]).t().contiguous())
m = blk.moe
for qe in (m.gate, m.up, m.down, m.s_gate, m.s_up, m.s_down):
W += [qe.w_q.contiguous(), qe.scales.contiguous(), qe.zeros.contiguous()]
wts.extend(W)
def sc(shape, dtype):
return torch.zeros(shape, dtype=dtype, device=dev)
nc = (20000 // 256) + 4
scratch = [
sc(6144, torch.bfloat16), # qs (also MLA q)
sc(4096, torch.bfloat16), # ks
sc(4096, torch.bfloat16), # vs
sc(4096, torch.bfloat16), # gs
sc(32, torch.bfloat16), # beta
sc(4096, torch.bfloat16), # o_att
sc(2304, torch.bfloat16), # h1
sc(2304, torch.float32), # acc0
sc(2304, torch.float32), # acc1
sc(1, torch.float32), # sumsq
sc(64, torch.float32), # logits
sc(32 * 512, torch.float32), # qlat
sc(32 * 512, torch.float32), # olat
sc(nc * 4 * 8 * 514, torch.float32), # attn partials
(cfg.rope_theta ** (-torch.arange(0, cfg.qk_rope, 2, dtype=torch.float32) / cfg.qk_rope)).to(dev),
sc(9 * 1024, torch.float32), # m
torch.zeros(2, dtype=torch.int32, device=dev), # barrier words
]
_ext.mega_init(wts, scratch)
self._nCTA = _ext.mega_occupancy()
self._keep = (wts, scratch)
self._built = True
def _step_mega(self, hidden, state):
cfg = self.cfg
L = state[3]["c_kv"].shape[0]
kda_in, kda_out = [], []
new_state = []
for i in range(3):
st = state[i]
S_out = torch.empty_like(st["S"])
cq_out = torch.empty_like(st["cq"])
ck_out = torch.empty_like(st["ck"])
cv_out = torch.empty_like(st["cv"])
kda_in += [st["S"].contiguous(), st["cq"].contiguous(), st["ck"].contiguous(), st["cv"].contiguous()]
kda_out += [S_out, cq_out, ck_out, cv_out]
new_state.append({"S": S_out, "cq": cq_out, "ck": ck_out, "cv": cv_out})
ckv = state[3]["c_kv"].contiguous()
kro = state[3]["k_rope"].contiguous()
ckv_out = torch.empty(L + 1, cfg.kv_lora, dtype=cfg.dtype, device=hidden.device)
kr_out = torch.empty(L + 1, cfg.qk_rope, dtype=cfg.dtype, device=hidden.device)
hout = torch.empty(cfg.hidden, dtype=cfg.dtype, device=hidden.device)
_ext.mega_step(hidden.contiguous(), hout, kda_in, kda_out, ckv, kro,
ckv_out, kr_out, L, self._nCTA, 999)
new_state.append({"c_kv": ckv_out, "k_rope": kr_out})
return hout, new_state
def step(self, hidden, state):
if _ext is not None:
if not self._built:
self._build_mega()
return self._step_mega(hidden, state)
for i, blk in enumerate(self.blocks):
hidden = blk.step(hidden, state[i])
return hidden, state
# --------------------------------------------------------------------------- #
# state / inputs
# --------------------------------------------------------------------------- #
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
20260716_150057_kinetic-claude_kinetic-0715_1m__02_kimi_linear_decode