KernelBench mega · H100
Kimi-Linear Decode Kimi K3 (256k)
manually audited: clean
H100 kinetic-0715 mega cell (6.8901x geomean speedup vs eager reference over ctx 2048/8192/16384: 6.18x / 5.33x / 9.92x). Genuine single-launch CUDA megakernel (framework.txt=cuda_raw): one load_inline-built __global__ mega_kernel launched via cudaLaunchCooperativeKernel exactly once per step() (NTHREADS=256, ncta = SM_count * occupancy, software gsync barriers across 18 phases). No torch.cuda.CUDAGraph, no torch.compile, no per-op kernel loop, no getattr/importlib obfuscation, no forbidden libs (transformers/flash_mla/vllm/marlin/bitsandbytes/reference/baseline etc.). Computes the full real fused decode on live inputs and live model state: 3x KDA (fused int4 dequant-GEMV q/k/v/g via fp16 magic nibble dequant + HFMA2, short-conv + silu with conv-window state update, gated-delta recurrence writing S in place, o_proj + residual), MLA (absorbed qlat path, rope, latent-cache attention with alien-row copy + append, Wv + o_proj + residual), per-block MoE (router softmax, top-8, weight renorm * routed_scaling, 8 routed + 1 shared expert gate/up/down int4 chains). State writeback matches check.py contract (S / c_kv / out cosine). No dropped terms for cosine gaming. Host _cache_reg is data_ptr-keyed MLA KV buffer capacity pooling only (grows c_kv/k_rope storage; copy_rows/ncop tell the kernel when to copy alien rows) — not output memoization; every step() still launches mega_step with live hidden and live S/conv/cache pointers. phase_limit hard-coded to -1 (full run) in step(). template_mutated= false; all template_files (check/benchmark/reference/baseline/problem.yaml/ shapes/PROMPT/sota) md5-identical across template_files, workspace, and ~/kernelbench.com/benchmarks/mega/problems/02_kimi_linear_decode. check.log PASS with out/S/cache cosines well above 0.98 (e.g. out 0.9945- 0.9999, S 0.9993-0.9999, cache 1.0000). Mega check.py is cosine-oracle (no KBH_NUMERIC_STRESS path on this problem); zero KBH_NUMERIC_STRESS=0 in transcript. Empirical recompute on H100 under flock: overwrite the same hidden buffer in place between two step() calls on cloned state — max_abs_delta=1.890625, cos(before,after)=0.259; vs reference on mutated inputs out_cos=0.999978 S0_cos=1.0 cache_cos=1.0. Transcript: only Read of own check.py (no Write/Edit to graders); sole intentional other-archive action was `cp .../07_w4a16_gemm/repo/.venv/bin/ninja` for build tooling; other timestamps (fp8_gemm / topk / sonic_moe) appear only as concurrent docker/ps noise — never reading another kimi_linear_decode solution. Publish tripwire may still auto-flag the incidental timestamps; human audit finds no answer contamination.
Kernel source (redacted)
"""Kimi-Linear W4A16 hybrid decode unit — single-launch CUDA megakernel.
The ENTIRE per-token forward is one CUDA __global__ launch (cooperative
grid of persistent CTAs with software grid barriers):
* every int4 dequant-GEMV fused (nibble -> fp16 magic -> HFMA2, weights
streamed exactly once through cp.async rings; never materialized as bf16)
* KDA: fused qkvg GEMV + beta, causal depthwise conv, gated-delta recurrence
with S-state update in place, o_proj streamed row-wise
* MLA: absorbed attention (qlat = q_nope . W_uk^T then latent-dot softmax
over the compressed cache), rope inline, combine + W_v + o_proj
* MoE: router + top-8 + 9 fused expert gate/up/down chains with atomics
* both RMSNorms and residual adds are fused per-phase.
An exact eager PyTorch path (`_step_eager`) is kept as an oracle/debug aid;
the timed path is `step(...)` -> one kernel launch per token.
"""
from __future__ import annotations
import os
import tempfile
from dataclasses import dataclass, field
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
# =========================================================================== #
# CUDA megakernel source (built once via load_inline)
# =========================================================================== #
_CUDA_SRC = r"""
// ==== merged megakernel source (common + kda + mla phases + kernel) ====
// Kimi-Linear W4A16 decode megakernel. ONE launch per step.
// Phases (barrier after each):
// KDA blk b (0..2): P1 fused qkvg int4 GEMV + beta | P2 conv+recurrence+o_proj | P3 MoE
// MLA blk 3 : P1 q,kv_a + k_rope rope | P2a qlat+attention+cache | P2b combine+Wv+o_proj | P3 MoE
// final: x_out store + invariant zeroing
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_pipeline.h>
#include <ATen/cuda/CUDAContext.h>
#include <type_traits>
#include <vector>
typedef unsigned char u8;
typedef unsigned int u32;
typedef unsigned long long u64;
#define NTHREADS 256
#define HIDDEN 2304
#define KHID2 1152
#define NBAR 18
// ---------------- scratch fp32 offsets ----------------
#define SC_PROJ 0 // 16416
#define SC_BETA 16416 // 32
#define SC_CKNEW 16448 // 512
#define SC_KRNEW 16960 // 64
#define SC_QLATF 17024 // (unused now)
#define SC_OUTA 17024 // 2304
#define SC_OUTM 19328 // 2304
#define SC_WVO 21632 // 64 items * 64 (Wv outputs)
#define SC_CONV 25728 // 3*4096 conv outputs
#define SC_OX 38016 // 4096 o values
#define SC_PART 42112 // partial attention: 164*2*16*514
#define SC_FLOATS (42112 + 164 * 2 * 16 * 514)
// ---------------- bf16 scratch offsets (elements) ----------------
#define BS_QLAT 0 // 32*576: per head [qlat(512) | q_rope_r(64)]
#define BS_FLOATS (32 * 576)
// ---------------- static weight pointer table ----------------
// KDA block b at b*38: [q k v g o]x3 = 15, beta, conv, anorm, mnorm, router,
// moe routed 9 (eg,eu,ed)x3, shared 9 (sg,su,sd)x3
// MLA block at 114: q(3) kva(3) kvb(3) o(3) anorm mnorm router, then moe same 18
#define T_KW(b, i) ((b) * 38 + (i))
#define T_MW(i) (114 + (i))
#define TAB_SIZE 152
struct MegaArgs {
const u64* tab;
const __nv_bfloat16* x_in;
__nv_bfloat16* x_out;
float* S[3];
__nv_bfloat16* cq[3];
__nv_bfloat16* ck[3];
__nv_bfloat16* cv[3];
const __nv_bfloat16* ckv_in;
const __nv_bfloat16* kr_in;
__nv_bfloat16* ckv_out;
__nv_bfloat16* kr_out;
float* scratch;
__nv_bfloat16* bscratch;
int* bar_ref; // NBAR pairs (cnt, flg)
int pos;
int gen;
int nc; // attention chunks
int copy_rows; // >0: alien copy rows
int ncop; // copy items (ceil(pos / 256)) when copy_rows > 0
int phase_limit;
};
static __device__ __forceinline__ void gsync(const MegaArgs& A, int idx) {
__syncthreads();
if (threadIdx.x == 0) {
__threadfence();
int target = (A.gen + 1) * (int)gridDim.x;
int a = atomicAdd(&A.bar_ref[idx * 2], 1);
if (a == target - 1) {
atomicExch(&A.bar_ref[idx * 2 + 1], A.gen + 1);
} else {
while (atomicAdd(&A.bar_ref[idx * 2 + 1], 0) < A.gen + 1) { }
}
__threadfence();
}
__syncthreads();
}
static __device__ __forceinline__ const u8* TAB_U8(const u64* tab, int i) { return (const u8*)tab[i]; }
static __device__ __forceinline__ const __nv_bfloat16* TAB_BF(const u64* tab, int i) {
return (const __nv_bfloat16*)tab[i];
}
// exact-once grid-stride helper: element ranges per CTA, threads within
#define GRID_FOR(var, total, ncta, cta) for (int var = (cta) * (((total) + (ncta) - 1) / (ncta)) + threadIdx.x; var < min(((cta) + 1) * (((total) + (ncta) - 1) / (ncta)), (int)(total)); var += NTHREADS)
// ---------------- fp16 magic dequant helpers ----------------
static __device__ __forceinline__ __half2 deq_lo(u32 v) {
u32 t = __byte_perm(v & 0x0F0F0F0Fu, 0x64646464u, 0x4140);
return __hsub2(*(__half2*)&t, __float2half2_rn(1024.f));
}
static __device__ __forceinline__ __half2 deq_lo_hi(u32 v) {
u32 t = __byte_perm(v & 0x0F0F0F0Fu, 0x64646464u, 0x4342);
return __hsub2(*(__half2*)&t, __float2half2_rn(1024.f));
}
static __device__ __forceinline__ __half2 deq_hi(u32 v) {
u32 t = __byte_perm((v >> 4) & 0x0F0F0F0Fu, 0x64646464u, 0x4140);
return __hsub2(*(__half2*)&t, __float2half2_rn(1024.f));
}
static __device__ __forceinline__ __half2 deq_hi_hi(u32 v) {
u32 t = __byte_perm((v >> 4) & 0x0F0F0F0Fu, 0x64646464u, 0x4342);
return __hsub2(*(__half2*)&t, __float2half2_rn(1024.f));
}
// ---------------- stage_xn: xs fp16 (K) + gsx(ng) ----------------
template <typename XT>
static __device__ void stage_xn(const XT* __restrict__ x,
const __nv_bfloat16* __restrict__ normw,
__half* __restrict__ xs, float* __restrict__ gsx,
float* __restrict__ red, int K, int ngroups) {
int tid = threadIdx.x;
float ss = 0.f;
for (int i = tid; i < K; i += 256) {
float v;
if constexpr (std::is_same_v<XT, __nv_bfloat16>) v = __bfloat162float(x[i]);
else v = x[i];
ss += v * v;
}
#pragma unroll
for (int o = 16; o > 0; o >>= 1) ss += __shfl_down_sync(0xffffffffu, ss, o);
if ((tid & 31) == 0) red[tid >> 5] = ss;
__syncthreads();
if (tid < 32) {
float v = (tid < 8) ? red[tid] : 0.f;
#pragma unroll
for (int o = 4; o > 0; o >>= 1) v += __shfl_down_sync(0xffffffffu, v, o);
if (tid == 0) red[0] = v;
}
__syncthreads();
float scale = rsqrtf(red[0] / (float)K + 1e-6f);
int gsz = K / ngroups;
for (int g = tid; g < ngroups; g += 256) {
float a = 0.f;
for (int j = 0; j < gsz; j++) {
int i = g * gsz + j;
float v;
if constexpr (std::is_same_v<XT, __nv_bfloat16>) v = __bfloat162float(x[i]);
else v = x[i];
__half h = __float2half(v * scale * __bfloat162float(normw[i]));
xs[i] = h;
a += __half2float(h);
}
gsx[g] = a;
}
__syncthreads();
}
// ---------------- gemvA (stream) ----------------
// y[n] += sum_k x[k] (q[k,n]-z)*s, rows [r0, r0+64*ng), cols [n0, n0+COLS).
// COLS in {64, 128}. ring: STAGES x (64*COLS) slab bytes.
// lane covers 4 cols. For COLS=64: lanes 0..15 row r, 16..31 row r+1 (2 rows/warp).
// zshare = 1/8 per warp handled: caller multiplies zacc total by 0.125 (COLS=128)
// or per explicit zdiv. We apply zacc * (1/zdiv) at item end.
// COLS=128: warp covers 1 row per pass, 8 passes: r = warp + rr*8
// COLS=64: warp covers 2 rows per pass (lane halves), 4 passes: r = warp*2 + (lane>>4) + rr*16
template <int COLS>
struct GVMap {};
template <> struct GVMap<128> { static constexpr int lane_cols(int l) { return l * 4; }
static constexpr int row_off(int w, int l, int rr) { return w + rr * 8; }
static constexpr int N_ITER = 8; };
template <> struct GVMap<64> { static constexpr int lane_cols(int l) { return (l & 15) * 4; }
static constexpr int row_off(int w, int l, int rr) { return w * 2 + (l >> 4) + rr * 16; }
static constexpr int N_ITER = 4; };
template <int STAGES, int COLS>
static __device__ void gemvA(const u8* __restrict__ w, const __nv_bfloat16* __restrict__ sc,
const __nv_bfloat16* __restrict__ zc,
const __half* __restrict__ xs, const float* __restrict__ gsx,
float* __restrict__ out, int outstride,
int N, int n0, int r0, int ng, int zdiv,
u8* ring, long ring_stride) {
const unsigned M2 = 0x64646464u;
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
const u8* srcbase = w + (long long)r0 * N + n0;
auto produce = [&](int gslot, int gidx) {
const u8* src = srcbase + (long long)gidx * 64 * N;
int t = threadIdx.x;
u8* dst = ring + (long long)gslot * ring_stride;
static const int CPO = COLS * 64 / 256 / 16; // cp.async ops per thread
#pragma unroll
for (int c = 0; c < CPO; c++) {
int tt = t + c * 256;
int row = tt / (COLS / 16), seg = tt % (COLS / 16);
__pipeline_memcpy_async(dst + row * COLS + seg * 16,
src + (long long)row * N + seg * 16, 16);
}
__pipeline_commit();
};
#pragma unroll
for (int p = 0; p < STAGES - 1; p++) {
if (p < ng) produce(p, p);
}
float gy[4] = {0.f, 0.f, 0.f, 0.f};
float zacc[4] = {0.f, 0.f, 0.f, 0.f};
int lcol = GVMap<COLS>::lane_cols(lane);
for (int g = 0; g < ng; g++) {
__pipeline_wait_prior(STAGES - 2);
__syncthreads();
const u8* slab = ring + (long long)(g % STAGES) * ring_stride;
int grp = (r0 + g * 64) / 64;
__half2 acc01 = __half2(__ushort_as_half(0), __ushort_as_half(0));
__half2 acc23 = acc01;
#pragma unroll
for (int rr = 0; rr < GVMap<COLS>::N_ITER; rr++) {
int r = GVMap<COLS>::row_off(warp, lane, rr);
unsigned v = *(const unsigned*)(slab + r * COLS + lcol);
__half2 e01 = deq_lo(v), e23 = deq_lo_hi(v), o01 = deq_hi(v), o23 = deq_hi_hi(v);
__half2 xe2 = __half2half2(xs[2 * (r0 + g * 64 + r)]);
__half2 xo2 = __half2half2(xs[2 * (r0 + g * 64 + r) + 1]);
acc01 = __hfma2(e01, xe2, __hfma2(o01, xo2, acc01));
acc23 = __hfma2(e23, xe2, __hfma2(o23, xo2, acc23));
}
if (g + STAGES - 1 < ng) produce((g + STAGES - 1) % STAGES, g + STAGES - 1);
float sx = gsx[grp];
uint2 sv = *(const uint2*)(sc + (long long)grp * N + n0 + lcol);
uint2 zv = *(const uint2*)(zc + (long long)grp * N + n0 + lcol);
const __nv_bfloat16* svv = (const __nv_bfloat16*)&sv;
const __nv_bfloat16* zvv = (const __nv_bfloat16*)&zv;
#pragma unroll
for (int j = 0; j < 2; j++) {
__half2 a = j == 0 ? acc01 : acc23;
float alo = __low2float(a), ahi = __high2float(a);
float sc0 = __bfloat162float(svv[2 * j]), zc0 = __bfloat162float(zvv[2 * j]);
float sc1 = __bfloat162float(svv[2 * j + 1]), zc1 = __bfloat162float(zvv[2 * j + 1]);
gy[j * 2] += alo * sc0;
gy[j * 2 + 1] += ahi * sc1;
zacc[j * 2] += zc0 * sx * sc0;
zacc[j * 2 + 1] += zc1 * sx * sc1;
}
__syncthreads();
}
float zf = 1.f / (float)zdiv;
#pragma unroll
for (int j = 0; j < 4; j++) atomicAdd(&out[(long long)(n0 + lcol + j) * outstride], gy[j] - zacc[j] * zf);
// drain async pipeline so a later gemv's produces can't be clobbered by stragglers
__pipeline_wait_prior(0);
__syncthreads();
}
// ---------------- gemv_rows ----------------
// Stream a small set of packed rows (ROWS in {16,32}) x N cols (N=2304-like).
// y[n] += wj * sum_r x[2r,2r+1] (q-z)s. w/sc/zc pre-offset to (row r0, group grp).
// xs: fp16 smem, relative to r0*2. sxo: sum of all 2*ROWS xs values (precomputed).
// warp w covers rows warp*(ROWS/8)..+(ROWS/8). z-share per warp = 1/zdiv (caller: 8).
template <int ROWS>
static __device__ void gemv_rows(const u8* __restrict__ w, const __nv_bfloat16* __restrict__ sc,
const __nv_bfloat16* __restrict__ zc,
const __half* __restrict__ xs, float sxo, int N,
float* __restrict__ out, float wj, int zdiv) {
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
constexpr int RPW = ROWS / 8;
const int CBT = N / 128;
float share = 1.f / (float)zdiv;
for (int cb = 0; cb < CBT; cb++) {
__half2 acc01 = __half2(__ushort_as_half(0), __ushort_as_half(0));
__half2 acc23 = acc01;
#pragma unroll
for (int rw = 0; rw < RPW; rw++) {
int r = warp * RPW + rw;
unsigned v = *(const unsigned*)(w + (long long)r * N + cb * 128 + lane * 4);
__half2 e01 = deq_lo(v), e23 = deq_lo_hi(v), o01 = deq_hi(v), o23 = deq_hi_hi(v);
__half2 xe2 = __half2half2(xs[2 * r]);
__half2 xo2 = __half2half2(xs[2 * r + 1]);
acc01 = __hfma2(e01, xe2, __hfma2(o01, xo2, acc01));
acc23 = __hfma2(e23, xe2, __hfma2(o23, xo2, acc23));
}
int n = cb * 128 + lane * 4;
uint2 sv = *(const uint2*)(sc + (long long)n);
uint2 zv = *(const uint2*)(zc + (long long)n);
const __nv_bfloat16* svv = (const __nv_bfloat16*)&sv;
const __nv_bfloat16* zvv = (const __nv_bfloat16*)&zv;
#pragma unroll
for (int j = 0; j < 2; j++) {
__half2 a = j == 0 ? acc01 : acc23;
float alo = __low2float(a), ahi = __high2float(a);
float sc0 = __bfloat162float(svv[2 * j]), zc0 = __bfloat162float(zvv[2 * j]);
float sc1 = __bfloat162float(svv[2 * j + 1]), zc1 = __bfloat162float(zvv[2 * j + 1]);
atomicAdd(&out[n + 2 * j], (alo - zc0 * sxo * share) * sc0 * wj);
atomicAdd(&out[n + 2 * j + 1], (ahi - zc1 * sxo * share) * sc1 * wj);
}
}
}
// Phase implementations: KDA blocks + MoE (shared by KDA/MLA).
// shared workspace layout inside `big` (33KB):
// p2_kda: kst[128] qst[128] gst[128] (fp32 x3) | vst[32] ost[8] | redsum[32][8] (fp32) | oxh[32] (fp16)
#define WK_KST ((float*)big)
#define WK_QST (WK_KST + 128)
#define WK_GST (WK_QST + 128)
#define WK_VST (WK_GST + 128)
#define WK_OST (WK_VST + 32)
#define WK_RED (WK_OST + 8) // 32*8 floats
#define WK_OXH ((__half*)(WK_RED + 256))
// ================================================================== //
// P1 KDA: fused qkvg int4 GEMV (256 items) + beta (1 item)
// ================================================================== //
static __device__ void p1_kda(const MegaArgs& A, int blk, __half* xs, float* gsx, float* red, u8* ring) {
if (blk == 0) stage_xn<__nv_bfloat16>(A.x_in, TAB_BF(A.tab, T_KW(0, 17)), xs, gsx, red, HIDDEN, 18);
else stage_xn<float>(A.scratch + SC_OUTM, TAB_BF(A.tab, T_KW(blk, 17)), xs, gsx, red, HIDDEN, 18);
for (int i = threadIdx.x; i < HIDDEN; i += 256) A.scratch[SC_OUTA + i] = 0.f;
__syncthreads();
int ncta = gridDim.x, cta = blockIdx.x;
for (int it = cta; it < 256; it += ncta) {
int nb = it & 127, kh = it >> 7;
int n0g = nb * 128;
int which = n0g >> 12;
int col = n0g & 4095;
const u8* w = TAB_U8(A.tab, T_KW(blk, which * 3));
const __nv_bfloat16* sc = TAB_BF(A.tab, T_KW(blk, which * 3 + 1));
const __nv_bfloat16* zc = TAB_BF(A.tab, T_KW(blk, which * 3 + 2));
gemvA<4, 128>(w + col, sc + col, zc + col, xs, gsx, A.scratch + SC_PROJ + n0g,
1, 4096, 0, kh * 576, 9, 8, ring, 8192);
}
if (cta == 0) {
const __nv_bfloat16* W = TAB_BF(A.tab, T_KW(blk, 15)); // (32, 2304)
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
int eg = warp * 4 + (lane >> 3);
int sl = lane & 7;
float part = 0.f;
for (int d = sl * 288; d < (sl + 1) * 288 && eg < 32; d++) {
part += __half2float(xs[d]) * __bfloat162float(W[eg * 2304 + d]);
}
#pragma unroll
for (int o = 4; o > 0; o >>= 1) part += __shfl_xor_sync(0xffffffffu, part, o, 8);
if ((threadIdx.x & 7) == 0 && eg < 32) A.scratch[SC_BETA + eg] = part;
}
}
// ================================================================== //
// P2R KDA: conv (all heads, race-free) + recurrence + o write (32 items)
// ox output: SC_OX (4096 fp32, o values post-fp16-rounding)
// ================================================================== //
static __device__ void p2r_kda(const MegaArgs& A, int blk, u8* big) {
float* kst = (float*)big; // 128
float* qst = (float*)(big + 512); // 128
float* gst = (float*)(big + 1024); // 128
float* ost = (float*)(big + 1536); // 8
float (*redsum)[8] = (float(*)[8])(big + 1568); // 32x8
int ncta = gridDim.x, cta = blockIdx.x;
GRID_FOR(i, HIDDEN, ncta, cta) {
float v = (blk == 0) ? __bfloat162float(A.x_in[i]) : A.scratch[SC_OUTM + i];
atomicAdd(&A.scratch[SC_OUTA + i], v);
if (blk != 0) A.scratch[SC_OUTM + i] = 0.f;
}
__syncthreads();
// ---- conv for all channels: 3 kinds x 4096 (fp32 out into SC_CONV) ----
__nv_bfloat16* cvst[3];
cvst[0] = A.cq[blk];
cvst[1] = A.ck[blk];
cvst[2] = A.cv[blk];
GRID_FOR(cc, 12288, ncta, cta) {
{
int kind = cc / 4096, ch = cc % 4096;
__nv_bfloat16* st = cvst[kind];
const __nv_bfloat16* cw = TAB_BF(A.tab, T_KW(blk, 16)) + (long long)kind * (4096 * 4) + ch * 4;
float outv = __bfloat162float(st[ch]) * __bfloat162float(cw[0])
+ __bfloat162float(st[4096 + ch]) * __bfloat162float(cw[1])
+ __bfloat162float(st[8192 + ch]) * __bfloat162float(cw[2]);
float cur = A.scratch[SC_PROJ + kind * 4096 + ch];
outv += cur * __bfloat162float(cw[3]);
outv = outv / (1.f + __expf(-outv));
st[ch] = st[4096 + ch];
st[4096 + ch] = st[8192 + ch];
st[8192 + ch] = __float2bfloat16(cur);
A.scratch[SC_CONV + cc] = outv;
}
}
__syncthreads();
// ---- recurrence per head ----
float* Sg = A.S[blk];
for (int it = cta; it < 32; it += ncta) {
int h = it;
__syncthreads();
for (int d = threadIdx.x; d < 128; d += 256) {
float gv = A.scratch[SC_PROJ + 3 * 4096 + h * 128 + d];
float sp = (gv > 20.f) ? gv : log1pf(__expf(gv));
gst[d] = __expf(-sp);
kst[d] = A.scratch[SC_CONV + 4096 + h * 128 + d];
qst[d] = A.scratch[SC_CONV + h * 128 + d] * 0.08838834764831845f;
}
if (threadIdx.x == 0) ost[0] = 1.f / (1.f + __expf(-A.scratch[SC_BETA + h]));
__syncthreads();
float beta = ost[0];
for (int dvb = 0; dvb < 4; dvb++) {
int dv0 = dvb * 32;
__syncthreads();
int dv = threadIdx.x & 31, dkp = threadIdx.x >> 5;
float Sl[16];
float kbase[16], qbase[16];
#pragma unroll
for (int j = 0; j < 16; j++) {
int dk = dkp * 16 + j;
Sl[j] = Sg[(long long)h * 16384 + dk * 128 + dv0 + dv] * gst[dk];
kbase[j] = kst[dk];
qbase[j] = qst[dk];
}
float part = 0.f;
#pragma unroll
for (int j = 0; j < 16; j++) part += Sl[j] * kbase[j];
redsum[dv][dkp] = part;
__syncthreads();
float errv = 0.f;
if (dkp == 0) {
float pred = 0.f;
#pragma unroll
for (int j = 0; j < 8; j++) pred += redsum[dv][j];
errv = A.scratch[SC_CONV + 2 * 4096 + h * 128 + dv0 + dv] - pred;
redsum[dv][0] = errv;
}
__syncthreads();
errv = redsum[dv][0];
float opart = 0.f;
#pragma unroll
for (int j = 0; j < 16; j++) {
int dk = dkp * 16 + j;
float s2 = Sl[j] + beta * kbase[j] * errv;
Sg[(long long)h * 16384 + dk * 128 + dv0 + dv] = s2;
opart += s2 * qbase[j];
}
redsum[dv][dkp] = opart;
__syncthreads();
if (dkp == 0) {
float o = 0.f;
#pragma unroll
for (int j = 0; j < 8; j++) o += redsum[dv][j];
A.scratch[SC_OX + h * 128 + dv0 + dv] = __half2float(__float2half(o));
}
}
}
}
// ================================================================== //
// P2O KDA: o_proj only (128 items)
// ================================================================== //
static __device__ void p2o_kda(const MegaArgs& A, int blk, u8* big) {
int ncta = gridDim.x, cta = blockIdx.x;
const u8* w = TAB_U8(A.tab, T_KW(blk, 12));
const __nv_bfloat16* sc = TAB_BF(A.tab, T_KW(blk, 13));
const __nv_bfloat16* zc = TAB_BF(A.tab, T_KW(blk, 14));
__half* oxh = (__half*)big; // 32
float* ost = (float*)(big + 64);
for (int it = cta; it < 128; it += ncta) {
int h = it >> 2, dvb = it & 3;
__syncthreads();
for (int t = threadIdx.x; t < 32; t += 256)
oxh[t] = __float2half(A.scratch[SC_OX + h * 128 + dvb * 32 + t]);
__syncthreads();
float sxo = 0.f;
if (threadIdx.x < 32) sxo = __half2float(oxh[threadIdx.x]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) sxo += __shfl_down_sync(0xffffffffu, sxo, o);
if (threadIdx.x == 0) ost[0] = sxo;
__syncthreads();
sxo = ost[0];
int r0 = h * 64 + dvb * 16;
gemv_rows<16>(w + (long long)r0 * 2304, sc + (long long)h * 2304, zc + (long long)h * 2304,
oxh, sxo, 2304, A.scratch + SC_OUTA, 1.f, 8);
__syncthreads();
}
}
// ================================================================== //
// P3 MoE (shared by KDA blocks and MLA block)
// tab layout per kind:
// KDA blk: router=19, routed at 20(eg),23(eu),26(ed), shared at 29(sg),32(su),35(sd)
// MLA : router=14, routed at 15(eg),18(eu),21(ed), shared at 24(sg),27(su),30(sd)
// ================================================================== //
struct MoETab {
const __nv_bfloat16* router;
const u8 *egw, *euw, *edw, *sgw, *suw, *sdw;
const __nv_bfloat16 *egs, *eus, *eds, *sgs, *sus, *sds;
const __nv_bfloat16 *egz, *euz, *edz, *sgz, *suz, *sdz;
const __nv_bfloat16* moe_norm;
};
static __device__ MoETab moe_tab(const MegaArgs& A, int blk) {
MoETab t;
int b = (blk == 3) ? 114 : blk * 38;
int ro = (blk == 3) ? 14 : 19;
int eg = (blk == 3) ? 15 : 20;
int eu = (blk == 3) ? 18 : 23;
int ed = (blk == 3) ? 21 : 26;
int sg = (blk == 3) ? 24 : 29;
int su = (blk == 3) ? 27 : 32;
int sd = (blk == 3) ? 30 : 35;
t.router = TAB_BF(A.tab, b + ro);
t.egw = TAB_U8(A.tab, b + eg);
t.egs = TAB_BF(A.tab, b + eg + 1);
t.egz = TAB_BF(A.tab, b + eg + 2);
t.euw = TAB_U8(A.tab, b + eu);
t.eus = TAB_BF(A.tab, b + eu + 1);
t.euz = TAB_BF(A.tab, b + eu + 2);
t.edw = TAB_U8(A.tab, b + ed);
t.eds = TAB_BF(A.tab, b + ed + 1);
t.edz = TAB_BF(A.tab, b + ed + 2);
t.sgw = TAB_U8(A.tab, b + sg);
t.sgs = TAB_BF(A.tab, b + sg + 1);
t.sgz = TAB_BF(A.tab, b + sg + 2);
t.suw = TAB_U8(A.tab, b + su);
t.sus = TAB_BF(A.tab, b + su + 1);
t.suz = TAB_BF(A.tab, b + su + 2);
t.sdw = TAB_U8(A.tab, b + sd);
t.sds = TAB_BF(A.tab, b + sd + 1);
t.sdz = TAB_BF(A.tab, b + sd + 2);
t.moe_norm = TAB_BF(A.tab, b + ((blk == 3) ? 13 : 18));
return t;
}
// P3: router + fused expert chain. h stash per CTA in smem.
static __device__ void p3_moe(const MegaArgs& A, int blk, __half* xs, float* gsx, float* red, u8* big) {
MoETab mt = moe_tab(A, blk);
stage_xn<float>(A.scratch + SC_OUTA, mt.moe_norm, xs, gsx, red, HIDDEN, 18);
float* rlog = (float*)big; // 64
float* topw = rlog + 64; // 8
int* tope = (int*)(topw + 8); // 8
float* hst = (float*)(tope + 8); // 64 (h moe, fp32 half of it)
__syncthreads();
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
// router logits: warp w covers e in {w*8..w*8+7}
{
#pragma unroll
for (int ee = 0; ee < 8; ee++) {
int e = warp * 8 + ee;
float part = 0.f;
const __nv_bfloat16* Wr = mt.router + (long long)e * HIDDEN;
for (int d = lane * 72; d < lane * 72 + 72; d++)
part += __half2float(xs[d]) * __bfloat162float(Wr[d]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) part += __shfl_xor_sync(0xffffffffu, part, o);
if (lane == 0 && e < 64) rlog[e] = part;
}
}
__syncthreads();
if (threadIdx.x == 0) {
float mx = rlog[0];
for (int e = 1; e < 64; e++) mx = fmaxf(mx, rlog[e]);
float pr[64];
float z = 0.f;
for (int e = 0; e < 64; e++) { pr[e] = __expf(rlog[e] - mx); z += pr[e]; }
for (int e = 0; e < 64; e++) pr[e] /= z;
bool used[64];
for (int e = 0; e < 64; e++) used[e] = false;
float wsum = 0.f;
int te[8]; float tw[8];
for (int j = 0; j < 8; j++) {
int best = -1; float bv = -1.f;
for (int e = 0; e < 64; e++) if (!used[e] && pr[e] > bv) { bv = pr[e]; best = e; }
if (best < 0) best = 0;
used[best] = true; te[j] = best; tw[j] = bv; wsum += bv;
}
for (int j = 0; j < 8; j++) { topw[j] = tw[j] / (wsum + 1e-9f) * 2.446f; tope[j] = te[j]; }
}
__syncthreads();
// residual duty: out_moe += out_attn slice (h)
int nctaR = gridDim.x;
GRID_FOR(i, HIDDEN, nctaR, blockIdx.x) atomicAdd(&A.scratch[SC_OUTM + i], A.scratch[SC_OUTA + i]);
// zero duties for next P1: s_proj + ckv/kr staging
GRID_FOR(i, 16416, nctaR, blockIdx.x) A.scratch[SC_PROJ + i] = 0.f;
GRID_FOR(i, 576, nctaR, blockIdx.x) A.scratch[SC_CKNEW + i] = 0.f;
int ncta = gridDim.x, cta = blockIdx.x;
for (int it = cta; it < 144; it += ncta) {
int j = it / 16, mb = it % 16;
bool shared_expert = (j == 8);
int e = shared_expert ? 0 : tope[j];
float wj = shared_expert ? 1.f : topw[j];
const u8* gw = shared_expert ? mt.sgw : mt.egw;
const __nv_bfloat16* gsc = shared_expert ? mt.sgs : mt.egs;
const __nv_bfloat16* gzc = shared_expert ? mt.sgz : mt.egz;
const u8* uw = shared_expert ? mt.suw : mt.euw;
const __nv_bfloat16* usc = shared_expert ? mt.sus : mt.eus;
const __nv_bfloat16* uzc = shared_expert ? mt.suz : mt.euz;
const u8* dw = shared_expert ? mt.sdw : mt.edw;
const __nv_bfloat16* dsc = shared_expert ? mt.sds : mt.eds;
const __nv_bfloat16* dzc = shared_expert ? mt.sdz : mt.edz;
long eoff = (long long)e * (1152LL * 1024);
long soff = (long long)e * (18LL * 1024);
long doff = (long long)e * (512LL * 2304);
long dsoff = (long long)e * (8LL * 2304);
for (int i = threadIdx.x; i < 128; i += 256) hst[i] = 0.f;
__syncthreads();
gemvA<4, 64>(gw + eoff + mb * 64, gsc + soff + mb * 64, gzc + soff + mb * 64,
xs, gsx, hst, 1, 1024, 0, 0, 18, 16, big + 16384, 4096);
gemvA<4, 64>(uw + eoff + mb * 64, usc + soff + mb * 64, uzc + soff + mb * 64,
xs, gsx, hst + 64, 1, 1024, 0, 0, 18, 16, big + 16384, 4096);
__half* hf = (__half*)(hst + 136); // 70 halves scratch
__syncthreads();
float sxo = 0.f;
for (int i = threadIdx.x; i < 64; i += 256) {
float g = hst[i], u = hst[64 + i];
float hv = (g / (1.f + __expf(-g))) * u;
hf[i] = __float2half(hv);
sxo += hv;
}
#pragma unroll
for (int o = 16; o > 0; o >>= 1) sxo += __shfl_xor_sync(0xffffffffu, sxo, o);
if ((threadIdx.x & 31) == 0 && threadIdx.x < 64) hst[132 + (threadIdx.x >> 5)] = sxo;
__syncthreads();
sxo = hst[132] + hst[133];
gemv_rows<32>(dw + doff + (long long)mb * 32 * 2304, dsc + dsoff + (long long)(mb / 2) * 2304,
dzc + dsoff + (long long)(mb / 2) * 2304,
hf, sxo, 2304, A.scratch + SC_OUTM, wj, 8);
__syncthreads();
}
}
// MLA phases: P1 | P2a (qlat+rope) | P2b (attention+copy+cachewrite) | P2c (combine+Wv+o_proj)
// ================================================================== //
// P1 MLA: q_proj (48 blocks x ks3) + kv_a c_kv (4+1 x ks3) + kv k_rope raw (ks3)
// ================================================================== //
static __device__ void p1_mla(const MegaArgs& A, __half* xs, float* gsx, float* red, u8* ring) {
stage_xn<float>(A.scratch + SC_OUTM, TAB_BF(A.tab, T_MW(12)), xs, gsx, red, HIDDEN, 18);
for (int i = threadIdx.x; i < HIDDEN; i += 256) A.scratch[SC_OUTA + i] = 0.f;
__syncthreads();
int ncta = gridDim.x, cta = blockIdx.x;
for (int it = cta; it < 159; it += ncta) {
if (it < 144) {
int nb = it % 48, kh = it / 48;
const u8* w = TAB_U8(A.tab, T_MW(0));
const __nv_bfloat16* sc = TAB_BF(A.tab, T_MW(1));
const __nv_bfloat16* zc = TAB_BF(A.tab, T_MW(2));
gemvA<4, 128>(w + nb * 128, sc + nb * 128, zc + nb * 128, xs, gsx,
A.scratch + SC_PROJ + nb * 128, 1, 6144, 0, kh * 384, 6, 8, ring, 8192);
} else if (it < 156) {
int nb = (it - 144) % 4, kh = (it - 144) / 4;
const u8* w = TAB_U8(A.tab, T_MW(3));
const __nv_bfloat16* sc = TAB_BF(A.tab, T_MW(4));
const __nv_bfloat16* zc = TAB_BF(A.tab, T_MW(5));
gemvA<4, 128>(w + nb * 128, sc + nb * 128, zc + nb * 128, xs, gsx,
A.scratch + SC_CKNEW + nb * 128, 1, 576, 0, kh * 384, 6, 8, ring, 8192);
} else {
int kh = it - 156;
const u8* w = TAB_U8(A.tab, T_MW(3));
const __nv_bfloat16* sc = TAB_BF(A.tab, T_MW(4));
const __nv_bfloat16* zc = TAB_BF(A.tab, T_MW(5));
gemvA<4, 64>(w + 512, sc + 512, zc + 512, xs, gsx,
A.scratch + SC_KRNEW, 1, 576, 0, kh * 384, 6, 16, ring, 4096);
}
}
}
// ================================================================== //
// P2a: qlat gemvB (transposed) + q_rope rope -> BS_QLAT (32 heads)
// ================================================================== //
static __device__ void p2a_mla(const MegaArgs& A, u8* big) {
__half* qh = (__half*)big; // 128
int ncta = gridDim.x, cta = blockIdx.x;
const u8* kb = TAB_U8(A.tab, T_MW(6));
const __nv_bfloat16* kbs = TAB_BF(A.tab, T_MW(7));
const __nv_bfloat16* kbz = TAB_BF(A.tab, T_MW(8));
for (int it = cta; it < 32; it += ncta) {
int h = it;
__syncthreads();
for (int t = threadIdx.x; t < 128; t += 256)
qh[t] = __float2half(A.scratch[SC_PROJ + h * 192 + t]);
if (threadIdx.x < 32) {
int i = threadIdx.x;
float e = A.scratch[SC_PROJ + h * 192 + 128 + 2 * i];
float o = A.scratch[SC_PROJ + h * 192 + 128 + 2 * i + 1];
float ang = A.pos * powf(10000.f, -(float)i / 32.f);
float cs = __cosf(ang), sn = __sinf(ang);
A.bscratch[BS_QLAT + h * 576 + 512 + 2 * i] = __float2bfloat16(e * cs - o * sn);
A.bscratch[BS_QLAT + h * 576 + 512 + 2 * i + 1] = __float2bfloat16(o * cs + e * sn);
}
__syncthreads();
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
int rsh = lane >> 4; // unused here (COLS n/a)
(void)rsh;
for (int g = 0; g < 4; g++) {
// load sc/z for my 4 cols once per group
int colb = h * 256 + lane * 4;
uint2 sv = *(const uint2*)(kbs + (long long)g * 8192 + colb);
uint2 zv = *(const uint2*)(kbz + (long long)g * 8192 + colb);
const __nv_bfloat16* svv = (const __nv_bfloat16*)&sv;
const __nv_bfloat16* zvv = (const __nv_bfloat16*)&zv;
float q0 = __half2float(qh[lane * 4]), q1 = __half2float(qh[lane * 4 + 1]);
float q2 = __half2float(qh[lane * 4 + 2]), q3 = __half2float(qh[lane * 4 + 3]);
float s0 = __bfloat162float(svv[0]) * q0, s1 = __bfloat162float(svv[1]) * q1;
float s2 = __bfloat162float(svv[2]) * q2, s3 = __bfloat162float(svv[3]) * q3;
float zl = __bfloat162float(zvv[0]) * s0 + __bfloat162float(zvv[1]) * s1
+ __bfloat162float(zvv[2]) * s2 + __bfloat162float(zvv[3]) * s3;
// block reduce dz over lanes
#pragma unroll
for (int o = 16; o > 0; o >>= 1) zl += __shfl_xor_sync(0xffffffffu, zl, o);
// rows for this warp in group: warp*8 + rr
#pragma unroll
for (int rr = 0; rr < 8; rr++) {
int r = g * 64 + warp * 8 + rr;
unsigned v = *(const unsigned*)(kb + (long long)r * 8192 + colb);
unsigned lo = v & 0x0F0F0F0Fu;
unsigned hi = (v >> 4) & 0x0F0F0F0Fu;
float pe = ((lo & 0xF) * s0 + ((lo >> 8) & 0xF) * s1 + ((lo >> 16) & 0xF) * s2 + ((lo >> 24) & 0xF) * s3);
float po = ((hi & 0xF) * s0 + ((hi >> 8) & 0xF) * s1 + ((hi >> 16) & 0xF) * s2 + ((hi >> 24) & 0xF) * s3);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
pe += __shfl_xor_sync(0xffffffffu, pe, o);
po += __shfl_xor_sync(0xffffffffu, po, o);
}
if (lane == 0) {
A.bscratch[BS_QLAT + h * 576 + 2 * r] = __float2bfloat16(pe - zl);
A.bscratch[BS_QLAT + h * 576 + 2 * r + 1] = __float2bfloat16(po - zl);
}
}
}
}
}
// ================================================================== //
// P2b: attention (nc*2 items) + alien copy (ncop) + cache write (1)
// qlatblk staged in smem `big` (16*576 bf16 = 18.4KB) + ccur (576 bf16) after it
// ================================================================== //
static __device__ void p2b_mla(const MegaArgs& A, u8* big) {
__nv_bfloat16* qlb = (__nv_bfloat16*)big; // 16*576
__nv_bfloat16* ccur = qlb + 16 * 576; // 576
float* ost = (float*)(ccur + 576); // small
int ncta = gridDim.x, cta = blockIdx.x;
int nc = A.nc;
int G = (A.pos + nc - 1) / nc;
const __nv_bfloat16* kb_tab = TAB_BF(A.tab, T_MW(6)); // unused
(void)kb_tab;
float scale = 0.07216878364870323f; // 192^-0.5
int total = nc * 2 + (A.copy_rows > 0 ? A.ncop : 0) + 1;
for (int it = cta; it < total; it += ncta) {
if (it < nc * 2) {
int c = it >> 1, hg = it & 1;
__syncthreads();
// stage qlatblk (18.4KB)
for (int i = threadIdx.x; i < 16 * 576 / 4; i += 256)
((uint2*)qlb)[i] = ((const uint2*)(A.bscratch + BS_QLAT + hg * 16 * 576))[i];
__syncthreads();
int t0 = c * G, t1 = min((c + 1) * G, A.pos);
int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
int h0 = warp * 2, h1 = warp * 2 + 1;
float o0[16], o1[16];
#pragma unroll
for (int j = 0; j < 16; j++) { o0[j] = 0.f; o1[j] = 0.f; }
float m0 = -1e30f, m1 = -1e30f, z0 = 0.f, z1 = 0.f;
int tstart = t0;
// chunk0: fold the current token from scratch first
if (c == 0) {
for (int i = threadIdx.x; i < 512; i += 256)
ccur[i] = __float2bfloat16(A.scratch[SC_CKNEW + i]);
if (threadIdx.x < 32) {
int i = threadIdx.x;
float e = A.scratch[SC_KRNEW + 2 * i];
float o = A.scratch[SC_KRNEW + 2 * i + 1];
float ang = A.pos * powf(10000.f, -(float)i / 32.f);
float cs = __cosf(ang), sn = __sinf(ang);
ccur[512 + 2 * i] = __float2bfloat16(e * cs - o * sn);
ccur[512 + 2 * i + 1] = __float2bfloat16(o * cs + e * sn);
}
__syncthreads();
// process virtual token from ccur
float cv0[16];
#pragma unroll
for (int j = 0; j < 16; j++) cv0[j] = __bfloat162float(ccur[lane * 16 + j]);
float r0 = __bfloat162float(ccur[512 + lane * 2]);
float r1 = __bfloat162float(ccur[512 + lane * 2 + 1]);
float p0 = 0.f, p1 = 0.f;
#pragma unroll
for (int j = 0; j < 16; j++) {
p0 += cv0[j] * __bfloat162float(qlb[h0 * 576 + lane * 16 + j]);
p1 += cv0[j] * __bfloat162float(qlb[h1 * 576 + lane * 16 + j]);
}
p0 += r0 * __bfloat162float(qlb[h0 * 576 + 512 + lane * 2]) + r1 * __bfloat162float(qlb[h0 * 576 + 512 + lane * 2 + 1]);
p1 += r0 * __bfloat162float(qlb[h1 * 576 + 512 + lane * 2]) + r1 * __bfloat162float(qlb[h1 * 576 + 512 + lane * 2 + 1]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
p0 += __shfl_xor_sync(0xffffffffu, p0, o);
p1 += __shfl_xor_sync(0xffffffffu, p1, o);
}
p0 *= scale; p1 *= scale;
m0 = p0; m1 = p1;
z0 = 1.f; z1 = 1.f;
#pragma unroll
for (int j = 0; j < 16; j++) {
o0[j] = cv0[j];
o1[j] = cv0[j];
}
}
for (int t = tstart; t < t1; t++) {
const __nv_bfloat16* kvr = A.ckv_in + (long long)t * 512;
const __nv_bfloat16* krr = A.kr_in + (long long)t * 64;
float cv[16];
#pragma unroll
for (int j = 0; j < 16; j++) cv[j] = __bfloat162float(kvr[lane * 16 + j]);
float r0 = __bfloat162float(krr[lane * 2]);
float r1 = __bfloat162float(krr[lane * 2 + 1]);
float p0 = 0.f, p1 = 0.f;
#pragma unroll
for (int j = 0; j < 16; j++) {
p0 += cv[j] * __bfloat162float(qlb[h0 * 576 + lane * 16 + j]);
p1 += cv[j] * __bfloat162float(qlb[h1 * 576 + lane * 16 + j]);
}
p0 += r0 * __bfloat162float(qlb[h0 * 576 + 512 + lane * 2]) + r1 * __bfloat162float(qlb[h0 * 576 + 512 + lane * 2 + 1]);
p1 += r0 * __bfloat162float(qlb[h1 * 576 + 512 + lane * 2]) + r1 * __bfloat162float(qlb[h1 * 576 + 512 + lane * 2 + 1]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
p0 += __shfl_xor_sync(0xffffffffu, p0, o);
p1 += __shfl_xor_sync(0xffffffffu, p1, o);
}
float s0 = p0 * scale, s1 = p1 * scale;
float mn0 = fmaxf(m0, s0), mn1 = fmaxf(m1, s1);
float f0 = __expf(m0 - mn0), f1 = __expf(m1 - mn1);
float e0 = __expf(s0 - mn0), e1 = __expf(s1 - mn1);
z0 = z0 * f0 + e0;
z1 = z1 * f1 + e1;
m0 = mn0; m1 = mn1;
#pragma unroll
for (int j = 0; j < 16; j++) {
o0[j] = o0[j] * f0 + e0 * cv[j];
o1[j] = o1[j] * f1 + e1 * cv[j];
}
}
// write partial slot
int slot = c * 2 + hg;
float* pd = A.scratch + SC_PART + (long long)(slot * 16 + h0) * 514;
#pragma unroll
for (int j = 0; j < 16; j++) pd[lane * 16 + j] = o0[j];
pd[512] = m0;
pd[513] = z0;
float* pd1 = A.scratch + SC_PART + (long long)(slot * 16 + h1) * 514;
#pragma unroll
for (int j = 0; j < 16; j++) pd1[lane * 16 + j] = o1[j];
pd1[512] = m1;
pd1[513] = z1;
} else if (A.copy_rows > 0 && it < nc * 2 + A.ncop) {
int cp = it - nc * 2;
int r0 = cp * 256, r1 = min(r0 + 256, A.pos);
for (int i = threadIdx.x; i < (r1 - r0) * 512 / 4; i += 256) {
((uint2*)(A.ckv_out + (long long)r0 * 512))[i] = ((const uint2*)(A.ckv_in + (long long)r0 * 512))[i];
}
for (int i = threadIdx.x; i < (r1 - r0) * 64 / 4; i += 256) {
((uint2*)(A.kr_out + (long long)r0 * 64))[i] = ((const uint2*)(A.kr_in + (long long)r0 * 64))[i];
}
} else if (it == total - 1) {
// cache writer
int p = A.pos;
for (int i = threadIdx.x; i < 512; i += 256)
A.ckv_out[(long long)p * 512 + i] = __float2bfloat16(A.scratch[SC_CKNEW + i]);
if (threadIdx.x < 32) {
int i = threadIdx.x;
float e = A.scratch[SC_KRNEW + 2 * i];
float o = A.scratch[SC_KRNEW + 2 * i + 1];
float ang = A.pos * powf(10000.f, -(float)i / 32.f);
float cs = __cosf(ang), sn = __sinf(ang);
A.kr_out[(long long)p * 64 + 2 * i] = __float2bfloat16(e * cs - o * sn);
A.kr_out[(long long)p * 64 + 2 * i + 1] = __float2bfloat16(o * cs + e * sn);
}
}
__syncthreads();
}
(void)ost;
}
// ================================================================== //
// P2c: combine attention partials + Wv GEMV + o_proj (64 items)
// ================================================================== //
static __device__ void p2c_mla(const MegaArgs& A, u8* big) {
int ncta = gridDim.x, cta = blockIdx.x;
GRID_FOR(i, HIDDEN, ncta, cta) {
atomicAdd(&A.scratch[SC_OUTA + i], A.scratch[SC_OUTM + i]);
A.scratch[SC_OUTM + i] = 0.f;
}
__syncthreads();
float* fpart = (float*)big; // 165 (f per chunk + zinv)
__half* ox = (__half*)(fpart + 168); // 512
__half* xf = ox + 512; // 70
float* gsxo = (float*)(xf + 72); // 4
float* ost = gsxo + 4; // 2
const u8* kb = TAB_U8(A.tab, T_MW(6));
const __nv_bfloat16* kbs = TAB_BF(A.tab, T_MW(7));
const __nv_bfloat16* kbz = TAB_BF(A.tab, T_MW(8));
const u8* ow = TAB_U8(A.tab, T_MW(9));
const __nv_bfloat16* ows = TAB_BF(A.tab, T_MW(10));
const __nv_bfloat16* owz = TAB_BF(A.tab, T_MW(11));
for (int it = cta; it < 64; it += ncta) {
int h = it >> 1, half = it & 1;
int hg = h >> 4, hid = h & 15;
__syncthreads();
// thread 0: combine m/z
if (threadIdx.x == 0) {
float M = -1e30f;
for (int c = 0; c < A.nc; c++) {
float m = A.scratch[SC_PART + (long long)((c * 2 + hg) * 16 + hid) * 514 + 512];
M = fmaxf(M, m);
}
float Z = 0.f;
for (int c = 0; c < A.nc; c++) {
const float* pd = A.scratch + SC_PART + (long long)((c * 2 + hg) * 16 + hid) * 514;
float f = __expf(pd[512] - M);
fpart[c] = f;
Z += f * pd[513];
}
fpart[164] = 1.f / Z;
}
__syncthreads();
float zinv = fpart[164];
int lane = threadIdx.x & 31;
float o16[16];
#pragma unroll
for (int j = 0; j < 16; j++) o16[j] = 0.f;
for (int c = 0; c < A.nc; c++) {
float f = fpart[c];
const float* pd = A.scratch + SC_PART + (long long)((c * 2 + hg) * 16 + hid) * 514;
#pragma unroll
for (int j = 0; j < 16; j++) o16[j] += f * pd[lane * 16 + j];
}
// stage ox (fp16) + group sums
__syncthreads();
#pragma unroll
for (int j = 0; j < 16; j++) ox[lane * 16 + j] = __float2half(o16[j] * zinv);
__syncthreads();
if (threadIdx.x < 32) {
int g = threadIdx.x >> 3;
if (g < 4) {
float a = 0.f;
int sl = threadIdx.x & 7;
#pragma unroll
for (int j = 0; j < 16; j++) a += __half2float(ox[g * 128 + sl * 16 + j]);
#pragma unroll
for (int o = 4; o > 0; o >>= 1) a += __shfl_xor_sync(0x1fffffffu, a, o, 8);
if (sl == 0) gsxo[g] = a;
}
}
__syncthreads();
// Wv: cols h*256+128+half*64 of kv_b (K=512 -> 2 packs per... K2=256, ng=4)
float* wvo = A.scratch + SC_WVO + it * 64;
for (int i = threadIdx.x; i < 64; i += 256) wvo[i] = 0.f;
__syncthreads();
int colv = h * 256 + 128 + half * 64;
gemvA<4, 64>(kb + colv, kbs + colv, kbz + colv, ox, gsxo, wvo, 1, 8192, 0, 0, 4, 16, big + 16384, 4096);
// o_proj: rows h*64 + half*32 .. +32
__syncthreads();
for (int i = threadIdx.x; i < 64; i += 256) xf[i] = __float2half(wvo[i]);
__syncthreads();
float sxo = 0.f;
for (int i = threadIdx.x; i < 64; i += 256) sxo += __half2float(xf[i]);
#pragma unroll
for (int o = 16; o > 0; o >>= 1) sxo += __shfl_xor_sync(0xffffffffu, sxo, o);
if ((threadIdx.x & 31) == 0 && threadIdx.x < 64) ost[threadIdx.x >> 5] = sxo;
__syncthreads();
sxo = ost[0] + ost[1];
gemv_rows<32>(ow + (long long)(h * 64 + half * 32) * 2304, ows + (long long)h * 2304,
owz + (long long)h * 2304, xf, sxo, 2304, A.scratch + SC_OUTA, 1.f, 8);
__syncthreads();
}
}
// Kimi-Linear W4A16 decode megakernel — one launch per decode step.
__global__ void __launch_bounds__(NTHREADS, 2) mega_kernel(MegaArgs A) {
__shared__ __half xs[HIDDEN];
__shared__ float gsx[32];
__shared__ float red[32];
__shared__ __align__(16) u8 big[34 * 1024];
u8* ring = big;
#define PHSYNC(i) do { gsync(A, i); if (A.phase_limit == (i)) return; } while (0)
// KDA blocks 0..2
#pragma unroll 1
for (int blk = 0; blk < 3; blk++) {
p1_kda(A, blk, xs, gsx, red, ring);
PHSYNC(blk * 4 + 0);
p2r_kda(A, blk, big);
PHSYNC(blk * 4 + 1);
p2o_kda(A, blk, big);
PHSYNC(blk * 4 + 2);
p3_moe(A, blk, xs, gsx, red, big);
PHSYNC(blk * 4 + 3);
}
// MLA block 3
p1_mla(A, xs, gsx, red, ring);
PHSYNC(12);
p2a_mla(A, big);
PHSYNC(13);
p2b_mla(A, big);
PHSYNC(14);
p2c_mla(A, big);
PHSYNC(15);
p3_moe(A, 3, xs, gsx, red, big);
PHSYNC(16);
// final: x_out store + invariant zero (no barrier needed at end)
GRID_FOR(i, HIDDEN, gridDim.x, blockIdx.x) {
float v = A.scratch[SC_OUTM + i];
A.x_out[i] = __float2bfloat16(v);
A.scratch[SC_OUTM + i] = 0.f;
A.scratch[SC_OUTA + i] = 0.f;
}
GRID_FOR(i, 16416, gridDim.x, blockIdx.x) A.scratch[SC_PROJ + i] = 0.f;
GRID_FOR(i, 576, gridDim.x, blockIdx.x) A.scratch[SC_CKNEW + i] = 0.f;
}
// ------------------------------------------------------------------ //
// host
// ------------------------------------------------------------------ //
static MegaArgs make_args(
torch::Tensor tab, torch::Tensor x_in, torch::Tensor x_out,
torch::Tensor S0, torch::Tensor S1, torch::Tensor S2,
torch::Tensor cq0, torch::Tensor cq1, torch::Tensor cq2,
torch::Tensor ck0, torch::Tensor ck1, torch::Tensor ck2,
torch::Tensor cv0, torch::Tensor cv1, torch::Tensor cv2,
torch::Tensor ckv_in, torch::Tensor kr_in,
torch::Tensor ckv_out, torch::Tensor kr_out,
torch::Tensor scratch, torch::Tensor bscratch, torch::Tensor bar,
long pos, long gen, long nc, long copy_rows, long ncop, long phase_limit) {
MegaArgs A;
A.tab = (const u64*)tab.data_ptr<int64_t>();
A.x_in = (const __nv_bfloat16*)x_in.data_ptr<at::BFloat16>();
A.x_out = (__nv_bfloat16*)x_out.data_ptr<at::BFloat16>();
A.S[0] = S0.data_ptr<float>();
A.S[1] = S1.data_ptr<float>();
A.S[2] = S2.data_ptr<float>();
A.cq[0] = (__nv_bfloat16*)cq0.data_ptr<at::BFloat16>();
A.cq[1] = (__nv_bfloat16*)cq1.data_ptr<at::BFloat16>();
A.cq[2] = (__nv_bfloat16*)cq2.data_ptr<at::BFloat16>();
A.ck[0] = (__nv_bfloat16*)ck0.data_ptr<at::BFloat16>();
A.ck[1] = (__nv_bfloat16*)ck1.data_ptr<at::BFloat16>();
A.ck[2] = (__nv_bfloat16*)ck2.data_ptr<at::BFloat16>();
A.cv[0] = (__nv_bfloat16*)cv0.data_ptr<at::BFloat16>();
A.cv[1] = (__nv_bfloat16*)cv1.data_ptr<at::BFloat16>();
A.cv[2] = (__nv_bfloat16*)cv2.data_ptr<at::BFloat16>();
A.ckv_in = (const __nv_bfloat16*)ckv_in.data_ptr<at::BFloat16>();
A.kr_in = (const __nv_bfloat16*)kr_in.data_ptr<at::BFloat16>();
A.ckv_out = (__nv_bfloat16*)ckv_out.data_ptr<at::BFloat16>();
A.kr_out = (__nv_bfloat16*)kr_out.data_ptr<at::BFloat16>();
A.scratch = scratch.data_ptr<float>();
A.bscratch = (__nv_bfloat16*)bscratch.data_ptr<at::BFloat16>();
A.bar_ref = bar.data_ptr<int>();
A.pos = (int)pos;
A.gen = (int)gen;
A.nc = (int)nc;
A.copy_rows = (int)copy_rows;
A.ncop = (int)ncop;
A.phase_limit = (int)phase_limit;
return A;
}
void mega_step(
torch::Tensor tab, torch::Tensor x_in, torch::Tensor x_out,
torch::Tensor S0, torch::Tensor S1, torch::Tensor S2,
torch::Tensor cq0, torch::Tensor cq1, torch::Tensor cq2,
torch::Tensor ck0, torch::Tensor ck1, torch::Tensor ck2,
torch::Tensor cv0, torch::Tensor cv1, torch::Tensor cv2,
torch::Tensor ckv_in, torch::Tensor kr_in,
torch::Tensor ckv_out, torch::Tensor kr_out,
torch::Tensor scratch, torch::Tensor bscratch, torch::Tensor bar,
long pos, long gen, long nc, long copy_rows, long ncop, long phase_limit, long ncta) {
MegaArgs A = make_args(tab, x_in, x_out, S0, S1, S2, cq0, cq1, cq2, ck0, ck1, ck2, cv0, cv1, cv2,
ckv_in, kr_in, ckv_out, kr_out, scratch, bscratch, bar,
pos, gen, nc, copy_rows, ncop, phase_limit);
auto stream = at::cuda::getCurrentCUDAStream();
void* kargs[] = {&A};
cudaError_t err = cudaLaunchCooperativeKernel((void*)mega_kernel, dim3((int)ncta), dim3(NTHREADS),
kargs, 0, stream.stream());
TORCH_CHECK(err == cudaSuccess, "cooperative launch failed: ", cudaGetErrorString(err));
}
long mega_occupancy() {
int n = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&n, mega_kernel, NTHREADS, 0);
return n;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("mega_step", &mega_step);
m.def("mega_occupancy", &mega_occupancy);
}
"""
def _build_ext():
cc = torch.cuda.get_device_capability(0)
arch = f"-gencode=arch=compute_{cc[0]}{cc[1]},code=sm_{cc[0]}{cc[1]}"
extdir = os.path.join(tempfile.gettempdir(), "kimi_linear_mega_ext")
os.makedirs(extdir, exist_ok=True)
return load_inline(
name="kimi_linear_megak",
cpp_sources="",
cuda_sources=_CUDA_SRC,
extra_cuda_cflags=["-O3", "--use_fast_math", arch],
verbose=False,
)
_EXT = None
def _ext():
global _EXT
if _EXT is None:
_EXT = _build_ext()
return _EXT
# --------------------------------------------------------------------------- #
# config / state helpers (signature-compatible with reference.py)
# --------------------------------------------------------------------------- #
@dataclass(frozen=True)
class Config:
hidden: int = 2304
kda_heads: int = 32
kda_head_dim: int = 128
short_conv: int = 4
mla_heads: int = 32
kv_lora: int = 512
qk_nope: int = 128
qk_rope: int = 64
v_head: int = 128
rope_theta: float = 10000.0
n_experts: int = 64
n_active: int = 8
n_shared: int = 1
moe_inter: int = 1024
routed_scaling: float = 2.446
group: int = 128
pattern: tuple = ("K", "K", "K", "M")
dtype: torch.dtype = field(default=torch.bfloat16)
def build_config(shape: dict) -> Config:
return Config(n_experts=int(shape.get("n_experts", 64)))
def 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
# --------------------------------------------------------------------------- #
# module tree identical to reference.py (buffer/parameter names must match)
# --------------------------------------------------------------------------- #
class QuantLinear(nn.Module):
def __init__(self, in_f: int, out_f: int, group: int = 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) -> torch.Tensor:
wu = torch.empty((self.in_f, self.out_f), dtype=torch.uint8, device=self.w_q.device)
wu[0::2] = self.w_q & 0xF
wu[1::2] = (self.w_q >> 4) & 0xF
s = self.scales.repeat_interleave(self.group, dim=0)
z = self.zeros.repeat_interleave(self.group, dim=0)
return (wu.to(torch.bfloat16) - z) * s
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:
wq, s, z = self.w_q[e], self.scales[e], self.zeros[e]
wu = torch.empty((self.in_f, self.out_f), dtype=torch.uint8, device=wq.device)
wu[0::2] = wq & 0xF
wu[1::2] = (wq >> 4) & 0xF
return (wu.to(torch.bfloat16) - z.repeat_interleave(self.group, 0)) * s.repeat_interleave(self.group, 0)
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)
class KDA(nn.Module):
def __init__(self, cfg: Config):
super().__init__()
self.cfg = cfg
H, Dk, d = cfg.kda_heads, cfg.kda_head_dim, cfg.hidden
self.q_proj = QuantLinear(d, H * Dk, cfg.group)
self.k_proj = QuantLinear(d, H * Dk, cfg.group)
self.v_proj = QuantLinear(d, H * Dk, cfg.group)
self.g_proj = QuantLinear(d, H * Dk, cfg.group)
self.beta_proj = nn.Linear(d, H, bias=False, dtype=cfg.dtype)
self.conv_w = nn.Parameter(torch.empty(3, H * Dk, cfg.short_conv, dtype=cfg.dtype))
self.o_proj = QuantLinear(H * Dk, d, cfg.group)
self.scale = Dk ** -0.5
class MLA(nn.Module):
def __init__(self, cfg: Config):
super().__init__()
self.cfg = cfg
H, d = cfg.mla_heads, cfg.hidden
self.q_proj = QuantLinear(d, H * (cfg.qk_nope + cfg.qk_rope), cfg.group)
self.kv_a = QuantLinear(d, cfg.kv_lora + cfg.qk_rope, cfg.group)
self.kv_b = QuantLinear(cfg.kv_lora, H * (cfg.qk_nope + cfg.v_head), cfg.group)
self.o_proj = QuantLinear(H * cfg.v_head, d, cfg.group)
self.scale = (cfg.qk_nope + cfg.qk_rope) ** -0.5
class MoE(nn.Module):
def __init__(self, cfg: Config):
super().__init__()
self.cfg = cfg
d, m, E = cfg.hidden, cfg.moe_inter, cfg.n_experts
self.router = nn.Linear(d, E, bias=False, dtype=cfg.dtype)
self.gate = QuantExperts(E, d, m, cfg.group)
self.up = QuantExperts(E, d, m, cfg.group)
self.down = QuantExperts(E, m, d, cfg.group)
self.s_gate = QuantExperts(cfg.n_shared, d, m, cfg.group)
self.s_up = QuantExperts(cfg.n_shared, d, m, cfg.group)
self.s_down = QuantExperts(cfg.n_shared, m, d, cfg.group)
class Block(nn.Module):
def __init__(self, cfg: Config, kind: str):
super().__init__()
self.kind = kind
self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
self.attn = KDA(cfg) if kind == "K" else MLA(cfg)
self.moe = MoE(cfg)
SC_PROJ, SC_CKNEW, SC_KRNEW = 0, 16448, 16960
SC_OUTA, SC_OUTM = 17024, 19328
SC_FLOATS = 42112 + 164 * 2 * 16 * 514
BS_FLOATS = 32 * 576
class Model(nn.Module):
def __init__(self, cfg: Config):
super().__init__()
self.cfg = cfg
self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern)
self._prep_done = False
# ---------------- megakernel prep ---------------- #
def _prep(self):
tab = []
for b in range(3):
blk = self.blocks[b]
a = blk.attn
for prj in (a.q_proj, a.k_proj, a.v_proj, a.g_proj, a.o_proj):
tab += [prj.w_q.data_ptr(), prj.scales.data_ptr(), prj.zeros.data_ptr()]
tab += [a.beta_proj.weight.data_ptr(), a.conv_w.data_ptr(),
blk.attn_norm.data_ptr(), blk.moe_norm.data_ptr(),
blk.moe.router.weight.data_ptr()]
m = blk.moe
for e in (m.gate, m.up, m.down, m.s_gate, m.s_up, m.s_down):
tab += [e.w_q.data_ptr(), e.scales.data_ptr(), e.zeros.data_ptr()]
blk = self.blocks[3]
a = blk.attn
for prj in (a.q_proj, a.kv_a, a.kv_b, a.o_proj):
tab += [prj.w_q.data_ptr(), prj.scales.data_ptr(), prj.zeros.data_ptr()]
tab += [blk.attn_norm.data_ptr(), blk.moe_norm.data_ptr(), blk.moe.router.weight.data_ptr()]
m = blk.moe
for e in (m.gate, m.up, m.down, m.s_gate, m.s_up, m.s_down):
tab += [e.w_q.data_ptr(), e.scales.data_ptr(), e.zeros.data_ptr()]
assert len(tab) <= 152
tab += [0] * (152 - len(tab))
dev = self.blocks[0].attn_norm.device
self._tab = torch.tensor(tab, dtype=torch.int64, device=dev)
self._scratch = torch.zeros(SC_FLOATS, dtype=torch.float32, device=dev)
self._bscratch = torch.zeros(BS_FLOATS, dtype=torch.bfloat16, device=dev)
self._bar = torch.zeros(18 * 2, dtype=torch.int32, device=dev)
self._gen = 0
self._cache_reg = {}
occ = _ext().mega_occupancy()
props = torch.cuda.get_device_properties(0)
self._ncta = props.multi_processor_count * max(1, min(2, int(occ)))
self._prep_done = True
def _caches(self, state):
st3 = state[3]
ckv, kr = st3["c_kv"], st3["k_rope"]
pos = ckv.shape[0]
ent = self._cache_reg.get(ckv.data_ptr())
if ent is not None and ent[2] > pos:
return ent[0], ent[1], 0, 0
cap = pos + 256
dev = ckv.device
ckv_new = torch.empty((cap, 512), dtype=torch.bfloat16, device=dev)
kr_new = torch.empty((cap, 64), dtype=torch.bfloat16, device=dev)
self._cache_reg[ckv_new.data_ptr()] = (ckv_new, kr_new, cap)
return ckv_new, kr_new, pos, (pos + 255) // 256
@staticmethod
def _nc(pos: int) -> int:
G = 32 if pos < 4096 else (64 if pos < 8192 else 128)
return min(164, max(1, -(-pos // G)))
def step(self, hidden, state):
if not self._prep_done:
self._prep()
pos = state[3]["c_kv"].shape[0]
ckv_out, kr_out, copy_rows, ncop = self._caches(state)
ckv_in, kr_in = state[3]["c_kv"], state[3]["k_rope"]
x_out = torch.empty_like(hidden)
args = (self._tab, hidden, x_out,
state[0]["S"], state[1]["S"], state[2]["S"],
state[0]["cq"], state[1]["cq"], state[2]["cq"],
state[0]["ck"], state[1]["ck"], state[2]["ck"],
state[0]["cv"], state[1]["cv"], state[2]["cv"],
ckv_in, kr_in, ckv_out, kr_out,
self._scratch, self._bscratch, self._bar,
pos, self._gen, self._nc(pos), copy_rows, ncop, -1, self._ncta)
self._gen += 1
_ext().mega_step(*args)
state[3] = {"c_kv": ckv_out[: pos + 1], "k_rope": kr_out[: pos + 1]}
return x_out, state
# ---------------- exact eager oracle (debug / comparison only) ---------------- #
def _kda_eager(self, blk, x, st):
cfg = self.cfg
H, Dk = cfg.kda_heads, cfg.kda_head_dim
a = blk.attn
q = a.q_proj(x)
k = a.k_proj(x)
v = a.v_proj(x)
def conv(val, prev, idx):
win = torch.cat([prev, val[None]], dim=0)
w = a.conv_w[idx].float().transpose(0, 1)
out = (win.float() * w).sum(0)
return F.silu(out).to(val.dtype), win[1:]
q, st["cq"] = conv(q, st["cq"], 0)
k, st["ck"] = conv(k, st["ck"], 1)
v, st["cv"] = conv(v, st["cv"], 2)
q = q.view(H, Dk).float() * a.scale
k = k.view(H, Dk).float()
v = v.view(H, Dk).float()
g = (-F.softplus(a.g_proj(x).float())).view(H, Dk)
beta = torch.sigmoid(a.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 a.o_proj(o.reshape(H * Dk).to(torch.bfloat16))
def _mla_eager(self, blk, x, st):
cfg = self.cfg
H = cfg.mla_heads
a = blk.attn
pos = st["c_kv"].shape[0]
q = a.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 = a.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 = a.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())) * a.scale
p = torch.softmax(scores, dim=0)
o = torch.einsum("lh,lhd->hd", p, v)
return a.o_proj(o.reshape(H * cfg.v_head).to(torch.bfloat16))
def _moe_eager(self, blk, x):
cfg = self.cfg
m = blk.moe
probs = torch.softmax(m.router(x).float(), dim=-1)
w, idx = torch.topk(probs, cfg.n_active)
w = w / (w.sum() + 1e-9) * cfg.routed_scaling
def ffn(xf, eg, eu, ed, e):
h = F.silu(xf @ eg.weight_bf(e).float()) * (xf @ eu.weight_bf(e).float())
return h @ ed.weight_bf(e).float()
xf = x.float()
out = x.new_zeros(cfg.hidden, dtype=torch.float32)
for j in range(cfg.n_active):
out = out + w[j] * ffn(xf, m.gate, m.up, m.down, int(idx[j]))
for sidx in range(cfg.n_shared):
out = out + ffn(xf, m.s_gate, m.s_up, m.s_down, sidx)
return out.to(torch.bfloat16)
def _step_eager(self, hidden, state):
for i, blk in enumerate(self.blocks):
if blk.kind == "K":
h = hidden + self._kda_eager(blk, _rmsnorm(hidden, blk.attn_norm), state[i])
else:
h = hidden + self._mla_eager(blk, _rmsnorm(hidden, blk.attn_norm), state[i])
hidden = h + self._moe_eager(blk, _rmsnorm(h, blk.moe_norm))
return hidden, state
20260715_204329_kinetic-claude_kinetic-0715_02_kimi_linear_decode