KernelBench mega · RTX PRO 6000
Kimi-Linear Decode DeepSeek V4.1 Flash
17.10×geomean speedup across shapes
DeepSeek V4.1 Flash built a real single-launch cooperative megakernel for the Kimi-Linear decode block: int4 dequant-GEMVs in registers, short conv, gated-delta recurrence, absorbed MLA, full MoE. Fast and honestly earned, but it computes the wrong answer inside the 0.98 cosine gate: a persistent atomicAdd accumulator is never zeroed, so every KDA layer inherits the previous one's activations.
harnessdeepseek-claude
Kernel source (redacted)
"""Kimi-Linear W4A16 hybrid decode -- single-launch megakernel solution.
The timed path is ONE custom CUDA kernel (CUDA_SRC below), launched exactly once
per step() and fusing the entire 4-block [K,K,K,M] forward: every int4
dequant-GEMV (q/k/v/g/o_proj, MLA q_proj/kv_a/kv_b/o_proj, all 9 MoE expert
gate/up/down), the short causal conv, the KDA recurrent state update, the MLA
absorbed latent-cache attention, the router/top-8/shared expert and both
RMSNorms + residual adds. No intermediate tensor ever touches DRAM except the
hand-rolled scratch the kernel itself uses for cross-CTA partial sums.
The module tree below exists only so that
`load_state_dict(reference.state_dict(), strict=True)` maps every reference
buffer/parameter name onto the raw storage the kernel reads; nothing in it
runs on the timed path. A plain-PyTorch debug path (_eager_step) is kept for
local development and as a fallback when the fused kernel cannot be used.
Design notes (see scratch/ for the development copy of the kernel):
* grid = #SMs, 512 threads/CTA, 83968 B dynamic smem, 1 CTA/SM. The SM caps
opt-in dynamic smem at 101376 B and the block sizes itself to fit one CTA;
the kernel uses 121 registers, so 512 threads is also the register-file
ceiling (2 CTAs would need both smem and regs to halve, neither of which
is close).
* 19 cooperative-grid barriers per token; every one is required by a real
cross-CTA producer/consumer edge, not a scheduling convenience.
* weights streamed once at int4; the zero-point correction is folded into
3 FMA + 2 BFE per (column, row-pair) instead of materialising bf16.
* MLA is done in absorbed form: q_nope @ Wb -> 512-d query per head, and
p @ c_kv, then @ Wb[:, v-part]; kv_b is never materialised.
"""
from __future__ import annotations
import sys
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.cpp_extension import load_inline
OP_TYPE = "kimi_linear_w4a16_decode"
EPS = 1.0e-6
GROUP_SIZE = 128
# ---8<--- CUDA_BEGIN
CUDA_SRC = r"""
// Kimi-Linear W4A16 hybrid decode megakernel -- one launch per token.
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#define D 2304
#define NH 32
#define DK 128
#define CTOT 4096
#define MOED 1024
#define NALL 9
#define NTS 512
#define NWARP 16
#define LRMAX 96
#define NCH 576
#define SCALE 0.07216878364870322f
#define QSCALE 0.08838834764831845f
#define EPS 1.0e-6f
// dynamic smem (floats). sm_120 caps opt-in dynamic smem at 101376 B and the
// sm at 102400 B total, so every buffer has to be small: the MLA KV rows are
// read straight out of the padded global cache rather than staged here.
#define O_X 0
#define O_HN 2304 // rmsnorm'd hidden
#define O_AUX 4608 // 9 * 1024 (router probs / moe silu scratch)
#define O_SC 13824 // 32 * LRMAX fp32 scores
#define O_M5X 16896 // MLA o_proj input staging (4096 f32)
#define SMEM_TOTAL 20992
#define POPRE (NH * 4 * DK) // second half of the K2 partials
#define FPRED (2 * NH * 4 * DK) // folded pred[j]
#define FOPRE (FPRED + NH * DK) // folded o[j] pre-scalar
#define BQKO (FOPRE + NH * DK) // per-head sum_i q[i]*beta[i]*k[i]
#define PREDSZ (BQKO + NH)
struct Params {
float *xb, *raw, *ob, *attn_out, *hb, *hhg, *hhu;
float *moe_out;
float *router, *betab, *selw;
int *sel, *cnt;
float *cw, *pm, *ps;
float *mlab, *mkv, *mob, *pred;
__nv_bfloat16 *qabs_g, *ckv_pad, *kr_pad;
const long long *wdesc;
const __nv_bfloat16 *attn_norm, *moe_norm, *beta_w, *router_w, *conv_w;
float *S0, *S1, *S2;
__nv_bfloat16 *cq0, *ck0, *cv0, *cq1, *ck1, *cv1, *cq2, *ck2, *cv2;
const __nv_bfloat16 *hidden_in;
__nv_bfloat16 *hidden_out;
const __nv_bfloat16 *src_ckv, *src_kr;
int pos, grid, ncopy;
};
__device__ __forceinline__ float b2f(__nv_bfloat16 x) { return __bfloat162float(x); }
__device__ __forceinline__ float softplus_(float x) {
return log1pf(__expf(-fabsf(x))) + fmaxf(x, 0.0f);
}
__device__ __forceinline__ float sigmoid_(float x) { return 1.0f / (1.0f + __expf(-x)); }
__device__ __forceinline__ float silu_(float x) { return x / (1.0f + __expf(-x)); }
// A 32-bit word holding two bf16 lanes. Shifting the low half into the exponent
// and clearing the high half's low bits is an exact bf16->fp32 widening, so no
// address is taken on a register and nothing lands in local memory.
__device__ __forceinline__ float bfp_lo(unsigned int u) { return __uint_as_float(u << 16); }
__device__ __forceinline__ float bfp_hi(unsigned int u) { return __uint_as_float(u & 0xFFFF0000u); }
__device__ __forceinline__ float ldcg(const float *p) {
float v;
asm volatile("ld.global.cg.f32 %0, [%1];" : "=f"(v) : "l"(p));
return v;
}
struct WD {
const unsigned char *wq;
const __nv_bfloat16 *sc, *zo;
int cols;
long long esw, ess, esz;
};
__device__ __forceinline__ WD wd_load(const long long *p) {
WD w;
w.wq = (const unsigned char *)(size_t)p[0];
w.sc = (const __nv_bfloat16 *)(size_t)p[1];
w.zo = (const __nv_bfloat16 *)(size_t)p[2];
w.cols = (int)p[4];
w.esw = p[5];
w.ess = p[6];
w.esz = p[7];
return w;
}
// Fused int4 dequant + GEMV.
template <int CB>
__device__ __forceinline__ void gemv_add(const unsigned char *__restrict__ wq,
int cols,
const __nv_bfloat16 *__restrict__ sc,
const __nv_bfloat16 *__restrict__ zo,
int g, int r0, int nr, int n0,
const float *__restrict__ xs, float wgt,
float *__restrict__ dst) {
const int lane = threadIdx.x & 31;
const int nbase = n0 + CB * lane;
float s[CB], zc[CB];
const __nv_bfloat16 *sp = sc + (size_t)g * cols + nbase;
const __nv_bfloat16 *zp = zo + (size_t)g * cols + nbase;
#pragma unroll
for (int j = 0; j < CB; ++j) {
s[j] = b2f(sp[j]) * wgt;
zc[j] = -b2f(zp[j]);
}
float rw[CB];
#pragma unroll
for (int j = 0; j < CB; ++j) rw[j] = 0.f;
const unsigned char *base = wq + (size_t)r0 * cols + nbase;
const unsigned char *rp8 = base;
#pragma unroll 8
for (int r = 0; r < nr; ++r) {
unsigned int bv;
if (CB == 4) bv = __ldcs((const unsigned int *)(rp8 + (size_t)r * cols));
else bv = __ldcs((const unsigned short *)(rp8 + (size_t)r * cols));
const float x0 = xs[2 * (r0 + r)];
const float x1 = xs[2 * (r0 + r) + 1];
const float xsm = x0 + x1;
#pragma unroll
for (int j = 0; j < CB; ++j) {
float lo = (float)((bv >> (8 * j)) & 0xF);
float hi = (float)((bv >> (8 * j + 4)) & 0xF);
rw[j] = fmaf(lo, x0, rw[j]);
rw[j] = fmaf(hi, x1, rw[j]);
rw[j] = fmaf(zc[j], xsm, rw[j]);
}
}
#pragma unroll
for (int j = 0; j < CB; ++j) atomicAdd(&dst[nbase + j], rw[j] * s[j]);
}
__device__ __forceinline__ void block_sum(float v, float *red, float *out) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
#pragma unroll
for (int o = 16; o > 0; o >>= 1) v += __shfl_xor_sync(0xffffffffu, v, o);
if (lane == 0) red[warp] = v;
__syncthreads();
if (threadIdx.x == 0) {
float u = 0.f;
for (int i = 0; i < NWARP; ++i) u += red[i];
*out = u;
}
__syncthreads();
}
// Router GEMV, softmax and top-8.
//
// The gate weight is bf16 [D, 64] and stays L2-resident, so every warp in the
// grid streams 32-row slabs of it and atomically accumulates the 64 output
// partials. Running it on CTA 0 alone cost 30-75us per block with the other 187
// CTAs parked at the following grid.sync.
#define RPAD 32
__device__ __forceinline__ void router_matvec(Params &P, int b, const float *shn,
int gw, int nw, int lane) {
const __nv_bfloat16 *rw = P.router_w + (size_t)b * 64 * D;
float *rp = P.router + (size_t)b * 64 * RPAD;
for (int t = gw; t < D / 32; t += nw) {
const __nv_bfloat16 *rr = rw + (size_t)t * 32 * 64;
const float *xk = shn + t * 32;
float a0 = 0.f, a1 = 0.f;
// One 128-byte row per iteration is a single fully-coalesced warp load;
// lane l owns outputs 2l and 2l+1, so no cross-lane reduction is needed.
#pragma unroll 8
for (int r = 0; r < 32; ++r) {
const unsigned int u = *(const unsigned int *)(rr + (size_t)r * 64 + 2 * lane);
const float x = xk[r];
a0 = fmaf(bfp_lo(u), x, a0);
a1 = fmaf(bfp_hi(u), x, a1);
}
atomicAdd(&rp[(2 * lane) * RPAD], a0);
atomicAdd(&rp[(2 * lane + 1) * RPAD], a1);
}
}
// Softmax over the 64 logits then top-8, run redundantly by every CTA after the
// router barrier. Because every CTA computes the same P.sel / P.selw no second
// grid.sync is needed before the MoE reads them. The 8-way selection is a warp
// argmax with the winner masked to -1; ties go to the lowest expert index.
__device__ __forceinline__ void router_finalize(Params &P, int b, int warp, int lane) {
// Only one CTA computes the top-8: all 188 writing the same 8 words is pure
// L2 write contention, and the grid.sync() right after the call already makes
// the result visible grid-wide.
if (blockIdx.x != 0 || warp != 0) return;
const float *rp = P.router + (size_t)b * 64 * RPAD;
float v0 = rp[2 * lane * RPAD];
float v1 = rp[(2 * lane + 1) * RPAD];
float m = fmaxf(v0, v1);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) m = fmaxf(m, __shfl_xor_sync(0xffffffffu, m, o));
float sv0 = __expf(v0 - m), sv1 = __expf(v1 - m);
float s = sv0 + sv1;
#pragma unroll
for (int o = 16; o > 0; o >>= 1) s += __shfl_xor_sync(0xffffffffu, s, o);
const float inv = 1.0f / s;
v0 = sv0 * inv;
v1 = sv1 * inv;
#pragma unroll
for (int j = 0; j < 8; ++j) {
float mm = fmaxf(v0, v1);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) mm = fmaxf(mm, __shfl_xor_sync(0xffffffffu, mm, o));
int w2 = (v0 == mm) ? 2 * lane : ((v1 == mm) ? 2 * lane + 1 : -1);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
const int t2 = __shfl_xor_sync(0xffffffffu, w2, o);
if (t2 >= 0 && (w2 < 0 || t2 < w2)) w2 = t2;
}
if (2 * lane == w2) v0 = -1.f;
else if (2 * lane + 1 == w2) v1 = -1.f;
if (lane == 0) { P.sel[j] = w2; P.selw[j] = mm; }
}
if (lane == 0) {
float tot = 0.f;
for (int j = 0; j < 8; ++j) tot += P.selw[j];
const float sc = 1.0f / (tot + 1e-9f) * 2.5f;
for (int j = 0; j < 8; ++j) P.selw[j] *= sc;
P.sel[8] = 0;
P.selw[8] = 1.0f;
}
}
// gate/up (9 slots, both projections) -> silu combine -> down
__device__ __forceinline__ void moe_body(Params &P, int b, const long long *wd, float *shn,
float *saux, float *dst, cg::grid_group &grid,
int gw, int nw, int tid) {
for (int t = gw; t < NALL * 2 * 8 * 18; t += nw) {
int rest = t;
const int g = rest % 18;
rest /= 18;
const int tile = rest % 8;
rest /= 8;
const int proj = rest & 1;
const int slot = rest >> 1;
const int rt = (slot < 8);
const int e = rt ? P.sel[slot] : 0;
// slot 8 is the shared expert: its gate/up live in descriptor slots 8/9,
// not in the routed experts' 5/6.
WD w = wd_load(wd + (rt ? (proj == 0 ? 5 : 6) : (proj == 0 ? 8 : 9)) * 8);
float *d = (proj == 0 ? P.hhg : P.hhu) + (size_t)slot * MOED;
gemv_add<4>(w.wq + (size_t)e * w.esw, w.cols, w.sc + (size_t)e * w.ess,
w.zo + (size_t)e * w.esz, g, g * 64, 64, tile * 128, shn, 1.0f, d);
}
// every CTA contributes to hhg/hhu with atomicAdd, so the combine needs a
// grid-wide barrier, not just __syncthreads.
grid.sync();
// Each CTA's down-warps between them touch at most two slots (the tasks are
// stratified 144 per slot and a CTA owns 16 consecutive ones), so filling the
// whole 9-slot array in every CTA is 9x the expf work and 9x the L2 traffic
// for nothing. Slots are addressed absolutely, so the slice stays coherent.
{
const int task0 = blockIdx.x * NWARP;
const int task1 = (task0 + NWARP < NALL * 18 * 8) ? task0 + NWARP : NALL * 18 * 8;
if (task0 < task1) {
const int s0 = task0 / (18 * 8);
const int s1 = (task1 - 1) / (18 * 8);
for (int sl = s0; sl <= s1; ++sl) {
const float *hg = P.hhg + (size_t)sl * MOED;
const float *hu = P.hhu + (size_t)sl * MOED;
float *ax = saux + (size_t)sl * MOED;
for (int i = tid; i < MOED; i += NTS) ax[i] = silu_(hg[i]) * hu[i];
}
}
}
__syncthreads();
for (int t = gw; t < NALL * 18 * 8; t += nw) {
int rest = t;
const int g = rest % 8;
rest /= 8;
const int tile = rest % 18;
const int slot = rest / 18;
const int e = (slot < 8) ? P.sel[slot] : 0;
const float wgt = (slot < 8) ? P.selw[slot] : 1.0f;
WD w = wd_load(wd + (slot < 8 ? 7 : 10) * 8);
gemv_add<4>(w.wq + (size_t)e * w.esw, w.cols, w.sc + (size_t)e * w.ess,
w.zo + (size_t)e * w.esz, g, g * 64, 64, tile * 128,
saux + slot * MOED, wgt, dst);
}
// This block's router accumulator has been consumed by router_finalize; clear
// it for the next token. Nothing re-reads it until a later launch.
float *rp = P.router + (size_t)b * 64 * RPAD;
for (int i = tid + blockIdx.x * NTS; i < 64 * RPAD; i += gridDim.x * NTS) rp[i] = 0.f;
}
// ===========================================================================
extern "C" __global__ void __launch_bounds__(NTS, 1) kimi_forward(Params P) {
cg::grid_group grid = cg::this_grid();
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int gw = blockIdx.x * NWARP + warp;
const int nw = gridDim.x * NWARP;
const int cta = blockIdx.x;
extern __shared__ float sm[];
float *sx = sm + O_X;
float *shn = sm + O_HN;
float *saux = sm + O_AUX;
float *ssc = sm + O_SC;
__shared__ float red[NWARP];
__shared__ float scal[2];
__shared__ float smrg[64];
__shared__ float scs[64];
__shared__ float qred[4][128][5];
float *skrq = &qred[0][0][0];
for (int b = 0; b < 4; ++b) {
const bool is_k = (b < 3);
const long long *wd = P.wdesc + (size_t)b * 12 * 8;
// ---------------- prologue: x, xn in smem ---------------------------
{
const __nv_bfloat16 *nm = (b == 3 ? P.moe_norm : P.attn_norm) + b * D;
float ss = 0.f;
for (int i = tid; i < D; i += NTS) {
float v = (b == 0) ? b2f(P.hidden_in[i])
: b2f(__float2bfloat16(P.hb[i] + P.moe_out[(size_t)(b - 1) * D + i]));
P.xb[i] = v;
sx[i] = v;
ss = fmaf(v, v, ss);
}
block_sum(ss, red, scal);
const float rstd = rsqrtf(scal[0] / (float)D + EPS);
for (int i = tid; i < D; i += NTS)
sx[i] = b2f(__float2bfloat16(sx[i] * rstd * b2f(nm[i])));
// materialise a freshly fed MLA cache into the padded buffers, split
// across the three KDA blocks so it is long done before the MLA stage.
if (P.ncopy > 0 && b < 3) {
const int p0 = (int)((long long)P.ncopy * b / 3);
const int p1 = (int)((long long)P.ncopy * (b + 1) / 3);
const int n4 = (p1 - p0) * 64;
const int k4 = (p1 - p0) * 8;
const int t0 = cta * NTS + tid;
const int stride = gridDim.x * NTS;
for (int i = t0; i < n4; i += stride)
((uint4 *)P.ckv_pad)[p0 * 64 + i] = ((const uint4 *)P.src_ckv)[p0 * 64 + i];
for (int i = t0; i < k4; i += stride)
((uint4 *)P.kr_pad)[p0 * 8 + i] = ((const uint4 *)P.src_kr)[p0 * 8 + i];
}
if (cta == 0) {
for (int i = tid; i < D; i += NTS) P.attn_out[i] = 0.f;
for (int i = tid; i < NALL * MOED; i += NTS) { P.hhg[i] = 0.f; P.hhu[i] = 0.f; }
for (int i = tid; i < D; i += NTS) P.moe_out[(size_t)((b + 1) & 3) * D + i] = 0.f;
if (b == 2) {
for (int i = tid; i < 6144; i += NTS) P.mlab[i] = 0.f;
for (int i = tid; i < 576; i += NTS) P.mkv[i] = 0.f;
for (int i = tid; i < NH * 512; i += NTS) P.cw[i] = 0.f;
for (int i = tid; i < 4096; i += NTS) P.mob[i] = 0.f;
}
}
__syncthreads();
}
if (is_k) {
// ================= KDA ==========================================
// K1: q,k,v,g GEMV + fused conv / activation
for (int t = gw; t < 4 * NH * 18 * 2; t += nw) {
const int half = t & 1;
int rest = t >> 1;
const int g = rest % 18;
rest /= 18;
const int head = rest % NH;
const int proj = rest / NH;
WD w = wd_load(wd + proj * 8);
gemv_add<4>(w.wq, w.cols, w.sc, w.zo, g, g * 64 + half * 32, 32,
head * 128, sx, 1.0f, P.raw + (size_t)proj * CTOT);
}
if (cta == 0 && warp < 4) {
const __nv_bfloat16 *bw = P.beta_w + (size_t)b * NH * D;
float acc[8];
#pragma unroll
for (int j = 0; j < 8; ++j) acc[j] = 0.f;
#pragma unroll 4
for (int k = lane; k < D; k += 32) {
const uint4 u4 = *(const uint4 *)(bw + (size_t)k * NH + warp * 8);
const __nv_bfloat16 *pp = (const __nv_bfloat16 *)&u4;
const float xk = sx[k];
#pragma unroll
for (int j = 0; j < 8; ++j) acc[j] = fmaf(b2f(pp[j]), xk, acc[j]);
}
#pragma unroll
for (int j = 0; j < 8; ++j)
#pragma unroll
for (int o = 16; o > 0; o >>= 1) acc[j] += __shfl_xor_sync(0xffffffffu, acc[j], o);
if (lane == 0)
for (int j = 0; j < 8; ++j) P.betab[warp * 8 + j] = sigmoid_(acc[j]);
}
grid.sync();
// Conv / activation. This used to run inside the GEMV loop, on whichever
// warp happened to land the 36th row-slab of a (proj, head); that forced a
// __threadfence + counter + poll per task, which cost more than the work
// it gated. After the barrier the accumulations are final, so one warp
// per (proj, head) can finish the whole thing with no ordering at all.
for (int u = gw; u < 4 * NH; u += nw) {
const int head = u % NH;
const int proj = u / NH;
float *dst = P.raw + (size_t)proj * CTOT;
const int c0 = head * 128 + 4 * lane;
if (proj == 3) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
const float v = ldcg(&dst[c0 + j]);
dst[c0 + j] = -softplus_(b2f(__float2bfloat16(v)));
}
} else {
__nv_bfloat16 *stt = (b == 0 ? (proj == 0 ? P.cq0 : proj == 1 ? P.ck0 : P.cv0)
: b == 1 ? (proj == 0 ? P.cq1 : proj == 1 ? P.ck1 : P.cv1)
: (proj == 0 ? P.cq2 : proj == 1 ? P.ck2 : P.cv2));
const __nv_bfloat16 *cwp = P.conv_w + ((size_t)b * 3 + proj) * CTOT * 4;
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int c = c0 + j;
const float v0 = b2f(stt[c]);
const float v1 = b2f(stt[CTOT + c]);
const float v2 = b2f(stt[2 * CTOT + c]);
const float val = b2f(__float2bfloat16(ldcg(&dst[c])));
const float acc = b2f(cwp[c * 4 + 0]) * v0 + b2f(cwp[c * 4 + 1]) * v1 +
b2f(cwp[c * 4 + 2]) * v2 + b2f(cwp[c * 4 + 3]) * val;
dst[c] = b2f(__float2bfloat16(silu_(acc)));
stt[c] = __float2bfloat16(v1);
stt[CTOT + c] = __float2bfloat16(v2);
stt[2 * CTOT + c] = __float2bfloat16(val);
}
}
}
// K2 below reads P.raw for every head, and this pass writes P.raw from a
// different CTA than the one that will read a given head, so the barrier
// has to stay grid-wide.
grid.sync();
// K2: recurrent update. The state is S[i][j] and the output contracts
// over i (reference: `o = (S * q[:,:,None]).sum(1)`), so
// pred[j] = sum_i d[i]*k[i]*S_old[i][j]
// o[j] = sum_i q[i]*d[i]*S_old[i][j] + (v[j]-pred[j])*sum_i q[i]*b[i]*k[i]
// S_new[i][j] = S_old[i][j]*d[i] + b[i]*k[i]*(v[j]-pred[j])
// Both reductions share the same S reads, so K2a does them together.
{
float *S = (b == 0 ? P.S0 : b == 1 ? P.S1 : P.S2);
const float *rq = P.raw;
const float *rk = P.raw + CTOT;
const float *rg = P.raw + 3 * CTOT;
// K2a: S^T * (d*k) and S^T * (q*d), 4 independent i-chunks per head
for (int t = gw; t < NH * 4; t += nw) {
const int head = t >> 2;
const int ic = t & 3;
const float *Srow = S + (size_t)head * DK * DK + ic * 32 * DK + 4 * lane;
const float *rqh = rq + head * DK + ic * 32;
const float *rkh = rk + head * DK + ic * 32;
const float *rgh = rg + head * DK + ic * 32;
float ap0 = 0.f, ap1 = 0.f, ap2 = 0.f, ap3 = 0.f;
float ao0 = 0.f, ao1 = 0.f, ao2 = 0.f, ao3 = 0.f;
#pragma unroll 4
for (int r = 0; r < 32; ++r) {
const uint4 sv = *(const uint4 *)(Srow + (size_t)r * DK);
const float d = __expf(rgh[r]);
const float dk = d * rkh[r];
const float qd = d * rqh[r];
const float *fv = (const float *)&sv;
ap0 = fmaf(fv[0], dk, ap0); ao0 = fmaf(fv[0], qd, ao0);
ap1 = fmaf(fv[1], dk, ap1); ao1 = fmaf(fv[1], qd, ao1);
ap2 = fmaf(fv[2], dk, ap2); ao2 = fmaf(fv[2], qd, ao2);
ap3 = fmaf(fv[3], dk, ap3); ao3 = fmaf(fv[3], qd, ao3);
}
float *pp = P.pred + (head * 4 + ic) * DK + 4 * lane;
float *po = P.pred + POPRE + (head * 4 + ic) * DK + 4 * lane;
pp[0] = ap0; pp[1] = ap1; pp[2] = ap2; pp[3] = ap3;
po[0] = ao0; po[1] = ao1; po[2] = ao2; po[3] = ao3;
}
grid.sync();
// K2b: fold the four i-chunks, and the per-head scalar
// sum_i q[i]*beta[i]*k[i].
for (int t = gw; t < NH * DK; t += nw) {
const int head = t >> 7;
const int j = t & 127;
float sp = 0.f, so = 0.f;
#pragma unroll
for (int ic = 0; ic < 4; ++ic) {
sp += P.pred[(head * 4 + ic) * DK + j];
so += P.pred[POPRE + (head * 4 + ic) * DK + j];
}
P.pred[FPRED + head * DK + j] = sp;
P.pred[FOPRE + head * DK + j] = so;
}
for (int t = gw; t < NH; t += nw) {
const int head = t;
const float *rqh = rq + head * DK;
const float *rkh = rk + head * DK;
float s = 0.f;
#pragma unroll
for (int l = 0; l < 4; ++l) {
const int i = lane + 32 * l;
s = fmaf(rqh[i], P.betab[head] * rkh[i], s);
}
#pragma unroll
for (int o = 16; o > 0; o >>= 1) s += __shfl_xor_sync(0xffffffffu, s, o);
if (lane == 0) P.pred[BQKO + head] = s;
}
grid.sync();
// K2c: emit o, then advance the state row by row (coalesced on j).
for (int t = gw; t < NH * DK; t += nw) {
const int head = t >> 7;
const int j = t & 127;
const float vj = rq[2 * CTOT + head * DK + j];
const float o = (P.pred[FOPRE + head * DK + j] +
(vj - P.pred[FPRED + head * DK + j]) * P.pred[BQKO + head]) *
QSCALE;
P.ob[head * DK + j] = b2f(__float2bfloat16(o));
}
for (int t = gw; t < NH * DK; t += nw) {
const int head = t >> 7;
const int i = t & 127;
float *Srow = S + (size_t)(head * DK + i) * DK;
uint4 s4 = *(const uint4 *)(Srow + 4 * lane);
const float d = __expf(rg[head * DK + i]);
const float bkr = P.betab[head] * rk[head * DK + i];
const float *rv = rq + 2 * CTOT + head * DK + 4 * lane;
const float *pp = P.pred + FPRED + head * DK + 4 * lane;
#pragma unroll
for (int j = 0; j < 4; ++j)
((float *)&s4)[j] = fmaf(((const float *)&s4)[j], d, bkr * (rv[j] - pp[j]));
*(uint4 *)(Srow + 4 * lane) = s4;
}
}
grid.sync();
// K3: o_proj (4096 -> 2304)
for (int i = tid; i < CTOT; i += NTS) saux[i] = P.ob[i];
__syncthreads();
{
WD w = wd_load(wd + 4 * 8);
for (int t = gw; t < 18 * 32 * 2; t += nw) {
const int half = t & 1;
const int rest = t >> 1;
const int g = rest % 32;
const int tile = rest / 32;
gemv_add<4>(w.wq, w.cols, w.sc, w.zo, g, g * 64 + half * 32, 32,
tile * 128, saux, 1.0f, P.attn_out);
}
}
grid.sync();
// K4: h, rmsnorm(moe_norm), router, topk
{
const __nv_bfloat16 *nm = P.moe_norm + b * D;
float ss = 0.f;
for (int i = tid; i < D; i += NTS) {
const float h = b2f(__float2bfloat16(P.xb[i] + P.attn_out[i]));
P.hb[i] = h;
ss = fmaf(h, h, ss);
}
block_sum(ss, red, scal);
const float rstd = rsqrtf(scal[0] / (float)D + EPS);
for (int i = tid; i < D; i += NTS)
shn[i] = b2f(__float2bfloat16(P.hb[i] * rstd * b2f(nm[i])));
}
__syncthreads();
router_matvec(P, b, shn, gw, nw, lane);
grid.sync();
router_finalize(P, b, warp, lane);
grid.sync();
moe_body(P, b, wd, shn, saux, P.moe_out + (size_t)b * D, grid, gw, nw, tid);
grid.sync();
} else {
// ================= MLA ===========================================
// M1: q_proj + kv_a
for (int t = gw; t < 96 * 18 * 2; t += nw) {
const int half = t & 1;
int rest = t >> 1;
const int g = rest % 18;
const int tile = rest / 18;
WD w = wd_load(wd);
gemv_add<2>(w.wq, w.cols, w.sc, w.zo, g, g * 64 + half * 32, 32,
tile * 64, sx, 1.0f, P.mlab);
__threadfence();
if (lane == 0) atomicAdd(&P.cnt[128 + tile / 3], 1);
}
for (int t = gw; t < 9 * 18 * 2; t += nw) {
const int half = t & 1;
int rest = t >> 1;
const int g = rest % 18;
const int tile = rest / 18;
WD w = wd_load(wd + 8);
gemv_add<2>(w.wq, w.cols, w.sc, w.zo, g, g * 64 + half * 32, 32,
tile * 64, sx, 1.0f, P.mkv);
__threadfence();
int old = 0;
if (lane == 0) old = atomicAdd(&P.cnt[224], 1);
old = __shfl_sync(0xffffffffu, old, 0);
if (old == 323) {
for (int i = lane; i < 512; i += 32)
P.ckv_pad[(size_t)P.pos * 512 + i] = __float2bfloat16(P.mkv[i]);
if (lane == 0) P.cnt[224] = 0;
}
}
if (cta < NH) {
if (tid == 0) {
volatile int *pp = (volatile int *)&P.cnt[128 + cta];
while (*pp < 108) {}
}
__syncthreads();
WD wk = wd_load(wd + 2 * 8);
const unsigned char *Wb = wk.wq;
// qabs[head][c] = sum_k q_nope[k] * Wb[c][head*256+k], with the weight
// dequantized as (nib - z[g][n]) * s[g][n]. The channel c = 4*cb + j
// sits inside quant group g = c>>7 = cb>>5 (4*cb never straddles a
// 128-row boundary), and the param column is n = head*256 + k, so s/z
// are common to all four accumulators of a thread: fold them in as
// acc_j = sum_k q*s*nib_j and ar = sum_k q*s*z.
// The fmaf chain walks k in the same order as before -- the attention
// softmax amplifies any reassociation of this dot product into a much
// larger output error at long context -- but now pulls the packed bytes
// four at a time, so a warp issues 8 loads instead of 32. Each of those
// still replays over 32 sectors (neighbouring cb are 16 KB apart).
const int cb = tid & 127;
const int kg = tid >> 7;
const __nv_bfloat16 *sp = wk.sc + (size_t)(cb >> 5) * 8192 + cta * 256 + 0;
const __nv_bfloat16 *zp = wk.zo + (size_t)(cb >> 5) * 8192 + cta * 256 + 0;
const float *qn = P.mlab + 192 * cta;
float acc0 = 0.f, acc1 = 0.f, acc2 = 0.f, acc3 = 0.f, ar = 0.f;
const unsigned char *r0 = Wb + (size_t)(2 * cb) * 8192 + cta * 256 + kg * 32;
const unsigned char *r1 = r0 + 8192;
// Four packed bytes per load instead of one: with neighbouring cb 16 KB
// apart each warp load still replays over 32 sectors, so the win is the
// 4x fewer load instructions. k is walked in the same order as before
// because the attention softmax magnifies any reassociation here.
#pragma unroll
for (int u = 0; u < 8; ++u) {
const unsigned int w0 = ((const unsigned int *)r0)[u];
const unsigned int w1 = ((const unsigned int *)r1)[u];
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int k = kg * 32 + 4 * u + j;
const float q = b2f(__float2bfloat16(ldcg(&qn[k])));
const float sv = b2f(sp[k]);
const float zv = b2f(zp[k]);
acc0 = fmaf((float)((w0 >> (8 * j)) & 0xF) * sv, q, acc0);
acc1 = fmaf((float)((w0 >> (8 * j + 4)) & 0xF) * sv, q, acc1);
acc2 = fmaf((float)((w1 >> (8 * j)) & 0xF) * sv, q, acc2);
acc3 = fmaf((float)((w1 >> (8 * j + 4)) & 0xF) * sv, q, acc3);
ar = fmaf(zv * sv, q, ar);
}
}
qred[kg][cb][0] = acc0;
qred[kg][cb][1] = acc1;
qred[kg][cb][2] = acc2;
qred[kg][cb][3] = acc3;
qred[kg][cb][4] = ar;
__syncthreads();
if (kg == 0) {
const float ars = qred[0][cb][4] + qred[1][cb][4] +
qred[2][cb][4] + qred[3][cb][4];
for (int j = 0; j < 4; ++j) {
const float s = qred[0][cb][j] + qred[1][cb][j] + qred[2][cb][j] + qred[3][cb][j];
// stride NCH so the rope half can live at [512, 576) of the same row
P.qabs_g[cta * NCH + 4 * cb + j] = __float2bfloat16(s - ars);
}
}
if (tid == 0) P.cnt[128 + cta] = 0;
}
grid.sync();
// M2: rope + staging + scores
const int ntok = P.pos + 1;
const int lr = min(((ntok + P.grid - 1) / P.grid + 3) & ~3, LRMAX);
const int l0 = cta * lr;
const int nrow = max(0, min(ntok, l0 + lr) - l0);
{
if (tid < 32) {
const float ang = (float)P.pos * exp2f(-13.287712379549449f * (float)tid / 32.0f);
scs[2 * tid] = cosf(ang);
scs[2 * tid + 1] = sinf(ang);
}
__syncthreads();
if (tid < 64) {
const int t = tid >> 1;
const float a = b2f(__float2bfloat16(P.mkv[512 + 2 * t]));
const float bb = b2f(__float2bfloat16(P.mkv[512 + 2 * t + 1]));
const float ct = scs[2 * t], st = scs[2 * t + 1];
const float r = (tid & 1) ? bb * ct + a * st : a * ct - bb * st;
// keep it bf16-rounded so this row matches what lands in kr_pad
skrq[tid] = b2f(__float2bfloat16(r));
if (cta == 0) P.kr_pad[(size_t)P.pos * 64 + tid] = __float2bfloat16(r);
}
for (int i = tid; i < NH * 32; i += NTS) {
const int h = i >> 5, t = i & 31;
const float a = b2f(__float2bfloat16(P.mlab[192 * h + 128 + 2 * t]));
const float bb = b2f(__float2bfloat16(P.mlab[192 * h + 128 + 2 * t + 1]));
const float ct = scs[2 * t], st = scs[2 * t + 1];
P.qabs_g[h * NCH + 512 + 2 * t] = __float2bfloat16(a * ct - bb * st);
P.qabs_g[h * NCH + 512 + 2 * t + 1] = __float2bfloat16(bb * ct + a * st);
}
__syncthreads();
// The query row's own k_rope was written to kr_pad by CTA 0 this stage
// and is not ordered against the other CTAs' reads, so keep it in smem.
const int rpos = (P.pos >= l0 && P.pos < l0 + nrow) ? (P.pos - l0) : -1;
{
const int th = tid / 22;
const int tl = tid % 22;
if (th < 8 && tl < nrow) {
const int h0 = th * 4;
// Rows stride by the 22 tl-slots rather than blocking four
// consecutive rows per thread: at short contexts (nrow=12) the
// blocked form left only 24 of 512 threads doing the score dots.
const int lt = tl;
float acc[4][4];
#pragma unroll
for (int i = 0; i < 4; ++i)
#pragma unroll
for (int k = 0; k < 4; ++k) acc[i][k] = 0.f;
// dot the 512 latent channels straight out of the padded global cache
for (int c0 = 0; c0 < 512; c0 += 8) {
float cv[4][8], qv[4][8];
#pragma unroll
for (int i = 0; i < 4; ++i) {
const int l = lt + i * 22;
if (l < nrow) {
const uint4 u =
*(const uint4 *)(P.ckv_pad + (size_t)(l0 + l) * 512 + c0);
const __nv_bfloat16 *pp = (const __nv_bfloat16 *)&u;
#pragma unroll
for (int j = 0; j < 8; ++j) cv[i][j] = b2f(pp[j]);
} else {
#pragma unroll
for (int j = 0; j < 8; ++j) cv[i][j] = 0.f;
}
}
#pragma unroll
for (int k = 0; k < 4; ++k) {
const uint4 u = *(const uint4 *)(P.qabs_g + (h0 + k) * NCH + c0);
const __nv_bfloat16 *pp = (const __nv_bfloat16 *)&u;
#pragma unroll
for (int j = 0; j < 8; ++j) qv[k][j] = b2f(pp[j]);
}
#pragma unroll
for (int i = 0; i < 4; ++i)
#pragma unroll
for (int k = 0; k < 4; ++k)
#pragma unroll
for (int j = 0; j < 8; ++j)
acc[i][k] = fmaf(cv[i][j], qv[k][j], acc[i][k]);
}
// ... then the 64 decoupled-rope channels
for (int c8 = 0; c8 < 8; ++c8) {
float cv[4][8], qv[4][8];
#pragma unroll
for (int i = 0; i < 4; ++i) {
const int l = lt + i * 22;
if (l < nrow) {
if (l == rpos) {
const float *sp = skrq + c8 * 8;
#pragma unroll
for (int j = 0; j < 8; ++j) cv[i][j] = sp[j];
} else {
const uint4 u = *(const uint4 *)(P.kr_pad + (size_t)(l0 + l) * 64 + c8 * 8);
const __nv_bfloat16 *pp = (const __nv_bfloat16 *)&u;
#pragma unroll
for (int j = 0; j < 8; ++j) cv[i][j] = b2f(pp[j]);
}
} else {
#pragma unroll
for (int j = 0; j < 8; ++j) cv[i][j] = 0.f;
}
}
#pragma unroll
for (int k = 0; k < 4; ++k) {
const uint4 u = *(const uint4 *)(P.qabs_g + (h0 + k) * NCH + 512 + c8 * 8);
const __nv_bfloat16 *pp = (const __nv_bfloat16 *)&u;
#pragma unroll
for (int j = 0; j < 8; ++j) qv[k][j] = b2f(pp[j]);
}
#pragma unroll
for (int i = 0; i < 4; ++i)
#pragma unroll
for (int k = 0; k < 4; ++k)
#pragma unroll
for (int j = 0; j < 8; ++j)
acc[i][k] = fmaf(cv[i][j], qv[k][j], acc[i][k]);
}
#pragma unroll
for (int i = 0; i < 4; ++i)
if (lt + i * 22 < nrow)
#pragma unroll
for (int k = 0; k < 4; ++k)
ssc[(h0 + k) * LRMAX + lt + i * 22] = acc[i][k] * SCALE;
}
}
__syncthreads();
for (int q = 0; q < 2; ++q) {
const int hd = warp * 2 + q;
float m = -1e30f;
for (int l = lane; l < nrow; l += 32) m = fmaxf(m, ssc[hd * LRMAX + l]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) m = fmaxf(m, __shfl_xor_sync(0xffffffffu, m, o));
float s = 0.f;
for (int l = lane; l < nrow; l += 32) s += __expf(ssc[hd * LRMAX + l] - m);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) s += __shfl_xor_sync(0xffffffffu, s, o);
if (lane == 0) {
// head-major partials: the M3 merge reads these with the CTA index
// in the fastest-varying lane position, so keep it contiguous.
P.pm[hd * P.grid + cta] = m;
P.ps[hd * P.grid + cta] = s;
}
}
}
grid.sync();
// M3: merge -> p -> cw
{
for (int q = 0; q < 2; ++q) {
const int hd = warp * 2 + q;
float m = -1e30f;
for (int j = lane; j < P.grid; j += 32) m = fmaxf(m, P.pm[hd * P.grid + j]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) m = fmaxf(m, __shfl_xor_sync(0xffffffffu, m, o));
float s = 0.f;
for (int j = lane; j < P.grid; j += 32)
s += P.ps[hd * P.grid + j] * __expf(P.pm[hd * P.grid + j] - m);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) s += __shfl_xor_sync(0xffffffffu, s, o);
if (lane == 0) {
smrg[hd] = m;
smrg[32 + hd] = 1.0f / (s + 1e-30f);
}
}
__syncthreads();
for (int i = tid; i < NH * LRMAX; i += NTS) {
const int hd = i / LRMAX, l = i % LRMAX;
ssc[i] = (l < nrow) ? __expf(ssc[i] - smrg[hd]) : 0.f;
}
__syncthreads();
// cw[hd][c] = sum_l p[hd][l] * ckv[l][c]. The old form looped the four
// heads per thread on the outside, so the same c_kv vector was fetched
// four times; it is the only global traffic in this phase. Hoisting the
// load keeps all four heads' accumulators live instead. ssc is now read
// by every lane of a warp at one address, i.e. a broadcast, not a
// bank conflict.
const int hg = tid >> 6;
const int cg = tid & 63;
float acc[4][8];
#pragma unroll
for (int q = 0; q < 4; ++q)
#pragma unroll
for (int j = 0; j < 8; ++j) acc[q][j] = 0.f;
for (int l = 0; l < nrow; ++l) {
const uint4 ca =
*(const uint4 *)(P.ckv_pad + (size_t)(l0 + l) * 512 + cg * 8);
const __nv_bfloat16 *pca = (const __nv_bfloat16 *)&ca;
float fv[8];
#pragma unroll
for (int j = 0; j < 8; ++j) fv[j] = b2f(pca[j]);
#pragma unroll
for (int q = 0; q < 4; ++q) {
const float pv = ssc[(hg * 4 + q) * LRMAX + l];
#pragma unroll
for (int j = 0; j < 8; ++j) acc[q][j] = fmaf(pv, fv[j], acc[q][j]);
}
}
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int hd = hg * 4 + q;
#pragma unroll
for (int j = 0; j < 8; ++j) atomicAdd(&P.cw[hd * 512 + cg * 8 + j], acc[q][j]);
}
}
grid.sync();
// M4: output absorption
{
WD w = wd_load(wd + 2 * 8);
for (int t = gw; t < NH * 16; t += nw) {
const int h = t >> 4;
const int gp = t & 15;
const int cb0 = gp * 16;
// This task covers kv_b input rows [32*gp, 32*gp+32), all inside the
// single quant group gg = gp>>2, so s/z depend only on the output
// column: fold them in as out = s*(acc - z*sum_k cw[k]).
const int gg = gp >> 2;
const int cn = h * 256 + 128 + 4 * lane;
const __nv_bfloat16 *sp = w.sc + (size_t)gg * 8192 + cn;
const __nv_bfloat16 *zp = w.zo + (size_t)gg * 8192 + cn;
const float inv = smrg[32 + h];
float sv[4], zv[4], acc[4];
float ar = 0.f;
#pragma unroll
for (int j = 0; j < 4; ++j) {
sv[j] = b2f(sp[j]);
zv[j] = b2f(zp[j]);
acc[j] = 0.f;
}
#pragma unroll 2
for (int cb = cb0; cb < cb0 + 16; ++cb) {
const float cl = P.cw[h * 512 + 2 * cb] * inv;
const float ch2 = P.cw[h * 512 + 2 * cb + 1] * inv;
const unsigned int bv = *(const unsigned int *)(w.wq + (size_t)cb * 8192 + cn);
acc[0] = fmaf((float)(bv & 0xF), cl, acc[0]);
acc[0] = fmaf((float)((bv >> 4) & 0xF), ch2, acc[0]);
acc[1] = fmaf((float)((bv >> 8) & 0xF), cl, acc[1]);
acc[1] = fmaf((float)((bv >> 12) & 0xF), ch2, acc[1]);
acc[2] = fmaf((float)((bv >> 16) & 0xF), cl, acc[2]);
acc[2] = fmaf((float)((bv >> 20) & 0xF), ch2, acc[2]);
acc[3] = fmaf((float)((bv >> 24) & 0xF), cl, acc[3]);
acc[3] = fmaf((float)((bv >> 28) & 0xF), ch2, acc[3]);
ar += cl + ch2;
}
#pragma unroll
for (int j = 0; j < 4; ++j)
atomicAdd(&P.mob[h * 128 + 4 * lane + j], (acc[j] - zv[j] * ar) * sv[j]);
}
}
grid.sync();
// M5: o_proj (4096 -> 2304)
for (int i = tid; i < CTOT; i += NTS)
sm[O_M5X + i] = b2f(__float2bfloat16(P.mob[i]));
__syncthreads();
{
WD w = wd_load(wd + 3 * 8);
for (int t = gw; t < 18 * 32 * 2; t += nw) {
const int half = t & 1;
const int rest = t >> 1;
const int g = rest % 32;
const int tile = rest / 32;
gemv_add<4>(w.wq, w.cols, w.sc, w.zo, g, g * 64 + half * 32, 32,
tile * 128, sm + O_M5X, 1.0f, P.attn_out);
}
}
grid.sync();
// M6: h, rmsnorm, router, topk
{
const __nv_bfloat16 *nm = P.moe_norm + b * D;
float ss = 0.f;
for (int i = tid; i < D; i += NTS) {
const float h = b2f(__float2bfloat16(P.xb[i] + P.attn_out[i]));
P.hb[i] = h;
ss = fmaf(h, h, ss);
}
block_sum(ss, red, scal);
const float rstd = rsqrtf(scal[0] / (float)D + EPS);
for (int i = tid; i < D; i += NTS)
shn[i] = b2f(__float2bfloat16(P.hb[i] * rstd * b2f(nm[i])));
}
__syncthreads();
router_matvec(P, b, shn, gw, nw, lane);
grid.sync();
router_finalize(P, b, warp, lane);
grid.sync();
moe_body(P, 3, wd, shn, saux, P.moe_out + 3 * D, grid, gw, nw, tid);
grid.sync();
}
}
for (int i = tid; i < D; i += NTS)
P.hidden_out[i] = __float2bfloat16(P.hb[i] + P.moe_out[3 * D + i]);
}
// ===========================================================================
// host side
// ===========================================================================
#include <cstdio>
#include <cstring>
static char g_err[256];
extern "C" const char *kimi_error() { return g_err; }
extern "C" int kimi_smem_bytes() { return SMEM_TOTAL * 4; }
extern "C" int kimi_threads() { return NTS; }
extern "C" int kimi_launch(const long long *p, int grid_want) {
g_err[0] = 0;
Params P;
memset(&P, 0, sizeof(P));
P.xb = (float *)p[0];
P.raw = (float *)p[1];
P.ob = (float *)p[2];
P.attn_out = (float *)p[3];
P.hb = (float *)p[4];
P.hhg = (float *)p[5];
P.hhu = (float *)p[6];
P.moe_out = (float *)p[7];
P.router = (float *)p[8];
P.betab = (float *)p[9];
P.selw = (float *)p[10];
P.sel = (int *)p[11];
P.cnt = (int *)p[12];
P.cw = (float *)p[13];
P.pm = (float *)p[14];
P.ps = (float *)p[15];
P.mlab = (float *)p[16];
P.mkv = (float *)p[17];
P.mob = (float *)p[18];
P.pred = (float *)p[48];
P.qabs_g = (__nv_bfloat16 *)p[19];
P.ckv_pad = (__nv_bfloat16 *)p[20];
P.kr_pad = (__nv_bfloat16 *)p[21];
P.wdesc = (const long long *)p[22];
P.attn_norm = (const __nv_bfloat16 *)p[23];
P.moe_norm = (const __nv_bfloat16 *)p[24];
P.beta_w = (const __nv_bfloat16 *)p[25];
P.router_w = (const __nv_bfloat16 *)p[26];
P.conv_w = (const __nv_bfloat16 *)p[27];
P.S0 = (float *)p[28];
P.S1 = (float *)p[29];
P.S2 = (float *)p[30];
P.cq0 = (__nv_bfloat16 *)p[31];
P.ck0 = (__nv_bfloat16 *)p[32];
P.cv0 = (__nv_bfloat16 *)p[33];
P.cq1 = (__nv_bfloat16 *)p[34];
P.ck1 = (__nv_bfloat16 *)p[35];
P.cv1 = (__nv_bfloat16 *)p[36];
P.cq2 = (__nv_bfloat16 *)p[37];
P.ck2 = (__nv_bfloat16 *)p[38];
P.cv2 = (__nv_bfloat16 *)p[39];
P.hidden_in = (const __nv_bfloat16 *)p[40];
P.hidden_out = (__nv_bfloat16 *)p[41];
P.src_ckv = (const __nv_bfloat16 *)p[42];
P.src_kr = (const __nv_bfloat16 *)p[43];
P.pos = (int)p[44];
P.grid = (int)p[45];
P.ncopy = (int)p[47];
const int smem = SMEM_TOTAL * 4;
static int smem_ok = 0;
if (!smem_ok) {
cudaError_t e = cudaFuncSetAttribute(
(const void *)kimi_forward, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
if (e != cudaSuccess) {
snprintf(g_err, sizeof(g_err), "cudaFuncSetAttribute: %s", cudaGetErrorString(e));
return 1;
}
smem_ok = 1;
}
int grid = grid_want;
if (grid <= 0) {
int dev = 0;
cudaGetDevice(&dev);
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
grid = prop.multiProcessorCount;
}
void *args[] = {&P};
cudaError_t e = cudaLaunchCooperativeKernel((const void *)kimi_forward, dim3(grid),
dim3(NTS), args, smem, nullptr);
if (e != cudaSuccess) {
snprintf(g_err, sizeof(g_err), "launch: %s", cudaGetErrorString(e));
return 2;
}
return 0;
}
"""
# ---8<--- CUDA_END
CPP_SRC = r"""
#include <torch/extension.h>
#include <cstdint>
extern "C" int kimi_launch(const long long *p, int grid);
extern "C" const char *kimi_error();
extern "C" int kimi_smem_bytes();
extern "C" int kimi_threads();
int64_t kimi_run(torch::Tensor ptrs, int64_t grid) {
auto c = ptrs.contiguous();
return (int64_t)kimi_launch((const long long *)c.data_ptr<int64_t>(), (int)grid);
}
std::string kimi_last_error() { return std::string(kimi_error()); }
int64_t kimi_smem() { return (int64_t)kimi_smem_bytes(); }
int64_t kimi_nthreads() { return (int64_t)kimi_threads(); }
"""
def _ensure_ninja_on_path():
"""torch's load_inline shells out to `ninja` by bare name; the venv it ships
in is not always on PATH, so put its bin dir there ourselves."""
import os
import shutil
if shutil.which("ninja"):
return
try:
import ninja # noqa: F401 (vendored build tool, not a banned lib)
bindir = getattr(ninja, "BIN_DIR", None)
if bindir and os.path.isdir(bindir):
os.environ["PATH"] = bindir + os.pathsep + os.environ.get("PATH", "")
except Exception:
pass
def _load_module():
_ensure_ninja_on_path()
return load_inline(
name="kimi_linear_megakernel",
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=["kimi_run", "kimi_last_error", "kimi_smem", "kimi_nthreads"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
# --------------------------------------------------------------------------- #
# W4A16 storage (names must mirror reference.py exactly)
# --------------------------------------------------------------------------- #
class QuantLinear(nn.Module):
def __init__(self, in_f, out_f, group=GROUP_SIZE):
super().__init__()
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):
w = torch.empty(self.in_f, self.out_f, dtype=torch.uint8, device=self.w_q.device)
w[0::2] = self.w_q & 0xF
w[1::2] = (self.w_q >> 4) & 0xF
z = self.zeros.repeat_interleave(self.group, 0).float()
s = self.scales.repeat_interleave(self.group, 0).float()
return ((w.float() - z) * s).to(torch.bfloat16)
def forward(self, x):
return (x.float() @ self.weight_bf().float()).to(torch.bfloat16)
class QuantExperts(nn.Module):
def __init__(self, n, in_f, out_f, group=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):
w = torch.empty(self.in_f, self.out_f, dtype=torch.uint8, device=self.w_q.device)
w[0::2] = self.w_q[e] & 0xF
w[1::2] = (self.w_q[e] >> 4) & 0xF
z = self.zeros[e].repeat_interleave(self.group, 0).float()
s = self.scales[e].repeat_interleave(self.group, 0).float()
return ((w.float() - z) * s).to(torch.bfloat16)
def _rmsnorm(x, w):
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, dim, theta, 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, cos, sin):
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
# --------------------------------------------------------------------------- #
class KDA(nn.Module):
def __init__(self, cfg):
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):
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):
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, gq, uq, dq, e):
h = F.silu(x.float() @ gq.weight_bf(e).float()) * (x.float() @ uq.weight_bf(e).float())
return h @ dq.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, kind):
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))
# --------------------------------------------------------------------------- #
# Model
# --------------------------------------------------------------------------- #
_SLOT_K = {"q_proj": 0, "k_proj": 1, "v_proj": 2, "g_proj": 3, "o_proj": 4,
"gate": 5, "up": 6, "down": 7, "s_gate": 8, "s_up": 9, "s_down": 10}
_SLOT_M = {"q_proj": 0, "kv_a": 1, "kv_b": 2, "o_proj": 3,
"gate": 5, "up": 6, "down": 7, "s_gate": 8, "s_up": 9, "s_down": 10}
_SMEM_TOTAL = 20992 # floats of dynamic smem the fused kernel asks for
_ELEM = {"float32": 4, "bfloat16": 2, "uint8": 1,
"torch.float32": 4, "torch.bfloat16": 2, "torch.uint8": 1}
class Model(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern)
self._ok = False
self._src = None
def step(self, hidden, state):
if not self._ok:
# First call: build the scratch buffers, weight descriptors and the
# fused CUDA module. If anything in that setup fails we degrade to
# the (correct but slow) eager reference path instead of crashing.
try:
self._prepare()
except Exception as e:
self._ok = False
self._fallback = f"{type(e).__name__}: {e}"
print(f"solution: fused kernel unavailable ({self._fallback}); "
f"using slow eager path", file=sys.stderr)
if self._ok:
return self._kernel_step(hidden, state)
return self._eager_step(hidden, state)
# -- debug-only eager path -------------------------------------------- #
def _fast_ok(self):
return self._ok
def _eager_step(self, hidden, state):
for i, blk in enumerate(self.blocks):
hidden = blk.step(hidden, state[i])
return hidden, state
# ---------------------------------------------------------------------- #
def _prepare(self):
dev = next(self.parameters()).device
if dev.type != "cuda":
return
self._dev = dev
import numpy as np
mod = _load_module()
self._mod = mod
if int(mod.kimi_smem()) != _SMEM_TOTAL * 4:
return
G = torch.cuda.get_device_properties(dev).multi_processor_count
self._grid = G
cap = G * 96
self._cap = cap
z = lambda n: torch.zeros(n, dtype=torch.float32, device=dev)
zi = lambda n: torch.zeros(n, dtype=torch.int32, device=dev)
zb = lambda n: torch.zeros(n, dtype=torch.bfloat16, device=dev)
self._buf = dict(
xb=z(2304), raw=z(4 * 4096), ob=z(4096), attn_out=z(2304), hb=z(2304),
hhg=z(9 * 1024), hhu=z(9 * 1024), moe_out=z(4 * 2304), router=z(4 * 64 * 32),
betab=z(32), selw=z(8), sel=zi(8), cnt=zi(256), cw=z(32 * 512),
pm=z(G * 32), ps=z(G * 32), mlab=z(6144), mkv=z(576), mob=z(4096), pred=z(41024),
qabs_g=zb(32 * 576),
ckv_pad=torch.zeros(cap, 512, dtype=torch.bfloat16, device=dev),
kr_pad=torch.zeros(cap, 64, dtype=torch.bfloat16, device=dev),
)
B = self._buf
self._pad_ckv = B["ckv_pad"]
self._pad_kr = B["kr_pad"]
self._hout = zb(2304)
self._const = []
# ---- weight descriptors -----------------------------------------
wd = torch.zeros(4, 12, 8, dtype=torch.int64)
kinds = list(self.cfg.pattern)
for b, (blk, kind) in enumerate(zip(self.blocks, kinds)):
table = _SLOT_K if kind == "K" else _SLOT_M
for name, slot in table.items():
if name in ("gate", "up", "down", "s_gate", "s_up", "s_down"):
mod_ = getattr(blk.moe, name)
experts = True
else:
mod_ = getattr(blk.attn, name)
experts = False
w = mod_.w_q
t = wd[b, slot]
t[0] = w.data_ptr()
t[1] = mod_.scales.data_ptr()
t[2] = mod_.zeros.data_ptr()
t[3] = mod_.in_f
t[4] = mod_.out_f
if experts:
t[5] = w.stride(0) * _ELEM[str(w.dtype)]
t[6] = mod_.scales.stride(0)
t[7] = mod_.zeros.stride(0)
self._wdesc = wd.to(dev)
an = torch.stack([b.attn_norm.detach() for b in self.blocks]).contiguous()
mn = torch.stack([b.moe_norm.detach() for b in self.blocks]).contiguous()
# nn.Linear weights are [out, in]; the kernel streams them as [in, out]
# so one contiguous 16-byte load covers 8 consecutive output features.
rw = torch.stack([b.moe.router.weight.detach() for b in self.blocks])
rw = rw.transpose(1, 2).contiguous()
bw = torch.stack([b.attn.beta_proj.weight.detach() for b in self.blocks[:3]])
bw = bw.transpose(1, 2).contiguous()
cw = torch.stack([b.attn.conv_w.detach() for b in self.blocks[:3]]).contiguous()
self._an, self._mn, self._rw, self._bw, self._cw = an, mn, rw, bw, cw
p = np.zeros(64, dtype=np.int64)
p[0] = B["xb"].data_ptr()
p[1] = B["raw"].data_ptr()
p[2] = B["ob"].data_ptr()
p[3] = B["attn_out"].data_ptr()
p[4] = B["hb"].data_ptr()
p[5] = B["hhg"].data_ptr()
p[6] = B["hhu"].data_ptr()
p[7] = B["moe_out"].data_ptr()
p[8] = B["router"].data_ptr()
p[9] = B["betab"].data_ptr()
p[10] = B["selw"].data_ptr()
p[11] = B["sel"].data_ptr()
p[12] = B["cnt"].data_ptr()
p[13] = B["cw"].data_ptr()
p[14] = B["pm"].data_ptr()
p[15] = B["ps"].data_ptr()
p[16] = B["mlab"].data_ptr()
p[17] = B["mkv"].data_ptr()
p[18] = B["mob"].data_ptr()
p[19] = B["qabs_g"].data_ptr()
p[20] = B["ckv_pad"].data_ptr()
p[21] = B["kr_pad"].data_ptr()
p[22] = self._wdesc.data_ptr()
p[23] = an.data_ptr()
p[24] = mn.data_ptr()
p[25] = bw.data_ptr()
p[26] = rw.data_ptr()
p[27] = cw.data_ptr()
p[48] = B["pred"].data_ptr()
p[41] = self._hout.data_ptr()
p[45] = G
self._p = p
self._mla_idx = kinds.index("M")
self._kda_idx = [i for i, k in enumerate(kinds) if k == "K"]
B["cnt"].zero_()
self._ok = True
def _kernel_step(self, hidden, state):
if not self._ok:
self._prepare()
if not self._ok:
return self._eager_step(hidden, state)
G, p, B = self._grid, self._p, self._buf
if not hidden.is_contiguous():
hidden = hidden.contiguous()
p[40] = hidden.data_ptr()
for i in self._kda_idx:
st = state[i]
p[28 + i] = st["S"].data_ptr()
p[31 + 3 * i] = st["cq"].data_ptr()
p[32 + 3 * i] = st["ck"].data_ptr()
p[33 + 3 * i] = st["cv"].data_ptr()
mst = state[self._mla_idx]
ckv = mst["c_kv"]
if ckv.data_ptr() != self._pad_ckv.data_ptr():
src = ckv if ckv.is_contiguous() else ckv.contiguous()
kr = mst["k_rope"]
kr = kr if kr.is_contiguous() else kr.contiguous()
self._src = (src, kr)
p[42] = src.data_ptr()
p[43] = kr.data_ptr()
p[47] = int(ckv.shape[0])
else:
p[47] = 0
pos = int(ckv.shape[0])
p[44] = pos
rc = int(self._mod.kimi_run(torch.from_numpy(p), G))
if rc:
raise RuntimeError(f"megakernel launch failed ({rc}): {self._mod.kimi_last_error()}")
n = pos + 1
mst["c_kv"] = self._pad_ckv[:n]
mst["k_rope"] = self._pad_kr[:n]
return self._hout, state
20260910_084202_deepseek-claude_deepseek-flash_02_kimi_linear_decode