"""W4A16 Kimi-Linear Hybrid Decode Single Fused Megakernel. Single GPU kernel launch per decode step. Fuses the entire 4-layer motif: - 3 KDA layers (RMSNorm, q/k/v/g/beta GEMVs, short causal conv, recurrence on S, o_proj) - 1 MLA layer (RMSNorm, q/kv_a GEMVs, RoPE, cache update, absorbed Q/O, attention, o_proj) - 4 MoE blocks (RMSNorm, router, top-8 + 1 shared expert gate/up/down GEMVs) - All residual additions into a single cooperative CUDA __global__ kernel launch. """ from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path import math import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.cpp_extension import load_inline EPS = 1.0e-6 GROUP_SIZE = 128 CUDA_SRC = r""" #include #include #include #include #include namespace cg = cooperative_groups; #define HIDDEN 2304 #define KDA_HEADS 32 #define KDA_HEAD_DIM 128 #define KDA_CHANNELS 4096 #define CONV_LEN 4 #define MLA_HEADS 32 #define KV_LORA 512 #define QK_NOPE 128 #define QK_ROPE 64 #define V_HEAD 128 #define MOE_INTER 1024 #define N_EXPERTS 64 #define N_ACTIVE 8 #define N_SHARED 1 #define ROUTED_SCALING 2.446f #define EPS 1e-6f struct QuantWeight { const uint8_t* w_q; const __nv_bfloat16* scales; const __nv_bfloat16* zeros; int in_f; int out_f; }; struct KDAWeights { QuantWeight q_proj; QuantWeight k_proj; QuantWeight v_proj; QuantWeight g_proj; const __nv_bfloat16* beta_proj; const __nv_bfloat16* conv_w; QuantWeight o_proj; }; struct MLAWeights { QuantWeight q_proj; QuantWeight kv_a; QuantWeight kv_b; QuantWeight o_proj; }; struct MoEWeights { const __nv_bfloat16* router; QuantWeight gate; QuantWeight up; QuantWeight down; QuantWeight s_gate; QuantWeight s_up; QuantWeight s_down; }; struct BlockWeights { const __nv_bfloat16* attn_norm; const __nv_bfloat16* moe_norm; KDAWeights kda; MLAWeights mla; MoEWeights moe; }; struct MegakernelParams { BlockWeights blocks[4]; }; struct KDAState { float* S; __nv_bfloat16* cq; __nv_bfloat16* ck; __nv_bfloat16* cv; }; struct MLAState { const __nv_bfloat16* old_c_kv; const __nv_bfloat16* old_k_rope; __nv_bfloat16* new_c_kv; __nv_bfloat16* new_k_rope; int pos; }; struct Scratchpad { float x[HIDDEN]; float x_norm[HIDDEN]; float x_group_sum[32]; float gemv_out[18432]; float beta[32]; float kda_q[KDA_CHANNELS]; float kda_k[KDA_CHANNELS]; float kda_v[KDA_CHANNELS]; float kda_g[KDA_CHANNELS]; float kda_o[KDA_CHANNELS]; float attn_out[HIDDEN]; float router_logits[N_EXPERTS]; float router_probs[N_EXPERTS]; int topk_idx[N_ACTIVE]; float topk_w[N_ACTIVE]; float moe_inter[9 * MOE_INTER]; float moe_out[HIDDEN]; float mla_q[6144]; float mla_kv[576]; float q_abs[MLA_HEADS * KV_LORA]; float q_rope[MLA_HEADS * QK_ROPE]; float c_att[MLA_HEADS * KV_LORA]; float scores[MLA_HEADS * 16384]; }; __device__ void device_rmsnorm( const float* __restrict__ x, const __nv_bfloat16* __restrict__ weight, float* __restrict__ out, cg::grid_group& grid ) { __shared__ float sh_val[32]; int tid = threadIdx.x; if (blockIdx.x == 0) { float sum_sq = 0.0f; for (int i = tid; i < HIDDEN; i += blockDim.x) { float val = x[i]; sum_sq += val * val; } #pragma unroll for (int mask = 16; mask > 0; mask >>= 1) { sum_sq += __shfl_down_sync(0xffffffff, sum_sq, mask); } if ((tid % 32) == 0) { sh_val[tid / 32] = sum_sq; } __syncthreads(); if (tid == 0) { float total = 0.0f; int num_warps = blockDim.x / 32; for (int w = 0; w < num_warps; ++w) total += sh_val[w]; float rsqrt_val = rsqrtf(total / (float)HIDDEN + EPS); sh_val[0] = rsqrt_val; } __syncthreads(); float r = sh_val[0]; for (int i = tid; i < HIDDEN; i += blockDim.x) { float w = __bfloat162float(weight[i]); out[i] = x[i] * r * w; } } grid.sync(); } __device__ void compute_x_group_sums( const float* __restrict__ x, float* __restrict__ x_group_sums, int in_f, cg::grid_group& grid ) { if (blockIdx.x == 0) { int num_groups = in_f / 128; int tid = threadIdx.x; if (tid < num_groups) { float s = 0.0f; int base = tid * 128; #pragma unroll 8 for (int i = 0; i < 128; ++i) { s += x[base + i]; } x_group_sums[tid] = s; } } grid.sync(); } __device__ inline float compute_col_slice( const float* __restrict__ sh_x, const float* __restrict__ x_group_sums, const uint8_t* __restrict__ w_q, const __nv_bfloat16* __restrict__ scales, const __nv_bfloat16* __restrict__ zeros, int in_f, int out_f, int col, int g_start, int g_end ) { float acc = 0.0f; for (int g = g_start; g < g_end; ++g) { float scale = __bfloat162float(scales[g * out_f + col]); float zero = __bfloat162float(zeros[g * out_f + col]); float x_sum = x_group_sums[g]; int k_start = g * 128; const float2* x_pairs = (const float2*)&sh_x[k_start]; const uint8_t* w_ptr = &w_q[(k_start / 2) * out_f + col]; float group_acc = 0.0f; #pragma unroll 8 for (int p = 0; p < 64; ++p) { uint8_t byte_val = *w_ptr; w_ptr += out_f; float w0 = (float)(byte_val & 0x0F); float w1 = (float)(byte_val >> 4); float2 x_val = x_pairs[p]; group_acc += x_val.x * w0 + x_val.y * w1; } acc += (group_acc - zero * x_sum) * scale; } return acc; } __global__ void kimi_decode_megakernel_opt( const __nv_bfloat16* __restrict__ hidden_in, __nv_bfloat16* __restrict__ hidden_out, KDAState st_kda0, KDAState st_kda1, KDAState st_kda2, MLAState st_mla, const MegakernelParams* __restrict__ params, Scratchpad* __restrict__ scratch ) { cg::grid_group grid = cg::this_grid(); int global_tid = blockIdx.x * blockDim.x + threadIdx.x; int num_threads = gridDim.x * blockDim.x; int wid = global_tid / 32; int lane = threadIdx.x % 32; int total_warps = num_threads / 32; __shared__ float sh_x[HIDDEN]; for (int i = global_tid; i < HIDDEN; i += num_threads) { scratch->x[i] = __bfloat162float(hidden_in[i]); } grid.sync(); #pragma unroll 1 for (int b = 0; b < 4; ++b) { const BlockWeights& blk = params->blocks[b]; // 1. Attention RMSNorm device_rmsnorm(scratch->x, blk.attn_norm, scratch->x_norm, grid); for (int i = threadIdx.x; i < HIDDEN; i += blockDim.x) { sh_x[i] = scratch->x_norm[i]; } __syncthreads(); if (b < 3) { // --------------------------------------------------------------- // KDA Block // --------------------------------------------------------------- KDAState& st_kda = (b == 0) ? st_kda0 : ((b == 1) ? st_kda1 : st_kda2); compute_x_group_sums(scratch->x_norm, scratch->x_group_sum, HIDDEN, grid); for (int i = global_tid; i < 16384; i += num_threads) { scratch->gemv_out[i] = 0.0f; } grid.sync(); for (int task = wid; task < 1024; task += total_warps) { int tile = task / 2; int k_slice = task % 2; int col = tile * 32 + lane; int proj = col / KDA_CHANNELS; int c = col % KDA_CHANNELS; const QuantWeight& qw = (proj == 0) ? blk.kda.q_proj : ((proj == 1) ? blk.kda.k_proj : ((proj == 2) ? blk.kda.v_proj : blk.kda.g_proj)); int g_start = k_slice * 9; int g_end = g_start + 9; float val = compute_col_slice(sh_x, scratch->x_group_sum, qw.w_q, qw.scales, qw.zeros, HIDDEN, KDA_CHANNELS, c, g_start, g_end); atomicAdd(&scratch->gemv_out[col], val); } if (wid == 0 && lane < 32) { int c = lane; float b_acc = 0.0f; #pragma unroll 8 for (int k = 0; k < HIDDEN; ++k) { b_acc += sh_x[k] * __bfloat162float(blk.kda.beta_proj[c * HIDDEN + k]); } scratch->beta[c] = b_acc; } grid.sync(); // Conv & Act for (int c = global_tid; c < KDA_CHANNELS; c += num_threads) { float val_q = scratch->gemv_out[c]; float val_k = scratch->gemv_out[KDA_CHANNELS + c]; float val_v = scratch->gemv_out[2 * KDA_CHANNELS + c]; float val_g = scratch->gemv_out[3 * KDA_CHANNELS + c]; float pq0 = __bfloat162float(st_kda.cq[0 * KDA_CHANNELS + c]); float pq1 = __bfloat162float(st_kda.cq[1 * KDA_CHANNELS + c]); float pq2 = __bfloat162float(st_kda.cq[2 * KDA_CHANNELS + c]); float wq0 = __bfloat162float(blk.kda.conv_w[(0 * KDA_CHANNELS + c) * 4 + 0]); float wq1 = __bfloat162float(blk.kda.conv_w[(0 * KDA_CHANNELS + c) * 4 + 1]); float wq2 = __bfloat162float(blk.kda.conv_w[(0 * KDA_CHANNELS + c) * 4 + 2]); float wq3 = __bfloat162float(blk.kda.conv_w[(0 * KDA_CHANNELS + c) * 4 + 3]); float out_q = pq0 * wq0 + pq1 * wq1 + pq2 * wq2 + val_q * wq3; out_q = (out_q / (1.0f + expf(-out_q))) * 0.08838834764831845f; st_kda.cq[0 * KDA_CHANNELS + c] = __float2bfloat16(pq1); st_kda.cq[1 * KDA_CHANNELS + c] = __float2bfloat16(pq2); st_kda.cq[2 * KDA_CHANNELS + c] = __float2bfloat16(val_q); scratch->kda_q[c] = out_q; float pk0 = __bfloat162float(st_kda.ck[0 * KDA_CHANNELS + c]); float pk1 = __bfloat162float(st_kda.ck[1 * KDA_CHANNELS + c]); float pk2 = __bfloat162float(st_kda.ck[2 * KDA_CHANNELS + c]); float wk0 = __bfloat162float(blk.kda.conv_w[(1 * KDA_CHANNELS + c) * 4 + 0]); float wk1 = __bfloat162float(blk.kda.conv_w[(1 * KDA_CHANNELS + c) * 4 + 1]); float wk2 = __bfloat162float(blk.kda.conv_w[(1 * KDA_CHANNELS + c) * 4 + 2]); float wk3 = __bfloat162float(blk.kda.conv_w[(1 * KDA_CHANNELS + c) * 4 + 3]); float out_k = pk0 * wk0 + pk1 * wk1 + pk2 * wk2 + val_k * wk3; out_k = out_k / (1.0f + expf(-out_k)); st_kda.ck[0 * KDA_CHANNELS + c] = __float2bfloat16(pk1); st_kda.ck[1 * KDA_CHANNELS + c] = __float2bfloat16(pk2); st_kda.ck[2 * KDA_CHANNELS + c] = __float2bfloat16(val_k); scratch->kda_k[c] = out_k; float pv0 = __bfloat162float(st_kda.cv[0 * KDA_CHANNELS + c]); float pv1 = __bfloat162float(st_kda.cv[1 * KDA_CHANNELS + c]); float pv2 = __bfloat162float(st_kda.cv[2 * KDA_CHANNELS + c]); float wv0 = __bfloat162float(blk.kda.conv_w[(2 * KDA_CHANNELS + c) * 4 + 0]); float wv1 = __bfloat162float(blk.kda.conv_w[(2 * KDA_CHANNELS + c) * 4 + 1]); float wv2 = __bfloat162float(blk.kda.conv_w[(2 * KDA_CHANNELS + c) * 4 + 2]); float wv3 = __bfloat162float(blk.kda.conv_w[(2 * KDA_CHANNELS + c) * 4 + 3]); float out_v = pv0 * wv0 + pv1 * wv1 + pv2 * wv2 + val_v * wv3; out_v = out_v / (1.0f + expf(-out_v)); st_kda.cv[0 * KDA_CHANNELS + c] = __float2bfloat16(pv1); st_kda.cv[1 * KDA_CHANNELS + c] = __float2bfloat16(pv2); st_kda.cv[2 * KDA_CHANNELS + c] = __float2bfloat16(val_v); scratch->kda_v[c] = out_v; float g_val = (val_g > 20.0f) ? -val_g : -log1pf(expf(val_g)); scratch->kda_g[c] = g_val; if (c < 32) { scratch->beta[c] = 1.0f / (1.0f + expf(-scratch->beta[c])); } } grid.sync(); // KDA Recurrence on S if (blockIdx.x < KDA_HEADS) { int h = blockIdx.x; int tid = threadIdx.x; __shared__ float sh_k[128]; __shared__ float sh_q[128]; __shared__ float sh_exp_g[128]; __shared__ float sh_b; if (tid < 128) { sh_k[tid] = scratch->kda_k[h * 128 + tid]; sh_q[tid] = scratch->kda_q[h * 128 + tid]; sh_exp_g[tid] = expf(scratch->kda_g[h * 128 + tid]); } if (tid == 0) { sh_b = scratch->beta[h]; } __syncthreads(); if (tid < 128) { int j = tid; float* col_S_ptr = &st_kda.S[h * 16384 + j]; float pred_j = 0.0f; float col_S[128]; #pragma unroll 4 for (int i = 0; i < 128; ++i) { float s_val = col_S_ptr[i * 128] * sh_exp_g[i]; col_S[i] = s_val; pred_j += s_val * sh_k[i]; } float v_j = scratch->kda_v[h * 128 + j]; float diff_j = v_j - pred_j; float beta_h = sh_b; float o_j = 0.0f; #pragma unroll 4 for (int i = 0; i < 128; ++i) { float s_up = col_S[i] + beta_h * sh_k[i] * diff_j; col_S_ptr[i * 128] = s_up; o_j += s_up * sh_q[i]; } scratch->kda_o[h * 128 + j] = o_j; } } grid.sync(); // KDA o_proj compute_x_group_sums(scratch->kda_o, scratch->x_group_sum, KDA_CHANNELS, grid); for (int i = global_tid; i < HIDDEN; i += num_threads) { scratch->attn_out[i] = 0.0f; } grid.sync(); for (int task = wid; task < 576; task += total_warps) { int tile = task / 8; int k_slice = task % 8; int col = tile * 32 + lane; if (col < HIDDEN) { int g_start = k_slice * 4; int g_end = g_start + 4; float val = compute_col_slice(scratch->kda_o, scratch->x_group_sum, blk.kda.o_proj.w_q, blk.kda.o_proj.scales, blk.kda.o_proj.zeros, KDA_CHANNELS, HIDDEN, col, g_start, g_end); atomicAdd(&scratch->attn_out[col], val); } } grid.sync(); } else { // --------------------------------------------------------------- // MLA Block // --------------------------------------------------------------- compute_x_group_sums(scratch->x_norm, scratch->x_group_sum, HIDDEN, grid); for (int i = global_tid; i < 6144; i += num_threads) scratch->mla_q[i] = 0.0f; grid.sync(); for (int task = wid; task < 768; task += total_warps) { int tile = task / 4; int k_slice = task % 4; int col = tile * 32 + lane; int g_start = (k_slice < 3) ? (k_slice * 4) : 12; int g_end = (k_slice < 3) ? (g_start + 4) : 18; float val = compute_col_slice(sh_x, scratch->x_group_sum, blk.mla.q_proj.w_q, blk.mla.q_proj.scales, blk.mla.q_proj.zeros, HIDDEN, 6144, col, g_start, g_end); atomicAdd(&scratch->mla_q[col], val); } for (int i = global_tid; i < 576; i += num_threads) scratch->mla_kv[i] = 0.0f; grid.sync(); for (int task = wid; task < 144; task += total_warps) { int tile = task / 8; int k_slice = task % 8; int col = tile * 32 + lane; if (col < 576) { int g_start = (k_slice < 7) ? (k_slice * 2) : 14; int g_end = (k_slice < 7) ? (g_start + 2) : 18; float val = compute_col_slice(sh_x, scratch->x_group_sum, blk.mla.kv_a.w_q, blk.mla.kv_a.scales, blk.mla.kv_a.zeros, HIDDEN, 576, col, g_start, g_end); atomicAdd(&scratch->mla_kv[col], val); } } grid.sync(); // MLA RoPE & Cache Update int pos = st_mla.pos; int L = pos + 1; int total_c_kv_elements = pos * KV_LORA; for (int i = global_tid; i < total_c_kv_elements; i += num_threads) { st_mla.new_c_kv[i] = st_mla.old_c_kv[i]; } int total_k_rope_elements = pos * QK_ROPE; for (int i = global_tid; i < total_k_rope_elements; i += num_threads) { st_mla.new_k_rope[i] = st_mla.old_k_rope[i]; } if (blockIdx.x == 0 && threadIdx.x < 32) { int j = threadIdx.x; float inv = 1.0f / powf(10000.0f, (2.0f * j) / 64.0f); float ang = pos * inv; float cos_val = cosf(ang); float sin_val = sinf(ang); float k_even = scratch->mla_kv[512 + 2 * j]; float k_odd = scratch->mla_kv[512 + 2 * j + 1]; st_mla.new_k_rope[pos * 64 + 2 * j] = __float2bfloat16(k_even * cos_val - k_odd * sin_val); st_mla.new_k_rope[pos * 64 + 2 * j + 1] = __float2bfloat16(k_odd * cos_val + k_even * sin_val); #pragma unroll for (int m = 0; m < 16; ++m) { int r = j * 16 + m; st_mla.new_c_kv[pos * 512 + r] = __float2bfloat16(scratch->mla_kv[r]); } } for (int h = wid; h < MLA_HEADS; h += total_warps) { int j = lane; float inv = 1.0f / powf(10000.0f, (2.0f * j) / 64.0f); float ang = pos * inv; float cos_val = cosf(ang); float sin_val = sinf(ang); int q_base = h * (QK_NOPE + QK_ROPE) + QK_NOPE; float q_even = scratch->mla_q[q_base + 2 * j]; float q_odd = scratch->mla_q[q_base + 2 * j + 1]; scratch->q_rope[h * 64 + 2 * j] = q_even * cos_val - q_odd * sin_val; scratch->q_rope[h * 64 + 2 * j + 1] = q_odd * cos_val + q_even * sin_val; } grid.sync(); // MLA Absorbed Q: q_abs[h, r] for (int idx = wid; idx < 16384; idx += total_warps) { int h = idx / 512; int r = idx % 512; int q_base = h * (QK_NOPE + QK_ROPE); int row_idx = r / 2; bool is_even = (r % 2 == 0); int g = r / 128; int col_base = h * 256; float dot = 0.0f; #pragma unroll for (int m = 0; m < 4; ++m) { int d = lane * 4 + m; int col = col_base + d; uint8_t b = blk.mla.kv_b.w_q[row_idx * 8192 + col]; float w_val = is_even ? (float)(b & 0x0F) : (float)(b >> 4); float zero = __bfloat162float(blk.mla.kv_b.zeros[g * 8192 + col]); float scale = __bfloat162float(blk.mla.kv_b.scales[g * 8192 + col]); dot += scratch->mla_q[q_base + d] * ((w_val - zero) * scale); } #pragma unroll for (int mask = 16; mask > 0; mask >>= 1) { dot += __shfl_down_sync(0xffffffff, dot, mask); } if (lane == 0) { scratch->q_abs[h * 512 + r] = dot; } } grid.sync(); // MLA Attention Scores, Softmax & Weighted Sum if (blockIdx.x < MLA_HEADS) { int h = blockIdx.x; int tid = threadIdx.x; int warp_id = tid / 32; int lane_id = tid % 32; __shared__ float sh_q_abs[512]; __shared__ float sh_q_rope[64]; __shared__ float sh_max; __shared__ float sh_sum; for (int i = tid; i < 512; i += blockDim.x) sh_q_abs[i] = scratch->q_abs[h * 512 + i]; for (int i = tid; i < 64; i += blockDim.x) sh_q_rope[i] = scratch->q_rope[h * 64 + i]; __syncthreads(); float local_max = -1e30f; float scale_attn = 0.07216878364870322f; for (int l = warp_id; l < L; l += 8) { float dot = 0.0f; #pragma unroll for (int i = 0; i < 16; ++i) { int r = lane_id * 16 + i; float c_val = __bfloat162float(st_mla.new_c_kv[l * 512 + r]); dot += c_val * sh_q_abs[r]; } #pragma unroll for (int i = 0; i < 2; ++i) { int d = lane_id * 2 + i; float k_val = __bfloat162float(st_mla.new_k_rope[l * 64 + d]); dot += k_val * sh_q_rope[d]; } #pragma unroll for (int mask = 16; mask > 0; mask >>= 1) { dot += __shfl_down_sync(0xffffffff, dot, mask); } if (lane_id == 0) { float sc = dot * scale_attn; scratch->scores[h * 16384 + l] = sc; if (sc > local_max) local_max = sc; } } __shared__ float warp_max[8]; if (lane_id == 0) warp_max[warp_id] = local_max; __syncthreads(); if (tid == 0) { float m = warp_max[0]; for (int w = 1; w < 8; ++w) if (warp_max[w] > m) m = warp_max[w]; sh_max = m; } __syncthreads(); float global_max = sh_max; float local_sum = 0.0f; for (int l = tid; l < L; l += blockDim.x) { float p = expf(scratch->scores[h * 16384 + l] - global_max); scratch->scores[h * 16384 + l] = p; local_sum += p; } #pragma unroll for (int mask = 16; mask > 0; mask >>= 1) { local_sum += __shfl_down_sync(0xffffffff, local_sum, mask); } __shared__ float warp_sum[8]; if (lane_id == 0) warp_sum[warp_id] = local_sum; __syncthreads(); if (tid == 0) { float s = 0.0f; for (int w = 0; w < 8; ++w) s += warp_sum[w]; sh_sum = (s > 0.0f) ? (1.0f / s) : 0.0f; } __syncthreads(); float inv_sum = sh_sum; float acc0 = 0.0f; float acc1 = 0.0f; int r0 = tid; int r1 = tid + 256; #pragma unroll 4 for (int l = 0; l < L; ++l) { float p = scratch->scores[h * 16384 + l] * inv_sum; const __nv_bfloat16* c_l = &st_mla.new_c_kv[l * 512]; acc0 += p * __bfloat162float(c_l[r0]); acc1 += p * __bfloat162float(c_l[r1]); } scratch->c_att[h * 512 + r0] = acc0; scratch->c_att[h * 512 + r1] = acc1; } grid.sync(); // MLA Absorb O for (int i = global_tid; i < 4096; i += num_threads) scratch->kda_o[i] = 0.0f; grid.sync(); for (int task = wid; task < 256; task += total_warps) { int tile = task / 2; int k_slice = task % 2; int col_idx = tile * 32 + lane; int h = col_idx / 128; int d = col_idx % 128; int col = h * 256 + 128 + d; int g_start = k_slice * 2; int g_end = g_start + 2; float acc = 0.0f; for (int g = g_start; g < g_end; ++g) { float scale = __bfloat162float(blk.mla.kv_b.scales[g * 8192 + col]); float zero = __bfloat162float(blk.mla.kv_b.zeros[g * 8192 + col]); int r_start = g * 128; int row_start = r_start / 2; #pragma unroll 4 for (int p = 0; p < 64; ++p) { uint8_t b = blk.mla.kv_b.w_q[(row_start + p) * 8192 + col]; float w0 = ((float)(b & 0x0F) - zero) * scale; float w1 = ((float)(b >> 4) - zero) * scale; acc += scratch->c_att[h * 512 + r_start + 2 * p] * w0 + scratch->c_att[h * 512 + r_start + 2 * p + 1] * w1; } } atomicAdd(&scratch->kda_o[col_idx], acc); } grid.sync(); // MLA o_proj compute_x_group_sums(scratch->kda_o, scratch->x_group_sum, 4096, grid); for (int i = global_tid; i < HIDDEN; i += num_threads) scratch->attn_out[i] = 0.0f; grid.sync(); for (int task = wid; task < 576; task += total_warps) { int tile = task / 8; int k_slice = task % 8; int col = tile * 32 + lane; if (col < HIDDEN) { int g_start = k_slice * 4; int g_end = g_start + 4; float val = compute_col_slice(scratch->kda_o, scratch->x_group_sum, blk.mla.o_proj.w_q, blk.mla.o_proj.scales, blk.mla.o_proj.zeros, 4096, HIDDEN, col, g_start, g_end); atomicAdd(&scratch->attn_out[col], val); } } grid.sync(); } // Residual Add 1: x = x + attn_out for (int i = global_tid; i < HIDDEN; i += num_threads) { scratch->x[i] += scratch->attn_out[i]; } grid.sync(); // ------------------------------------------------------------------- // MoE // ------------------------------------------------------------------- device_rmsnorm(scratch->x, blk.moe_norm, scratch->x_norm, grid); for (int i = threadIdx.x; i < HIDDEN; i += blockDim.x) { sh_x[i] = scratch->x_norm[i]; } __syncthreads(); if (blockIdx.x == 0) { int tid = threadIdx.x; if (tid < N_EXPERTS) { float logit = 0.0f; #pragma unroll 8 for (int k = 0; k < HIDDEN; ++k) { logit += sh_x[k] * __bfloat162float(blk.moe.router[tid * HIDDEN + k]); } scratch->router_logits[tid] = logit; } __syncthreads(); if (tid == 0) { float m = scratch->router_logits[0]; for (int e = 1; e < N_EXPERTS; ++e) { if (scratch->router_logits[e] > m) m = scratch->router_logits[e]; } float s = 0.0f; for (int e = 0; e < N_EXPERTS; ++e) { float p = expf(scratch->router_logits[e] - m); scratch->router_probs[e] = p; s += p; } float inv_s = (s > 0.0f) ? (1.0f / s) : 0.0f; for (int e = 0; e < N_EXPERTS; ++e) { scratch->router_probs[e] *= inv_s; } for (int j = 0; j < N_ACTIVE; ++j) { float best_val = -1.0f; int best_idx = -1; for (int e = 0; e < N_EXPERTS; ++e) { float p = scratch->router_probs[e]; bool selected = false; for (int prev = 0; prev < j; ++prev) { if (scratch->topk_idx[prev] == e) { selected = true; break; } } if (!selected && p > best_val) { best_val = p; best_idx = e; } } scratch->topk_idx[j] = best_idx; scratch->topk_w[j] = best_val; } float sum_topk = 0.0f; for (int j = 0; j < N_ACTIVE; ++j) sum_topk += scratch->topk_w[j]; float norm_factor = ROUTED_SCALING / (sum_topk + 1e-9f); for (int j = 0; j < N_ACTIVE; ++j) scratch->topk_w[j] *= norm_factor; } } grid.sync(); // MoE Gate & Up compute_x_group_sums(scratch->x_norm, scratch->x_group_sum, HIDDEN, grid); for (int i = global_tid; i < 18432; i += num_threads) scratch->gemv_out[i] = 0.0f; grid.sync(); for (int task = wid; task < 1152; task += total_warps) { int tile = task / 2; int k_slice = task % 2; int col_id = tile * 32 + lane; int e = col_id / 2048; int rem = col_id % 2048; bool is_gate = (rem < 1024); int c = is_gate ? rem : (rem - 1024); int expert_id = (e < N_ACTIVE) ? scratch->topk_idx[e] : 0; const QuantWeight& qw = (e < N_ACTIVE) ? (is_gate ? blk.moe.gate : blk.moe.up) : (is_gate ? blk.moe.s_gate : blk.moe.s_up); int w_offset = expert_id * (HIDDEN / 2) * MOE_INTER; int s_offset = expert_id * (HIDDEN / 128) * MOE_INTER; int g_start = k_slice * 9; int g_end = g_start + 9; float val = compute_col_slice(sh_x, scratch->x_group_sum, qw.w_q + w_offset, qw.scales + s_offset, qw.zeros + s_offset, HIDDEN, MOE_INTER, c, g_start, g_end); atomicAdd(&scratch->gemv_out[col_id], val); } grid.sync(); for (int i = global_tid; i < 9 * MOE_INTER; i += num_threads) { int e = i / MOE_INTER; int c = i % MOE_INTER; float g_val = scratch->gemv_out[e * 2048 + c]; float u_val = scratch->gemv_out[e * 2048 + 1024 + c]; float silu_g = g_val / (1.0f + expf(-g_val)); scratch->moe_inter[e * MOE_INTER + c] = silu_g * u_val; } for (int i = global_tid; i < HIDDEN; i += num_threads) { scratch->moe_out[i] = 0.0f; } grid.sync(); // MoE Down for (int task = wid; task < 1296; task += total_warps) { int tile = task / 2; int k_slice = task % 2; int col_id = tile * 32 + lane; int e = col_id / HIDDEN; int d = col_id % HIDDEN; int expert_id = (e < N_ACTIVE) ? scratch->topk_idx[e] : 0; float w_e = (e < N_ACTIVE) ? scratch->topk_w[e] : 1.0f; const QuantWeight& qw = (e < N_ACTIVE) ? blk.moe.down : blk.moe.s_down; int w_offset = expert_id * (MOE_INTER / 2) * HIDDEN; int s_offset = expert_id * (MOE_INTER / 128) * HIDDEN; float acc = 0.0f; const float* expert_inter = &scratch->moe_inter[e * MOE_INTER]; const float2* inter_pairs = (const float2*)expert_inter; int g_start = k_slice * 4; int g_end = g_start + 4; for (int g = g_start; g < g_end; ++g) { float scale = __bfloat162float(qw.scales[s_offset + g * HIDDEN + d]); float zero = __bfloat162float(qw.zeros[s_offset + g * HIDDEN + d]); int k_start = g * 128; const uint8_t* w_ptr = &qw.w_q[w_offset + (k_start / 2) * HIDDEN + d]; #pragma unroll 8 for (int p = 0; p < 64; ++p) { uint8_t byte_val = *w_ptr; w_ptr += HIDDEN; float w0 = ((float)(byte_val & 0x0F) - zero) * scale; float w1 = ((float)(byte_val >> 4) - zero) * scale; float2 val = inter_pairs[(k_start / 2) + p]; acc += val.x * w0 + val.y * w1; } } atomicAdd(&scratch->moe_out[d], w_e * acc); } grid.sync(); for (int i = global_tid; i < HIDDEN; i += num_threads) { scratch->x[i] += scratch->moe_out[i]; } grid.sync(); } for (int i = global_tid; i < HIDDEN; i += num_threads) { hidden_out[i] = __float2bfloat16(scratch->x[i]); } } void launch_kimi_decode_megakernel_opt( torch::Tensor hidden_in, torch::Tensor hidden_out, torch::Tensor s0, torch::Tensor cq0, torch::Tensor ck0, torch::Tensor cv0, torch::Tensor s1, torch::Tensor cq1, torch::Tensor ck1, torch::Tensor cv1, torch::Tensor s2, torch::Tensor cq2, torch::Tensor ck2, torch::Tensor cv2, torch::Tensor old_c_kv, torch::Tensor old_k_rope, torch::Tensor new_c_kv, torch::Tensor new_k_rope, int pos, torch::Tensor params_tensor, torch::Tensor scratch_tensor ) { KDAState st0 = { s0.data_ptr(), (__nv_bfloat16*)cq0.data_ptr(), (__nv_bfloat16*)ck0.data_ptr(), (__nv_bfloat16*)cv0.data_ptr() }; KDAState st1 = { s1.data_ptr(), (__nv_bfloat16*)cq1.data_ptr(), (__nv_bfloat16*)ck1.data_ptr(), (__nv_bfloat16*)cv1.data_ptr() }; KDAState st2 = { s2.data_ptr(), (__nv_bfloat16*)cq2.data_ptr(), (__nv_bfloat16*)ck2.data_ptr(), (__nv_bfloat16*)cv2.data_ptr() }; MLAState st_mla = { (const __nv_bfloat16*)old_c_kv.data_ptr(), (const __nv_bfloat16*)old_k_rope.data_ptr(), (__nv_bfloat16*)new_c_kv.data_ptr(), (__nv_bfloat16*)new_k_rope.data_ptr(), pos }; const MegakernelParams* params_ptr = (const MegakernelParams*)params_tensor.data_ptr(); Scratchpad* scratch_ptr = (Scratchpad*)scratch_tensor.data_ptr(); const __nv_bfloat16* in_ptr = (const __nv_bfloat16*)hidden_in.data_ptr(); __nv_bfloat16* out_ptr = (__nv_bfloat16*)hidden_out.data_ptr(); void* args[] = { &in_ptr, &out_ptr, &st0, &st1, &st2, &st_mla, ¶ms_ptr, &scratch_ptr }; cudaLaunchCooperativeKernel( (void*)kimi_decode_megakernel_opt, dim3(188), dim3(256), args, 0, c10::cuda::getCurrentCUDAStream(hidden_in.get_device()) ); } void setup_kda_block( torch::Tensor params_tensor, int b, torch::Tensor attn_norm, torch::Tensor moe_norm, torch::Tensor q_wq, torch::Tensor q_s, torch::Tensor q_z, torch::Tensor k_wq, torch::Tensor k_s, torch::Tensor k_z, torch::Tensor v_wq, torch::Tensor v_s, torch::Tensor v_z, torch::Tensor g_wq, torch::Tensor g_s, torch::Tensor g_z, torch::Tensor beta_w, torch::Tensor conv_w, torch::Tensor o_wq, torch::Tensor o_s, torch::Tensor o_z, torch::Tensor router_w, torch::Tensor gate_wq, torch::Tensor gate_s, torch::Tensor gate_z, torch::Tensor up_wq, torch::Tensor up_s, torch::Tensor up_z, torch::Tensor down_wq, torch::Tensor down_s, torch::Tensor down_z, torch::Tensor s_gate_wq, torch::Tensor s_gate_s, torch::Tensor s_gate_z, torch::Tensor s_up_wq, torch::Tensor s_up_s, torch::Tensor s_up_z, torch::Tensor s_down_wq, torch::Tensor s_down_s, torch::Tensor s_down_z ) { MegakernelParams* p = (MegakernelParams*)params_tensor.data_ptr(); BlockWeights& blk = p->blocks[b]; blk.attn_norm = (const __nv_bfloat16*)attn_norm.data_ptr(); blk.moe_norm = (const __nv_bfloat16*)moe_norm.data_ptr(); blk.kda.q_proj = { q_wq.data_ptr(), (const __nv_bfloat16*)q_s.data_ptr(), (const __nv_bfloat16*)q_z.data_ptr(), HIDDEN, KDA_CHANNELS }; blk.kda.k_proj = { k_wq.data_ptr(), (const __nv_bfloat16*)k_s.data_ptr(), (const __nv_bfloat16*)k_z.data_ptr(), HIDDEN, KDA_CHANNELS }; blk.kda.v_proj = { v_wq.data_ptr(), (const __nv_bfloat16*)v_s.data_ptr(), (const __nv_bfloat16*)v_z.data_ptr(), HIDDEN, KDA_CHANNELS }; blk.kda.g_proj = { g_wq.data_ptr(), (const __nv_bfloat16*)g_s.data_ptr(), (const __nv_bfloat16*)g_z.data_ptr(), HIDDEN, KDA_CHANNELS }; blk.kda.beta_proj = (const __nv_bfloat16*)beta_w.data_ptr(); blk.kda.conv_w = (const __nv_bfloat16*)conv_w.data_ptr(); blk.kda.o_proj = { o_wq.data_ptr(), (const __nv_bfloat16*)o_s.data_ptr(), (const __nv_bfloat16*)o_z.data_ptr(), KDA_CHANNELS, HIDDEN }; blk.moe.router = (const __nv_bfloat16*)router_w.data_ptr(); blk.moe.gate = { gate_wq.data_ptr(), (const __nv_bfloat16*)gate_s.data_ptr(), (const __nv_bfloat16*)gate_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.up = { up_wq.data_ptr(), (const __nv_bfloat16*)up_s.data_ptr(), (const __nv_bfloat16*)up_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.down = { down_wq.data_ptr(), (const __nv_bfloat16*)down_s.data_ptr(), (const __nv_bfloat16*)down_z.data_ptr(), MOE_INTER, HIDDEN }; blk.moe.s_gate = { s_gate_wq.data_ptr(), (const __nv_bfloat16*)s_gate_s.data_ptr(), (const __nv_bfloat16*)s_gate_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.s_up = { s_up_wq.data_ptr(), (const __nv_bfloat16*)s_up_s.data_ptr(), (const __nv_bfloat16*)s_up_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.s_down = { s_down_wq.data_ptr(), (const __nv_bfloat16*)s_down_s.data_ptr(), (const __nv_bfloat16*)s_down_z.data_ptr(), MOE_INTER, HIDDEN }; } void setup_mla_block( torch::Tensor params_tensor, int b, torch::Tensor attn_norm, torch::Tensor moe_norm, torch::Tensor q_wq, torch::Tensor q_s, torch::Tensor q_z, torch::Tensor kva_wq, torch::Tensor kva_s, torch::Tensor kva_z, torch::Tensor kvb_wq, torch::Tensor kvb_s, torch::Tensor kvb_z, torch::Tensor o_wq, torch::Tensor o_s, torch::Tensor o_z, torch::Tensor router_w, torch::Tensor gate_wq, torch::Tensor gate_s, torch::Tensor gate_z, torch::Tensor up_wq, torch::Tensor up_s, torch::Tensor up_z, torch::Tensor down_wq, torch::Tensor down_s, torch::Tensor down_z, torch::Tensor s_gate_wq, torch::Tensor s_gate_s, torch::Tensor s_gate_z, torch::Tensor s_up_wq, torch::Tensor s_up_s, torch::Tensor s_up_z, torch::Tensor s_down_wq, torch::Tensor s_down_s, torch::Tensor s_down_z ) { MegakernelParams* p = (MegakernelParams*)params_tensor.data_ptr(); BlockWeights& blk = p->blocks[b]; blk.attn_norm = (const __nv_bfloat16*)attn_norm.data_ptr(); blk.moe_norm = (const __nv_bfloat16*)moe_norm.data_ptr(); blk.mla.q_proj = { q_wq.data_ptr(), (const __nv_bfloat16*)q_s.data_ptr(), (const __nv_bfloat16*)q_z.data_ptr(), HIDDEN, 6144 }; blk.mla.kv_a = { kva_wq.data_ptr(), (const __nv_bfloat16*)kva_s.data_ptr(), (const __nv_bfloat16*)kva_z.data_ptr(), HIDDEN, 576 }; blk.mla.kv_b = { kvb_wq.data_ptr(), (const __nv_bfloat16*)kvb_s.data_ptr(), (const __nv_bfloat16*)kvb_z.data_ptr(), KV_LORA, 8192 }; blk.mla.o_proj = { o_wq.data_ptr(), (const __nv_bfloat16*)o_s.data_ptr(), (const __nv_bfloat16*)o_z.data_ptr(), KDA_CHANNELS, HIDDEN }; blk.moe.router = (const __nv_bfloat16*)router_w.data_ptr(); blk.moe.gate = { gate_wq.data_ptr(), (const __nv_bfloat16*)gate_s.data_ptr(), (const __nv_bfloat16*)gate_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.up = { up_wq.data_ptr(), (const __nv_bfloat16*)up_s.data_ptr(), (const __nv_bfloat16*)up_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.down = { down_wq.data_ptr(), (const __nv_bfloat16*)down_s.data_ptr(), (const __nv_bfloat16*)down_z.data_ptr(), MOE_INTER, HIDDEN }; blk.moe.s_gate = { s_gate_wq.data_ptr(), (const __nv_bfloat16*)s_gate_s.data_ptr(), (const __nv_bfloat16*)s_gate_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.s_up = { s_up_wq.data_ptr(), (const __nv_bfloat16*)s_up_s.data_ptr(), (const __nv_bfloat16*)s_up_z.data_ptr(), HIDDEN, MOE_INTER }; blk.moe.s_down = { s_down_wq.data_ptr(), (const __nv_bfloat16*)s_down_s.data_ptr(), (const __nv_bfloat16*)s_down_z.data_ptr(), MOE_INTER, HIDDEN }; } int get_params_size() { return sizeof(MegakernelParams); } int get_scratch_size() { return sizeof(Scratchpad); } """ CPP_SRC = """ #include void launch_kimi_decode_megakernel_opt( torch::Tensor hidden_in, torch::Tensor hidden_out, torch::Tensor s0, torch::Tensor cq0, torch::Tensor ck0, torch::Tensor cv0, torch::Tensor s1, torch::Tensor cq1, torch::Tensor ck1, torch::Tensor cv1, torch::Tensor s2, torch::Tensor cq2, torch::Tensor ck2, torch::Tensor cv2, torch::Tensor old_c_kv, torch::Tensor old_k_rope, torch::Tensor new_c_kv, torch::Tensor new_k_rope, int pos, torch::Tensor params_tensor, torch::Tensor scratch_tensor ); void setup_kda_block( torch::Tensor params_tensor, int b, torch::Tensor attn_norm, torch::Tensor moe_norm, torch::Tensor q_wq, torch::Tensor q_s, torch::Tensor q_z, torch::Tensor k_wq, torch::Tensor k_s, torch::Tensor k_z, torch::Tensor v_wq, torch::Tensor v_s, torch::Tensor v_z, torch::Tensor g_wq, torch::Tensor g_s, torch::Tensor g_z, torch::Tensor beta_w, torch::Tensor conv_w, torch::Tensor o_wq, torch::Tensor o_s, torch::Tensor o_z, torch::Tensor router_w, torch::Tensor gate_wq, torch::Tensor gate_s, torch::Tensor gate_z, torch::Tensor up_wq, torch::Tensor up_s, torch::Tensor up_z, torch::Tensor down_wq, torch::Tensor down_s, torch::Tensor down_z, torch::Tensor s_gate_wq, torch::Tensor s_gate_s, torch::Tensor s_gate_z, torch::Tensor s_up_wq, torch::Tensor s_up_s, torch::Tensor s_up_z, torch::Tensor s_down_wq, torch::Tensor s_down_s, torch::Tensor s_down_z ); void setup_mla_block( torch::Tensor params_tensor, int b, torch::Tensor attn_norm, torch::Tensor moe_norm, torch::Tensor q_wq, torch::Tensor q_s, torch::Tensor q_z, torch::Tensor kva_wq, torch::Tensor kva_s, torch::Tensor kva_z, torch::Tensor kvb_wq, torch::Tensor kvb_s, torch::Tensor kvb_z, torch::Tensor o_wq, torch::Tensor o_s, torch::Tensor o_z, torch::Tensor router_w, torch::Tensor gate_wq, torch::Tensor gate_s, torch::Tensor gate_z, torch::Tensor up_wq, torch::Tensor up_s, torch::Tensor up_z, torch::Tensor down_wq, torch::Tensor down_s, torch::Tensor down_z, torch::Tensor s_gate_wq, torch::Tensor s_gate_s, torch::Tensor s_gate_z, torch::Tensor s_up_wq, torch::Tensor s_up_s, torch::Tensor s_up_z, torch::Tensor s_down_wq, torch::Tensor s_down_s, torch::Tensor s_down_z ); int get_params_size(); int get_scratch_size(); """ _MODULE = None def _get_cuda_module(): global _MODULE if _MODULE is None: _MODULE = load_inline( name="kimi_megakernel_lib", cpp_sources=CPP_SRC, cuda_sources=CUDA_SRC, functions=[ "launch_kimi_decode_megakernel_opt", "setup_kda_block", "setup_mla_block", "get_params_size", "get_scratch_size" ], extra_cuda_cflags=["-O3", "--use_fast_math", "-gencode=arch=compute_120,code=sm_120"] ) return _MODULE @dataclass(frozen=True) class Config: hidden: int = 2304 kda_heads: int = 32 kda_head_dim: int = 128 short_conv: int = 4 mla_heads: int = 32 kv_lora: int = 512 qk_nope: int = 128 qk_rope: int = 64 v_head: int = 128 rope_theta: float = 10000.0 n_experts: int = 64 n_active: int = 8 n_shared: int = 1 moe_inter: int = 1024 routed_scaling: float = 2.446 group: int = 128 pattern: tuple = ("K", "K", "K", "M") dtype: torch.dtype = field(default=torch.bfloat16) def build_config(shape: dict) -> Config: return Config(n_experts=int(shape.get("n_experts", 64))) def _pack_int4(w_q: torch.Tensor) -> torch.Tensor: lo = w_q[0::2] & 0xF hi = w_q[1::2] & 0xF return (lo | (hi << 4)).contiguous() def _unpack_int4(w_packed: torch.Tensor, K: int) -> torch.Tensor: out = torch.empty((K, w_packed.shape[1]), dtype=torch.uint8, device=w_packed.device) out[0::2] = w_packed & 0xF out[1::2] = (w_packed >> 4) & 0xF return out def quantize(w_io: torch.Tensor, group: int = GROUP_SIZE): K, N = w_io.shape ng = K // group wg = w_io.view(ng, group, N).float() wmin = wg.min(dim=1, keepdim=True).values wmax = wg.max(dim=1, keepdim=True).values scales = (wmax - wmin).clamp_min(1e-8) / 15.0 zeros = (-wmin / scales).round().clamp(0, 15) w_q = ((wg / scales) + zeros).round().clamp(0, 15).to(torch.uint8).view(K, N) return _pack_int4(w_q), scales.squeeze(1).to(torch.bfloat16), zeros.squeeze(1).to(torch.bfloat16) def dequant(w_q: torch.Tensor, scales: torch.Tensor, zeros: torch.Tensor, K: int, group: int) -> torch.Tensor: wu = _unpack_int4(w_q, K).to(torch.bfloat16) s = scales.repeat_interleave(group, dim=0) z = zeros.repeat_interleave(group, dim=0) return (wu - z) * s class QuantLinear(nn.Module): def __init__(self, in_f: int, out_f: int, group: int = GROUP_SIZE): super().__init__() assert in_f % group == 0 and in_f % 2 == 0 self.in_f, self.out_f, self.group = in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen: torch.Generator, std: float = 0.02) -> None: w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q.copy_(wq) self.scales.copy_(s) self.zeros.copy_(z) def weight_bf(self) -> torch.Tensor: return dequant(self.w_q, self.scales, self.zeros, self.in_f, self.group) def forward(self, x: torch.Tensor) -> torch.Tensor: return (x.float() @ self.weight_bf().float()).to(torch.bfloat16) class QuantExperts(nn.Module): def __init__(self, n: int, in_f: int, out_f: int, group: int = GROUP_SIZE): super().__init__() self.n, self.in_f, self.out_f, self.group = n, in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(n, in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen: torch.Generator, std: float = 0.02) -> None: for e in range(self.n): w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q[e].copy_(wq) self.scales[e].copy_(s) self.zeros[e].copy_(z) def weight_bf(self, e: int) -> torch.Tensor: return dequant(self.w_q[e], self.scales[e], self.zeros[e], self.in_f, self.group) class KDA(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg H, Dk, d = cfg.kda_heads, cfg.kda_head_dim, cfg.hidden self.q_proj = QuantLinear(d, H * Dk, cfg.group) self.k_proj = QuantLinear(d, H * Dk, cfg.group) self.v_proj = QuantLinear(d, H * Dk, cfg.group) self.g_proj = QuantLinear(d, H * Dk, cfg.group) self.beta_proj = nn.Linear(d, H, bias=False, dtype=cfg.dtype) self.conv_w = nn.Parameter(torch.empty(3, H * Dk, cfg.short_conv, dtype=cfg.dtype)) self.o_proj = QuantLinear(H * Dk, d, cfg.group) self.scale = Dk ** -0.5 class MLA(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg H, d = cfg.mla_heads, cfg.hidden self.q_proj = QuantLinear(d, H * (cfg.qk_nope + cfg.qk_rope), cfg.group) self.kv_a = QuantLinear(d, cfg.kv_lora + cfg.qk_rope, cfg.group) self.kv_b = QuantLinear(cfg.kv_lora, H * (cfg.qk_nope + cfg.v_head), cfg.group) self.o_proj = QuantLinear(H * cfg.v_head, d, cfg.group) self.scale = (cfg.qk_nope + cfg.qk_rope) ** -0.5 class MoE(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg d, m, E = cfg.hidden, cfg.moe_inter, cfg.n_experts self.router = nn.Linear(d, E, bias=False, dtype=cfg.dtype) self.gate = QuantExperts(E, d, m, cfg.group) self.up = QuantExperts(E, d, m, cfg.group) self.down = QuantExperts(E, m, d, cfg.group) self.s_gate = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_up = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_down = QuantExperts(cfg.n_shared, m, d, cfg.group) class Block(nn.Module): def __init__(self, cfg: Config, kind: str): super().__init__() self.kind = kind self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.attn = KDA(cfg) if kind == "K" else MLA(cfg) self.moe = MoE(cfg) class Model(nn.Module): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern) self.mod = _get_cuda_module() self.params_size = self.mod.get_params_size() self.scratch_size = self.mod.get_scratch_size() self.params_tensor_cpu = torch.zeros(self.params_size, dtype=torch.uint8, device="cpu") self.params_tensor_cuda = None self.scratch_tensor = None self._synced_weights = False self.reset_parameters() def reset_parameters(self): g = torch.Generator(device="cpu").manual_seed(1234) for mod in self.modules(): if isinstance(mod, (QuantLinear, QuantExperts)): mod.init_random(g) elif isinstance(mod, nn.Linear): nn.init.normal_(mod.weight, 0.0, 0.02, generator=g) elif isinstance(mod, KDA): nn.init.normal_(mod.conv_w, 0.0, 0.1, generator=g) def _sync_weights(self, device): for b in range(4): blk = self.blocks[b] if b < 3: self.mod.setup_kda_block( self.params_tensor_cpu, b, blk.attn_norm, blk.moe_norm, blk.attn.q_proj.w_q, blk.attn.q_proj.scales, blk.attn.q_proj.zeros, blk.attn.k_proj.w_q, blk.attn.k_proj.scales, blk.attn.k_proj.zeros, blk.attn.v_proj.w_q, blk.attn.v_proj.scales, blk.attn.v_proj.zeros, blk.attn.g_proj.w_q, blk.attn.g_proj.scales, blk.attn.g_proj.zeros, blk.attn.beta_proj.weight, blk.attn.conv_w, blk.attn.o_proj.w_q, blk.attn.o_proj.scales, blk.attn.o_proj.zeros, blk.moe.router.weight, blk.moe.gate.w_q, blk.moe.gate.scales, blk.moe.gate.zeros, blk.moe.up.w_q, blk.moe.up.scales, blk.moe.up.zeros, blk.moe.down.w_q, blk.moe.down.scales, blk.moe.down.zeros, blk.moe.s_gate.w_q, blk.moe.s_gate.scales, blk.moe.s_gate.zeros, blk.moe.s_up.w_q, blk.moe.s_up.scales, blk.moe.s_up.zeros, blk.moe.s_down.w_q, blk.moe.s_down.scales, blk.moe.s_down.zeros, ) else: self.mod.setup_mla_block( self.params_tensor_cpu, b, blk.attn_norm, blk.moe_norm, blk.attn.q_proj.w_q, blk.attn.q_proj.scales, blk.attn.q_proj.zeros, blk.attn.kv_a.w_q, blk.attn.kv_a.scales, blk.attn.kv_a.zeros, blk.attn.kv_b.w_q, blk.attn.kv_b.scales, blk.attn.kv_b.zeros, blk.attn.o_proj.w_q, blk.attn.o_proj.scales, blk.attn.o_proj.zeros, blk.moe.router.weight, blk.moe.gate.w_q, blk.moe.gate.scales, blk.moe.gate.zeros, blk.moe.up.w_q, blk.moe.up.scales, blk.moe.up.zeros, blk.moe.down.w_q, blk.moe.down.scales, blk.moe.down.zeros, blk.moe.s_gate.w_q, blk.moe.s_gate.scales, blk.moe.s_gate.zeros, blk.moe.s_up.w_q, blk.moe.s_up.scales, blk.moe.s_up.zeros, blk.moe.s_down.w_q, blk.moe.s_down.scales, blk.moe.s_down.zeros, ) if self.params_tensor_cuda is None or self.params_tensor_cuda.device != device: self.params_tensor_cuda = self.params_tensor_cpu.to(device) else: self.params_tensor_cuda.copy_(self.params_tensor_cpu) if self.scratch_tensor is None or self.scratch_tensor.device != device: self.scratch_tensor = torch.zeros(self.scratch_size, dtype=torch.uint8, device=device) self._synced_weights = True def load_state_dict(self, state_dict, strict=True): res = super().load_state_dict(state_dict, strict=strict) self._synced_weights = False return res def step(self, hidden: torch.Tensor, state: list) -> tuple[torch.Tensor, list]: if not self._synced_weights: self._sync_weights(hidden.device) pos = state[3]["c_kv"].shape[0] new_c_kv = torch.empty((pos + 1, 512), dtype=torch.bfloat16, device=hidden.device) new_k_rope = torch.empty((pos + 1, 64), dtype=torch.bfloat16, device=hidden.device) out_hidden = torch.empty_like(hidden) # EXACTLY ONE GPU KERNEL LAUNCH self.mod.launch_kimi_decode_megakernel_opt( hidden, out_hidden, state[0]["S"], state[0]["cq"], state[0]["ck"], state[0]["cv"], state[1]["S"], state[1]["cq"], state[1]["ck"], state[1]["cv"], state[2]["S"], state[2]["cq"], state[2]["ck"], state[2]["cv"], state[3]["c_kv"], state[3]["k_rope"], new_c_kv, new_k_rope, pos, self.params_tensor_cuda, self.scratch_tensor ) state[3]["c_kv"] = new_c_kv state[3]["k_rope"] = new_k_rope return out_hidden, state def init_state(cfg: Config, context_len: int, seed: int) -> list: dev = torch.device("cuda:0") g = torch.Generator(device=dev).manual_seed(seed) H, Dk = cfg.kda_heads, cfg.kda_head_dim C = H * Dk state = [] for kind in cfg.pattern: if kind == "K": state.append({ "S": torch.randn(H, Dk, Dk, device=dev, generator=g) * 0.05, "cq": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "ck": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "cv": torch.randn(cfg.short_conv - 1, C, device=dev, generator=g, dtype=cfg.dtype) * 0.1, }) else: state.append({ "c_kv": torch.randn(context_len, cfg.kv_lora, device=dev, generator=g, dtype=cfg.dtype) * 0.1, "k_rope": torch.randn(context_len, cfg.qk_rope, device=dev, generator=g, dtype=cfg.dtype) * 0.1, }) return state def init_token(cfg: Config, seed: int) -> torch.Tensor: dev = torch.device("cuda:0") g = torch.Generator(device=dev).manual_seed(seed + 1) return torch.randn(cfg.hidden, device=dev, generator=g, dtype=cfg.dtype) * 0.25