"""Fused W4A16 Kimi-Linear decode megakernel (batch-1). One cooperative CUDA launch per step. Int4 unpack + group dequant stay inside the GEMV; MLA uses the absorbed latent form so the kv_b projection is not materialized over the cache. """ from __future__ import annotations import os import torch import torch.nn as nn import torch.nn.functional as F GROUP = 128 HIDDEN = 2304 KDA_H = 32 KDA_D = 128 KDA_C = KDA_H * KDA_D # 4096 MLA_Q = 32 * (128 + 64) # 6144 KV_A = 512 + 64 KV_B_N = 32 * (128 + 128) # 8192 MOE_INTER = 1024 N_EXPERTS = 64 N_ACTIVE = 8 CACHE_CAP = 17000 THREADS = 256 NPT = 4 N_TILE = THREADS * NPT # 1024 K_TILE = 128 MAX_BLOCKS = 1024 # Pointer pack layout. Dynamic slots are rewritten every step; static slots once. # 0 hidden_in, 1 hidden_out, 2 ckv_src, 3 krope_src, 4 ckv_dst, 5 krope_dst # 6..8 S, 9..11 cq, 12..14 ck, 15..17 cv # 18 ws_norm, 19 ws_q, 20 ws_k, 21 ws_v, 22 ws_g, 23 ws_beta, 24 ws_o, 25 ws_kv # 26 ws_qabs, 27 ws_qrope, 28 ws_ctx, 29 ws_scores, 30 ws_gate, 31 ws_up # 32 ws_hid, 33 ws_moe, 34 exp_idx, 35 exp_w, 36 ws_partial # KDA layer li base = 37 + li*19 # MLA base = 94 # MoE layer li base = 108 + li*19 PACK_N = 192 KDA_BASE = 37 KDA_STRIDE = 19 MLA_BASE = 94 MOE_BASE = 108 MOE_STRIDE = 19 def _ext(): global _EXT if _EXT is not None: return _EXT import ctypes import hashlib import subprocess cuda_src = r""" #include #include #include #include #include namespace cg = cooperative_groups; static constexpr int HIDDEN = 2304; static constexpr int KDA_C = 4096; static constexpr int KDA_H = 32; static constexpr int KDA_D = 128; static constexpr int THREADS = 256; static constexpr int NPT = 4; static constexpr int N_TILE = 1024; static constexpr int K_TILE = 128; static constexpr int MAX_BLOCKS = 1024; static constexpr int KDA_BASE = 37; static constexpr int KDA_STRIDE = 19; static constexpr int MLA_BASE = 94; static constexpr int MOE_BASE = 108; static constexpr int MOE_STRIDE = 19; static constexpr int MOE_INTER = 1024; static constexpr float ATTN_SCALE = 0.07216878364870322f; static constexpr float KDA_SCALE = 0.08838834764831845f; static constexpr float ROUTED_SCALE = 2.446f; static constexpr float ROPE_THETA = 10000.f; struct Pack { int64_t p[192]; int pos; int pad; }; __device__ __forceinline__ float bf2f(uint16_t u) { uint32_t x = (uint32_t)u << 16; float f; asm volatile("mov.b32 %0, %1;" : "=f"(f) : "r"(x)); return f; } __device__ __forceinline__ uint16_t f2bf(float f) { uint32_t x; asm volatile("mov.b32 %0, %1;" : "=r"(x) : "f"(f)); uint32_t lsb = (x >> 16) & 1u; x += 0x7fffu + lsb; return (uint16_t)(x >> 16); } __device__ __forceinline__ float round_bf(float f) { return bf2f(f2bf(f)); } __device__ __forceinline__ float ld_bf(const uint16_t* p) { return bf2f(__ldg(p)); } __device__ __forceinline__ uint4 ld4(const void* p) { uint4 v; asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];" : "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w) : "l"(p)); return v; } __device__ __forceinline__ void st_bf(uint16_t* p, float f) { *p = f2bf(f); } __device__ __forceinline__ float warp_sum(float v) { #pragma unroll for (int m = 16; m > 0; m >>= 1) v += __shfl_xor_sync(0xffffffff, v, m); return v; } __device__ __forceinline__ float warp_max(float v) { #pragma unroll for (int m = 16; m > 0; m >>= 1) v = fmaxf(v, __shfl_xor_sync(0xffffffff, v, m)); return v; } __device__ float block_sum(float v, float* sm) { int lane = threadIdx.x & 31; int wid = threadIdx.x >> 5; v = warp_sum(v); if (lane == 0) sm[wid] = v; __syncthreads(); int nwarps = blockDim.x >> 5; float t = (threadIdx.x < nwarps) ? sm[threadIdx.x] : 0.f; if (wid == 0) t = warp_sum(t); if (threadIdx.x == 0) sm[0] = t; __syncthreads(); t = sm[0]; __syncthreads(); return t; } __device__ float block_max(float v, float* sm) { int lane = threadIdx.x & 31; int wid = threadIdx.x >> 5; v = warp_max(v); if (lane == 0) sm[wid] = v; __syncthreads(); int nwarps = blockDim.x >> 5; float t = (threadIdx.x < nwarps) ? sm[threadIdx.x] : -1e30f; if (wid == 0) t = warp_max(t); if (threadIdx.x == 0) sm[0] = t; __syncthreads(); t = sm[0]; __syncthreads(); return t; } __device__ __forceinline__ float silu(float x) { float s = (x >= 0.f) ? (1.f / (1.f + expf(-x))) : (expf(x) / (1.f + expf(x))); return x * s; } __device__ __forceinline__ float softplus_neg_exp(float x) { // exp(-softplus(x)), PyTorch softplus threshold 20 if (x > 20.f) return expf(-x); return 1.f / (1.f + expf(x)); } // One K-group (128) of an int4 GEMV. // 4 groups of 64 threads each own the same 1024 outputs (16 / thread, 16-byte loads) // and split the 64 packed K-pairs. Partials reduce in smem, one atomic per output. __device__ void gemv4( const uint8_t* __restrict__ wq, const uint16_t* __restrict__ scales, const uint16_t* __restrict__ zeros, const float* __restrict__ x, float* __restrict__ y, int N, int n0, int n_end, int k0, int y_delta, float y_scale, int round_x, char* smem ) { float* xs = reinterpret_cast(smem); int tid = threadIdx.x; if (tid < K_TILE) { float xv = x[k0 + tid]; xs[tid] = round_x ? round_bf(xv) : xv; } __syncthreads(); const int group = tid >> 6; // 0..3 const int sub = tid & 63; // 0..63 const int n_base = n0 + sub * 16; const bool active = (n_base + 15 < n_end); float acc[16]; #pragma unroll for (int j = 0; j < 16; ++j) acc[j] = 0.f; if (active) { int g = k0 >> 7; float s[16], z[16]; #pragma unroll for (int j = 0; j < 16; ++j) { s[j] = ld_bf(scales + (int64_t)g * N + n_base + j); z[j] = ld_bf(zeros + (int64_t)g * N + n_base + j); } const uint8_t* row = wq + ((int64_t)(k0 >> 1) * N) + n_base; int i = group; uint4 v = ld4(row + (int64_t)i * N); #pragma unroll 1 for (; i + 4 < 64; i += 4) { uint4 nxt = ld4(row + (int64_t)(i + 4) * N); float x0 = xs[i * 2]; float x1 = xs[i * 2 + 1]; uint32_t wds[4] = {v.x, v.y, v.z, v.w}; #pragma unroll for (int w = 0; w < 4; ++w) { uint32_t p = wds[w]; #pragma unroll for (int b = 0; b < 4; ++b) { int j = w * 4 + b; uint32_t byte = (p >> (8 * b)) & 0xFFu; float q0 = float(byte & 0xFu); float q1 = float(byte >> 4); acc[j] = fmaf((q0 - z[j]) * s[j], x0, acc[j]); acc[j] = fmaf((q1 - z[j]) * s[j], x1, acc[j]); } } v = nxt; } { float x0 = xs[i * 2]; float x1 = xs[i * 2 + 1]; uint32_t wds[4] = {v.x, v.y, v.z, v.w}; #pragma unroll for (int w = 0; w < 4; ++w) { uint32_t p = wds[w]; #pragma unroll for (int b = 0; b < 4; ++b) { int j = w * 4 + b; uint32_t byte = (p >> (8 * b)) & 0xFFu; float q0 = float(byte & 0xFu); float q1 = float(byte >> 4); acc[j] = fmaf((q0 - z[j]) * s[j], x0, acc[j]); acc[j] = fmaf((q1 - z[j]) * s[j], x1, acc[j]); } } } } __syncthreads(); // Reduce the 4 K-groups. xs is dead; reuse smem as [64][16]. float* red = reinterpret_cast(smem); if (group == 0) { #pragma unroll for (int j = 0; j < 16; ++j) red[sub * 16 + j] = active ? acc[j] : 0.f; } __syncthreads(); if (group != 0 && active) { #pragma unroll for (int j = 0; j < 16; ++j) atomicAdd(red + sub * 16 + j, acc[j]); } __syncthreads(); if (group == 0 && active) { #pragma unroll for (int j = 0; j < 16; ++j) atomicAdd(y + (n_base + j - y_delta), red[sub * 16 + j] * y_scale); } __syncthreads(); } __device__ void zero_f(float* p, int n) { int stride = blockDim.x * gridDim.x; for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += stride) p[i] = 0.f; } __device__ void rmsnorm(const uint16_t* x, const uint16_t* w, float* y, float* sm) { int tid = threadIdx.x; float ss = 0.f; for (int i = tid; i < HIDDEN; i += blockDim.x) { float v = ld_bf(x + i); ss = fmaf(v, v, ss); } ss = block_sum(ss, sm); float inv = rsqrtf(ss / float(HIDDEN) + 1e-6f); for (int i = tid; i < HIDDEN; i += blockDim.x) { float v = ld_bf(x + i) * inv * ld_bf(w + i); y[i] = round_bf(v); } } __device__ void residual(uint16_t* h, const float* add, int n) { int stride = blockDim.x * gridDim.x; for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += stride) { float s = ld_bf(h + i) + round_bf(add[i]); st_bf(h + i, s); } } __device__ void bf16_gemv_out(const float* x, const uint16_t* W, float* y, int M, int K, float* sm, int do_sigmoid) { int tid = threadIdx.x; for (int m = 0; m < M; ++m) { float acc = 0.f; const uint16_t* row = W + (int64_t)m * K; for (int i = tid; i < K; i += blockDim.x) acc = fmaf(x[i], ld_bf(row + i), acc); acc = block_sum(acc, sm); if (tid == 0) { acc = round_bf(acc); if (do_sigmoid) acc = 1.f / (1.f + expf(-acc)); y[m] = acc; } __syncthreads(); } } __device__ void kda_head(int h, float* ws_q, float* ws_k, float* ws_v, float* ws_g, float* ws_beta, float* ws_o, float* S, uint16_t* cq, uint16_t* ck, uint16_t* cv, const uint16_t* conv, char* smem) { int tx = threadIdx.x; float* qs = reinterpret_cast(smem); float* ks = qs + 128; float* vs = ks + 128; float* eg = vs + 128; float* pred = eg + 128; int base = h * KDA_D; if (tx < KDA_D) { int c = base + tx; // short conv on q, k, v. conv layout (3, 4096, 4) auto conv_one = [&](float* ws, uint16_t* prev, int idx) { float val = round_bf(ws[c]); uint16_t p0 = prev[0 * KDA_C + c]; uint16_t p1 = prev[1 * KDA_C + c]; uint16_t p2 = prev[2 * KDA_C + c]; const uint16_t* cw = conv + ((int64_t)idx * KDA_C + c) * 4; float acc = bf2f(p0) * ld_bf(cw + 0) + bf2f(p1) * ld_bf(cw + 1) + bf2f(p2) * ld_bf(cw + 2) + val * ld_bf(cw + 3); ws[c] = round_bf(silu(acc)); prev[0 * KDA_C + c] = p1; prev[1 * KDA_C + c] = p2; prev[2 * KDA_C + c] = f2bf(val); }; conv_one(ws_q, cq, 0); conv_one(ws_k, ck, 1); conv_one(ws_v, cv, 2); qs[tx] = ws_q[c] * KDA_SCALE; ks[tx] = ws_k[c]; vs[tx] = ws_v[c]; eg[tx] = softplus_neg_exp(round_bf(ws_g[c])); } __syncthreads(); if (tx < KDA_D) { float* Sh = S + (int64_t)h * KDA_D * KDA_D; int j = tx; float p = 0.f; for (int i = 0; i < KDA_D; ++i) p = fmaf(Sh[(int64_t)i * KDA_D + j] * eg[i], ks[i], p); pred[j] = p; } __syncthreads(); if (tx < KDA_D) { float* Sh = S + (int64_t)h * KDA_D * KDA_D; int j = tx; float beta = ws_beta[h]; float delta = vs[j] - pred[j]; float o = 0.f; for (int i = 0; i < KDA_D; ++i) { float s = Sh[(int64_t)i * KDA_D + j] * eg[i] + beta * ks[i] * delta; Sh[(int64_t)i * KDA_D + j] = s; o = fmaf(s, qs[i], o); } ws_o[base + j] = o; } __syncthreads(); } __device__ void apply_rope_store(float e, float o, int pair, int pos, uint16_t* dst) { float inv = expf(-logf(ROPE_THETA) * (2.f * pair) / 64.f); float ang = float(pos) * inv; float c = cosf(ang), s = sinf(ang); st_bf(dst + 2 * pair, e * c - o * s); st_bf(dst + 2 * pair + 1, o * c + e * s); } __device__ void absorb_q(int h, const uint8_t* wq, const uint16_t* sc, const uint16_t* zc, const float* q, float* qabs, char* smem) { float* qn = reinterpret_cast(smem); int tx = threadIdx.x; int n0 = h * 256; constexpr int N = 8192; if (tx < 128) qn[tx] = round_bf(q[h * 192 + tx]); __syncthreads(); for (int k = tx; k < 512; k += blockDim.x) { int kp = k >> 1; int hi = k & 1; int g = k >> 7; float acc = 0.f; const uint8_t* row = wq + (int64_t)kp * N + n0; const uint16_t* sg = sc + (int64_t)g * N + n0; const uint16_t* zg = zc + (int64_t)g * N + n0; for (int d = 0; d < 128; ++d) { uint8_t b = row[d]; float nib = hi ? float(b >> 4) : float(b & 0xF); float w = (nib - ld_bf(zg + d)) * ld_bf(sg + d); acc = fmaf(w, qn[d], acc); } qabs[(int64_t)k * 32 + h] = acc; } __syncthreads(); } __device__ void mla_scores(const Pack& a, int L, int pos, char* smem) { const uint16_t* ckv_src = (const uint16_t*)a.p[2]; const uint16_t* kr_src = (const uint16_t*)a.p[3]; uint16_t* ckv_dst = (uint16_t*)a.p[4]; uint16_t* kr_dst = (uint16_t*)a.p[5]; const float* qabs = (const float*)a.p[26]; const float* qrope = (const float*)a.p[27]; float* scores = (float*)a.p[29]; // qabs tile [64][32] then per-warp ckv chunk [8][64] float* qsm = reinterpret_cast(smem); uint16_t* cchunk = reinterpret_cast(smem + 8192); int tx = threadIdx.x; int warp = tx >> 5; int lane = tx & 31; const bool same_c = (ckv_src == ckv_dst); const bool same_k = (kr_src == kr_dst); for (int lbase = blockIdx.x * 8; lbase < L; lbase += gridDim.x * 8) { int l = lbase + warp; bool valid = l < L; float acc = 0.f; for (int k0 = 0; k0 < 512; k0 += 64) { for (int i = tx; i < 64 * 32; i += blockDim.x) { int kk = i >> 5; int hh = i & 31; qsm[i] = qabs[(int64_t)(k0 + kk) * 32 + hh]; } if (valid) { bool copy = (l < pos) && !same_c; const uint16_t* cptr = (copy ? ckv_src : ckv_dst) + (int64_t)l * 512 + k0; uint16_t v0 = cptr[lane]; uint16_t v1 = cptr[32 + lane]; cchunk[warp * 64 + lane] = v0; cchunk[warp * 64 + 32 + lane] = v1; if (copy) { uint16_t* dst = ckv_dst + (int64_t)l * 512 + k0; dst[lane] = v0; dst[32 + lane] = v1; } } __syncthreads(); if (valid) { #pragma unroll for (int k = 0; k < 64; k += 4) { acc = fmaf(bf2f(cchunk[warp * 64 + k]), qsm[k * 32 + lane], acc); acc = fmaf(bf2f(cchunk[warp * 64 + k + 1]), qsm[(k + 1) * 32 + lane], acc); acc = fmaf(bf2f(cchunk[warp * 64 + k + 2]), qsm[(k + 2) * 32 + lane], acc); acc = fmaf(bf2f(cchunk[warp * 64 + k + 3]), qsm[(k + 3) * 32 + lane], acc); } } __syncthreads(); } if (valid) { bool kcopy = (l < pos) && !same_k; const uint16_t* kp = (kcopy ? kr_src : kr_dst) + (int64_t)l * 64; float rd = 0.f; for (int d = 0; d < 64; ++d) rd = fmaf(bf2f(kp[d]), qrope[lane * 64 + d], rd); if (kcopy) { kr_dst[(int64_t)l * 64 + lane] = kp[lane]; kr_dst[(int64_t)l * 64 + 32 + lane] = kp[32 + lane]; } scores[(int64_t)lane * L + l] = (acc + rd) * ATTN_SCALE; } __syncthreads(); } } // Parallel softmax over L. scores is [head][L]. partials live in ws_partial. __device__ void softmax_partial_max(float* scores, int L, float* pmax, float* sm) { int npos = (L + gridDim.x - 1) / gridDim.x; int l0 = blockIdx.x * npos; if (l0 > L) l0 = L; int l1 = l0 + npos; if (l1 > L) l1 = L; for (int h = 0; h < 32; ++h) { float m = -1e30f; const float* col = scores + (int64_t)h * L; for (int l = l0 + threadIdx.x; l < l1; l += blockDim.x) m = fmaxf(m, col[l]); m = block_max(m, sm); if (threadIdx.x == 0) pmax[blockIdx.x * 32 + h] = m; } } __device__ void softmax_reduce_max(float* pmax, int nb, float* gmax) { if (blockIdx.x == 0) { for (int h = threadIdx.x; h < 32; h += blockDim.x) { float m = -1e30f; for (int b = 0; b < nb; ++b) m = fmaxf(m, pmax[(int64_t)b * 32 + h]); gmax[h] = m; } } } __device__ void softmax_partial_sum(float* scores, int L, const float* gmax, float* psum, float* sm) { int npos = (L + gridDim.x - 1) / gridDim.x; int l0 = blockIdx.x * npos; if (l0 > L) l0 = L; int l1 = l0 + npos; if (l1 > L) l1 = L; for (int h = 0; h < 32; ++h) { float m = gmax[h]; float s = 0.f; const float* col = scores + (int64_t)h * L; for (int l = l0 + threadIdx.x; l < l1; l += blockDim.x) s += expf(col[l] - m); s = block_sum(s, sm); if (threadIdx.x == 0) psum[blockIdx.x * 32 + h] = s; } } __device__ void softmax_reduce_sum(float* psum, int nb, float* gsum) { if (blockIdx.x == 0) { for (int h = threadIdx.x; h < 32; h += blockDim.x) { float s = 0.f; for (int b = 0; b < nb; ++b) s += psum[(int64_t)b * 32 + h]; gsum[h] = s; } } } __device__ void softmax_write(float* scores, int L, const float* gmax, const float* gsum) { int npos = (L + gridDim.x - 1) / gridDim.x; int l0 = blockIdx.x * npos; if (l0 > L) l0 = L; int l1 = l0 + npos; if (l1 > L) l1 = L; for (int h = 0; h < 32; ++h) { float m = gmax[h]; float inv = 1.f / gsum[h]; float* col = scores + (int64_t)h * L; for (int l = l0 + threadIdx.x; l < l1; l += blockDim.x) col[l] = expf(col[l] - m) * inv; } } __device__ void weighted_ctx(const Pack& a, int L) { const uint16_t* ckv = (const uint16_t*)a.p[4]; const float* prob = (const float*)a.p[29]; float* partial = (float*)a.p[36] + (int64_t)blockIdx.x * (32 * 512); int tid = threadIdx.x; for (int pass = 0; pass < 2; ++pass) { int k = tid + pass * 256; float acc[32]; #pragma unroll for (int h = 0; h < 32; ++h) acc[h] = 0.f; for (int l = blockIdx.x; l < L; l += gridDim.x) { float cv = ld_bf(ckv + (int64_t)l * 512 + k); #pragma unroll for (int h = 0; h < 32; ++h) acc[h] = fmaf(prob[(int64_t)h * L + l], cv, acc[h]); } #pragma unroll for (int h = 0; h < 32; ++h) partial[h * 512 + k] = acc[h]; } } __device__ void reduce_ctx(const Pack& a) { float* partial = (float*)a.p[36]; float* ctx = (float*)a.p[28]; int nout = 32 * 512; int stride = blockDim.x * gridDim.x; int nb = gridDim.x; for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < nout; i += stride) { float s = 0.f; for (int b = 0; b < nb; ++b) s += partial[(int64_t)b * nout + i]; ctx[i] = s; } } __device__ void router_topk(const float* x, const uint16_t* W, int* exp_idx, float* exp_w, char* smem) { float* red = reinterpret_cast(smem); float* logits = red + 64; int tid = threadIdx.x; for (int e = 0; e < 64; ++e) { float acc = 0.f; const uint16_t* row = W + (int64_t)e * HIDDEN; for (int i = tid; i < HIDDEN; i += blockDim.x) acc = fmaf(x[i], ld_bf(row + i), acc); acc = block_sum(acc, red); if (tid == 0) logits[e] = round_bf(acc); __syncthreads(); } if (tid == 0) { float m = -1e30f; for (int i = 0; i < 64; ++i) m = fmaxf(m, logits[i]); float s = 0.f; for (int i = 0; i < 64; ++i) { logits[i] = expf(logits[i] - m); s += logits[i]; } float invs = 1.f / s; for (int i = 0; i < 64; ++i) logits[i] *= invs; int ids[64]; for (int i = 0; i < 64; ++i) ids[i] = i; for (int k = 0; k < 8; ++k) { int best = k; for (int j = k + 1; j < 64; ++j) { if (logits[j] > logits[best] || (logits[j] == logits[best] && ids[j] < ids[best])) best = j; } float tv = logits[k]; logits[k] = logits[best]; logits[best] = tv; int ti = ids[k]; ids[k] = ids[best]; ids[best] = ti; } float sum = 0.f; for (int k = 0; k < 8; ++k) sum += logits[k]; float inv = ROUTED_SCALE / (sum + 1e-9f); for (int k = 0; k < 8; ++k) { exp_idx[k] = ids[k]; exp_w[k] = logits[k] * inv; } exp_w[8] = 1.f; } } __device__ void gemv_tasks( int ntasks, int n_tiles, int k_tiles, int n_limit, const uint8_t* wq, const uint16_t* sc, const uint16_t* zc, const float* x, float* y, int N, int K, int y_delta, float y_scale, int round_x, char* smem ) { for (int task = blockIdx.x; task < ntasks; task += gridDim.x) { int nt = task / k_tiles; int kt = task - nt * k_tiles; int n0 = nt * N_TILE; int n_end = n0 + N_TILE; if (n_end > n_limit) n_end = n_limit; int k0 = kt * K_TILE; if (n0 < n_limit && k0 < K) gemv4(wq, sc, zc, x, y, N, n0, n_end, k0, y_delta, y_scale, round_x, smem); else __syncthreads(); } } // Multi-matrix / multi-expert task loop. task space is caller-defined via a device lambda-like switch. // Implemented explicitly at call sites for clarity. __device__ unsigned long long d_clk[256]; __device__ int d_tag[256]; __device__ int d_clk_n; __device__ int g_did_repack; __device__ __forceinline__ void stamp(int tag) { if (blockIdx.x == 0 && threadIdx.x == 0) { int i = d_clk_n; if (i < 256) { d_clk[i] = clock64(); d_tag[i] = tag; } d_clk_n = i + 1; } } #define GSYNC() do { stamp(1); grid.sync(); stamp(2); } while (0) // Repack int4 weights so each lane's K stream is contiguous uint32s: // dst[tile][lane][kp] = src[kp][tile*128 + lane*4]. Tile is 128 outputs. __device__ void repack_matrix(const uint8_t* src, uint8_t* dst, int K, int N, int E) { int kpairs = K >> 1; int tiles = N >> 7; int64_t mat_bytes = (int64_t)kpairs * N; int nwarps = gridDim.x * (blockDim.x >> 5); int wgid = blockIdx.x * (blockDim.x >> 5) + (threadIdx.x >> 5); int lane = threadIdx.x & 31; int jobs = tiles * kpairs; for (int e = 0; e < E; ++e) { const uint8_t* es = src + e * mat_bytes; uint8_t* ed = dst + e * mat_bytes; for (int job = wgid; job < jobs; job += nwarps) { int tile = job / kpairs; int kp = job - tile * kpairs; const uint32_t* s = reinterpret_cast(es + ((int64_t)kp * N + tile * 128)); uint32_t* d2 = reinterpret_cast(ed) + ((int64_t)tile * 32 + lane) * kpairs + kp; *d2 = s[lane]; } } } __device__ int64_t repack_off_kda(int layer, int mat) { // q,k,v,g,o each 1152*4096 = 4718592, o is the same size return (int64_t)(layer * 5 + mat) * 4718592LL; } __device__ int64_t repack_off_mla_q() { return 3LL * 5 * 4718592LL; } __device__ int64_t repack_off_mla_kvb() { return repack_off_mla_q() + 1152LL * 6144; } __device__ int64_t repack_off_mla_o() { return repack_off_mla_kvb() + 256LL * 8192; } __device__ int64_t repack_off_moe(int layer, int mat) { // after MLA q + kvb + o. gate/up/down/sg/su/sd int64_t base = repack_off_mla_o() + 2048LL * 2304; // per layer: 3*75497472 + 3*1179648 = 230031360 int64_t layer_base = base + (int64_t)layer * 230031360LL; if (mat == 0) return layer_base; // gate 64*1152*1024 if (mat == 1) return layer_base + 75497472LL; // up if (mat == 2) return layer_base + 2 * 75497472LL; // down if (mat == 3) return layer_base + 3 * 75497472LL; // sg if (mat == 4) return layer_base + 3 * 75497472LL + 1179648LL; // su return layer_base + 3 * 75497472LL + 2 * 1179648LL; // sd } __device__ void repack_all(const Pack& a, uint8_t* dst) { for (int layer = 0; layer < 3; ++layer) { int kb = KDA_BASE + layer * KDA_STRIDE; repack_matrix((const uint8_t*)a.p[kb + 2], dst + repack_off_kda(layer, 0), 2304, 4096, 1); repack_matrix((const uint8_t*)a.p[kb + 5], dst + repack_off_kda(layer, 1), 2304, 4096, 1); repack_matrix((const uint8_t*)a.p[kb + 8], dst + repack_off_kda(layer, 2), 2304, 4096, 1); repack_matrix((const uint8_t*)a.p[kb + 11], dst + repack_off_kda(layer, 3), 2304, 4096, 1); repack_matrix((const uint8_t*)a.p[kb + 14], dst + repack_off_kda(layer, 4), 4096, 2304, 1); } repack_matrix((const uint8_t*)a.p[MLA_BASE + 2], dst + repack_off_mla_q(), 2304, 6144, 1); repack_matrix((const uint8_t*)a.p[MLA_BASE + 8], dst + repack_off_mla_kvb(), 512, 8192, 1); repack_matrix((const uint8_t*)a.p[MLA_BASE + 11], dst + repack_off_mla_o(), 4096, 2304, 1); for (int layer = 0; layer < 4; ++layer) { int mb = MOE_BASE + layer * MOE_STRIDE; repack_matrix((const uint8_t*)a.p[mb + 1], dst + repack_off_moe(layer, 0), 2304, 1024, 64); repack_matrix((const uint8_t*)a.p[mb + 4], dst + repack_off_moe(layer, 1), 2304, 1024, 64); repack_matrix((const uint8_t*)a.p[mb + 7], dst + repack_off_moe(layer, 2), 1024, 2304, 64); repack_matrix((const uint8_t*)a.p[mb + 10], dst + repack_off_moe(layer, 3), 2304, 1024, 1); repack_matrix((const uint8_t*)a.p[mb + 13], dst + repack_off_moe(layer, 4), 2304, 1024, 1); repack_matrix((const uint8_t*)a.p[mb + 16], dst + repack_off_moe(layer, 5), 1024, 2304, 1); } } // Warp GEMV over one 128-output tile and one 128-K group, sequential per-lane weights. __device__ void gemv_warp( const uint8_t* repacked, const uint16_t* scales, const uint16_t* zeros, const float* x, float* y, int K, int N, int tile, int k0, float y_scale, int round_x, int y_delta, float* xs ) { int lane = threadIdx.x & 31; int kpairs = K >> 1; int kp0 = k0 >> 1; int g = k0 >> 7; int n = tile * 128 + lane * 4; #pragma unroll for (int t = 0; t < 4; ++t) { float xv = x[k0 + lane + t * 32]; xs[lane + t * 32] = round_x ? round_bf(xv) : xv; } __syncwarp(); float s[4], z[4], acc[4]; #pragma unroll for (int j = 0; j < 4; ++j) { s[j] = ld_bf(scales + (int64_t)g * N + n + j); z[j] = ld_bf(zeros + (int64_t)g * N + n + j); acc[j] = 0.f; } const uint32_t* mine = reinterpret_cast(repacked) + ((int64_t)tile * 32 + lane) * kpairs + kp0; #pragma unroll 1 for (int i = 0; i < 64; i += 4) { uint4 v = ld4(mine + i); uint32_t pw[4] = {v.x, v.y, v.z, v.w}; #pragma unroll for (int t = 0; t < 4; ++t) { float x0 = xs[(i + t) * 2]; float x1 = xs[(i + t) * 2 + 1]; uint32_t p = pw[t]; #pragma unroll for (int j = 0; j < 4; ++j) { uint32_t byte = (p >> (8 * j)) & 0xFFu; float q0 = float(byte & 0xFu); float q1 = float(byte >> 4); acc[j] = fmaf((q0 - z[j]) * s[j], x0, acc[j]); acc[j] = fmaf((q1 - z[j]) * s[j], x1, acc[j]); } } } #pragma unroll for (int j = 0; j < 4; ++j) atomicAdd(y + (n + j - y_delta), acc[j] * y_scale); } __global__ void __launch_bounds__(256, 2) kimi_mega_kernel(Pack a) { cg::grid_group grid = cg::this_grid(); if (blockIdx.x == 0 && threadIdx.x == 0) d_clk_n = 0; GSYNC(); stamp(3); __shared__ __align__(16) char smem[12288]; float* red = reinterpret_cast(smem); const int pos = a.pos; const int L = pos + 1; uint16_t* h = (uint16_t*)a.p[1]; const uint16_t* hin = (const uint16_t*)a.p[0]; float* ws_norm = (float*)a.p[18]; float* ws_q = (float*)a.p[19]; float* ws_k = (float*)a.p[20]; float* ws_v = (float*)a.p[21]; float* ws_g = (float*)a.p[22]; float* ws_beta = (float*)a.p[23]; float* ws_o = (float*)a.p[24]; float* ws_kv = (float*)a.p[25]; float* ws_qabs = (float*)a.p[26]; float* ws_qrope = (float*)a.p[27]; float* ws_ctx = (float*)a.p[28]; float* ws_scores = (float*)a.p[29]; float* ws_gate = (float*)a.p[30]; float* ws_up = (float*)a.p[31]; float* ws_hid = (float*)a.p[32]; float* ws_moe = (float*)a.p[33]; int* exp_idx = (int*)a.p[34]; float* exp_w = (float*)a.p[35]; int stride = blockDim.x * gridDim.x; for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < HIDDEN; i += stride) h[i] = hin[i]; GSYNC(); int* repack_flag = (int*)a.p[191]; (void)repack_flag; stamp(4); for (int layer = 0; layer < 4; ++layer) { const bool kda = layer < 3; int kb = KDA_BASE + layer * KDA_STRIDE; int mb = MOE_BASE + layer * MOE_STRIDE; const uint16_t* attn_norm = (const uint16_t*)a.p[kda ? kb : MLA_BASE]; const uint16_t* moe_norm = (const uint16_t*)a.p[kda ? kb + 1 : MLA_BASE + 1]; if (blockIdx.x == 0) rmsnorm(h, attn_norm, ws_norm, red); GSYNC(); stamp(10 + layer); if (kda) { zero_f(ws_q, KDA_C); zero_f(ws_k, KDA_C); zero_f(ws_v, KDA_C); zero_f(ws_g, KDA_C); GSYNC(); // q,k,v,g : K=2304 N=4096 int n_tiles = KDA_C / N_TILE; // 4 int k_tiles = HIDDEN / K_TILE; // 18 int per = n_tiles * k_tiles; // 72 int tasks = 4 * per; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int mat = task / per; int rem = task - mat * per; int nt = rem / k_tiles; int kt = rem - nt * k_tiles; int off = kb + 2 + mat * 3; gemv4((const uint8_t*)a.p[off], (const uint16_t*)a.p[off + 1], (const uint16_t*)a.p[off + 2], ws_norm, mat == 0 ? ws_q : mat == 1 ? ws_k : mat == 2 ? ws_v : ws_g, KDA_C, nt * N_TILE, nt * N_TILE + N_TILE, kt * K_TILE, 0, 1.f, 0, smem); } GSYNC(); if (blockIdx.x == 0) bf16_gemv_out(ws_norm, (const uint16_t*)a.p[kb + 17], ws_beta, KDA_H, HIDDEN, red, 1); GSYNC(); if (blockIdx.x < KDA_H) kda_head(blockIdx.x, ws_q, ws_k, ws_v, ws_g, ws_beta, ws_o, (float*)a.p[6 + layer], (uint16_t*)a.p[9 + layer], (uint16_t*)a.p[12 + layer], (uint16_t*)a.p[15 + layer], (const uint16_t*)a.p[kb + 18], smem); GSYNC(); // o_proj K=4096 N=2304 zero_f(ws_moe, HIDDEN); GSYNC(); { int n_tiles = (HIDDEN + N_TILE - 1) / N_TILE; // 3 int k_tiles = KDA_C / K_TILE; // 32 int tasks = n_tiles * k_tiles; int off = kb + 14; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int nt = task / k_tiles; int kt = task - nt * k_tiles; int n0 = nt * N_TILE; int n_end = n0 + N_TILE < HIDDEN ? n0 + N_TILE : HIDDEN; gemv4((const uint8_t*)a.p[off], (const uint16_t*)a.p[off + 1], (const uint16_t*)a.p[off + 2], ws_o, ws_moe, HIDDEN, n0, n_end, kt * K_TILE, 0, 1.f, 1, smem); } } GSYNC(); stamp(20 + layer); residual(h, ws_moe, HIDDEN); GSYNC(); stamp(21 + layer); } else { // MLA q (6144) + kv_a (576) zero_f(ws_q, 6144); zero_f(ws_kv, 576); GSYNC(); { int k_tiles = HIDDEN / K_TILE; // 18 int q_nt = 6144 / N_TILE; // 6 int tasks = q_nt * k_tiles + k_tiles; // q tiles + kv_a tiles int qoff = MLA_BASE + 2; int aoff = MLA_BASE + 5; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { if (task < q_nt * k_tiles) { int nt = task / k_tiles; int kt = task - nt * k_tiles; gemv4((const uint8_t*)a.p[qoff], (const uint16_t*)a.p[qoff + 1], (const uint16_t*)a.p[qoff + 2], ws_norm, ws_q, 6144, nt * N_TILE, nt * N_TILE + N_TILE, kt * K_TILE, 0, 1.f, 0, smem); } else { int kt = task - q_nt * k_tiles; gemv4((const uint8_t*)a.p[aoff], (const uint16_t*)a.p[aoff + 1], (const uint16_t*)a.p[aoff + 2], ws_norm, ws_kv, 576, 0, 576, kt * K_TILE, 0, 1.f, 0, smem); } } } GSYNC(); // rope + append new token if (blockIdx.x == 0) { int tx = threadIdx.x; uint16_t* ckv_dst = (uint16_t*)a.p[4]; uint16_t* kr_dst = (uint16_t*)a.p[5]; for (int d = tx; d < 512; d += blockDim.x) st_bf(ckv_dst + (int64_t)pos * 512 + d, ws_kv[d]); if (tx < 32) { int hh = tx; float logt = logf(ROPE_THETA); for (int i = 0; i < 32; ++i) { float e = round_bf(ws_q[hh * 192 + 128 + 2 * i]); float o = round_bf(ws_q[hh * 192 + 128 + 2 * i + 1]); float inv = expf(-logt * (2.f * i) / 64.f); float ang = float(pos) * inv; float c = cosf(ang), s = sinf(ang); ws_qrope[hh * 64 + 2 * i] = round_bf(e * c - o * s); ws_qrope[hh * 64 + 2 * i + 1] = round_bf(o * c + e * s); } } if (tx == 0) { float logt = logf(ROPE_THETA); for (int i = 0; i < 32; ++i) { float e = round_bf(ws_kv[512 + 2 * i]); float o = round_bf(ws_kv[512 + 2 * i + 1]); float inv = expf(-logt * (2.f * i) / 64.f); float ang = float(pos) * inv; float c = cosf(ang), s = sinf(ang); st_bf(kr_dst + (int64_t)pos * 64 + 2 * i, e * c - o * s); st_bf(kr_dst + (int64_t)pos * 64 + 2 * i + 1, o * c + e * s); } } } GSYNC(); if (blockIdx.x < 32) { int boff = MLA_BASE + 8; absorb_q(blockIdx.x, (const uint8_t*)a.p[boff], (const uint16_t*)a.p[boff + 1], (const uint16_t*)a.p[boff + 2], ws_q, ws_qabs, smem); } GSYNC(); stamp(50); mla_scores(a, L, pos, smem); stamp(51); GSYNC(); { float* pmax = (float*)a.p[36]; float* psum = pmax + (int64_t)gridDim.x * 32; float* gmax = psum + (int64_t)gridDim.x * 32; float* gsum = gmax + 32; softmax_partial_max(ws_scores, L, pmax, red); GSYNC(); softmax_reduce_max(pmax, gridDim.x, gmax); GSYNC(); softmax_partial_sum(ws_scores, L, gmax, psum, red); GSYNC(); softmax_reduce_sum(psum, gridDim.x, gsum); GSYNC(); softmax_write(ws_scores, L, gmax, gsum); } stamp(52); GSYNC(); stamp(53); weighted_ctx(a, L); stamp(54); GSYNC(); reduce_ctx(a); GSYNC(); // v absorb: 32 heads * 4 k-tiles, N-slice of 128 at column h*256+128 zero_f(ws_o, KDA_C); GSYNC(); { int boff = MLA_BASE + 8; int tasks = 32 * 4; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int hh = task >> 2; int kt = task & 3; int n0 = hh * 256 + 128; gemv4((const uint8_t*)a.p[boff], (const uint16_t*)a.p[boff + 1], (const uint16_t*)a.p[boff + 2], ws_ctx + hh * 512, ws_o + hh * 128, 8192, n0, n0 + 128, kt * K_TILE, n0, 1.f, 0, smem); } } GSYNC(); zero_f(ws_moe, HIDDEN); GSYNC(); { int off = MLA_BASE + 11; int n_tiles = (HIDDEN + N_TILE - 1) / N_TILE; int k_tiles = KDA_C / K_TILE; int tasks = n_tiles * k_tiles; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int nt = task / k_tiles; int kt = task - nt * k_tiles; int n0 = nt * N_TILE; int n_end = n0 + N_TILE < HIDDEN ? n0 + N_TILE : HIDDEN; gemv4((const uint8_t*)a.p[off], (const uint16_t*)a.p[off + 1], (const uint16_t*)a.p[off + 2], ws_o, ws_moe, HIDDEN, n0, n_end, kt * K_TILE, 0, 1.f, 1, smem); } } GSYNC(); residual(h, ws_moe, HIDDEN); GSYNC(); } // MoE if (blockIdx.x == 0) rmsnorm(h, moe_norm, ws_norm, red); GSYNC(); if (blockIdx.x == 0) router_topk(ws_norm, (const uint16_t*)a.p[mb], exp_idx, exp_w, smem); GSYNC(); zero_f(ws_gate, 9 * MOE_INTER); zero_f(ws_up, 9 * MOE_INTER); GSYNC(); { constexpr int k_tiles = HIDDEN / K_TILE; int tasks = 9 * 2 * k_tiles; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int slot = task / (2 * k_tiles); int rem = task - slot * 2 * k_tiles; int mat = rem / k_tiles; int kt = rem - mat * k_tiles; int e = (slot < 8) ? exp_idx[slot] : 0; int base = (slot < 8) ? (mb + 1 + mat * 3) : (mb + 10 + mat * 3); const uint8_t* wq = (const uint8_t*)a.p[base]; const uint16_t* sc = (const uint16_t*)a.p[base + 1]; const uint16_t* zc = (const uint16_t*)a.p[base + 2]; int64_t wq_e = (int64_t)e * (HIDDEN / 2) * MOE_INTER; int64_t sc_e = (int64_t)e * (HIDDEN / K_TILE) * MOE_INTER; float* y = (mat == 0 ? ws_gate : ws_up) + slot * MOE_INTER; gemv4(wq + wq_e, sc + sc_e, zc + sc_e, ws_norm, y, MOE_INTER, 0, MOE_INTER, kt * K_TILE, 0, 1.f, 0, smem); } } GSYNC(); { int n = 9 * MOE_INTER; for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += stride) ws_hid[i] = silu(ws_gate[i]) * ws_up[i]; } GSYNC(); zero_f(ws_moe, HIDDEN); GSYNC(); { constexpr int K = MOE_INTER; constexpr int k_tiles = K / K_TILE; // 8 int n_tiles = (HIDDEN + N_TILE - 1) / N_TILE; // 3 int per = n_tiles * k_tiles; int tasks = 9 * per; for (int task = blockIdx.x; task < tasks; task += gridDim.x) { int slot = task / per; int rem = task - slot * per; int nt = rem / k_tiles; int kt = rem - nt * k_tiles; int e = (slot < 8) ? exp_idx[slot] : 0; int base = (slot < 8) ? (mb + 7) : (mb + 16); const uint8_t* wq = (const uint8_t*)a.p[base]; const uint16_t* sc = (const uint16_t*)a.p[base + 1]; const uint16_t* zc = (const uint16_t*)a.p[base + 2]; int64_t wq_e = (int64_t)e * (K / 2) * HIDDEN; int64_t sc_e = (int64_t)e * (K / K_TILE) * HIDDEN; int n0 = nt * N_TILE; int n_end = n0 + N_TILE < HIDDEN ? n0 + N_TILE : HIDDEN; float scale = (slot < 8) ? exp_w[slot] : 1.f; gemv4(wq + wq_e, sc + sc_e, zc + sc_e, ws_hid + slot * K, ws_moe, HIDDEN, n0, n_end, kt * K_TILE, 0, scale, 0, smem); } } GSYNC(); residual(h, ws_moe, HIDDEN); GSYNC(); stamp(80 + layer); } stamp(99); } static char errbuf[512]; extern "C" int launch_kimi_mega(const int64_t* pack, int pos, cudaStream_t stream) { Pack pk; for (int i = 0; i < 192; ++i) pk.p[i] = pack[i]; pk.pos = pos; pk.pad = 0; int sms = 0; cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, 0); int bpsm = 0; cudaError_t oe = cudaOccupancyMaxActiveBlocksPerMultiprocessor(&bpsm, kimi_mega_kernel, THREADS, 0); if (oe != cudaSuccess) { snprintf(errbuf, sizeof(errbuf), "occupancy: %s", cudaGetErrorString(oe)); return 2; } int blocks = bpsm * sms; if (blocks > MAX_BLOCKS) blocks = MAX_BLOCKS; if (blocks < 1) { snprintf(errbuf, sizeof(errbuf), "no resident blocks (bpsm=%d)", bpsm); return 3; } static int announced = 0; if (!announced) { announced = 1; fprintf(stderr, "kimi_mega blocks=%d bpsm=%d sms=%d\n", blocks, bpsm, sms); } void* args[] = {&pk}; cudaError_t e = cudaLaunchCooperativeKernel((void*)kimi_mega_kernel, dim3(blocks), dim3(THREADS), args, 0, stream); if (e != cudaSuccess) { snprintf(errbuf, sizeof(errbuf), "launch: %s (blocks=%d bpsm=%d)", cudaGetErrorString(e), blocks, bpsm); return 4; } return 0; } extern "C" const char* kimi_last_error() { return errbuf; } extern "C" int kimi_copy_clocks(unsigned long long* host, int* tags, int n) { int cnt = 0; cudaMemcpyFromSymbol(&cnt, d_clk_n, sizeof(int)); if (cnt > n) cnt = n; if (cnt > 0) { cudaMemcpyFromSymbol(host, d_clk, sizeof(unsigned long long) * cnt); cudaMemcpyFromSymbol(tags, d_tag, sizeof(int) * cnt); } return cnt; } """ digest = hashlib.sha256(cuda_src.encode()).hexdigest()[:16] cache = os.path.join(os.path.dirname(os.path.abspath(__file__)), ".build") os.makedirs(cache, exist_ok=True) so = os.path.join(cache, f"kimi_mega_{digest}.so") cu = os.path.join(cache, f"kimi_mega_{digest}.cu") if not os.path.exists(so): with open(cu, "w") as f: f.write(cuda_src) nvcc = os.environ.get("NVCC", "/usr/local/cuda/bin/nvcc") subprocess.check_call( [ nvcc, "-O3", "-std=c++17", "-arch=sm_120", "-Xcompiler", "-fPIC", "--shared", "-Xptxas=-v", "-o", so, cu, ] ) lib = ctypes.CDLL(so) lib.launch_kimi_mega.argtypes = [ctypes.c_void_p, ctypes.c_int, ctypes.c_void_p] lib.launch_kimi_mega.restype = ctypes.c_int lib.kimi_last_error.restype = ctypes.c_char_p lib.kimi_copy_clocks.argtypes = [ctypes.POINTER(ctypes.c_uint64), ctypes.POINTER(ctypes.c_int), ctypes.c_int] lib.kimi_copy_clocks.restype = ctypes.c_int lib._copy = lib.kimi_copy_clocks class _Wrap: def clocks(self): buf = (ctypes.c_uint64 * 256)() tags = (ctypes.c_int * 256)() n = lib._copy(buf, tags, 256) return [(tags[i], buf[i]) for i in range(n)] def launch(self, pack, pos): stream = torch.cuda.current_stream().cuda_stream rc = lib.launch_kimi_mega(ctypes.c_void_p(pack.data_ptr()), int(pos), ctypes.c_void_p(stream)) if rc != 0: err = lib.kimi_last_error() raise RuntimeError(err.decode() if err else f"kimi_mega rc={rc}") _EXT = _Wrap() return _EXT _EXT = None class QuantLinear(nn.Module): def __init__(self, in_f: int, out_f: int, group: int = GROUP): 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)) class QuantExperts(nn.Module): def __init__(self, n: int, in_f: int, out_f: int, group: int = GROUP): 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)) class KDA(nn.Module): def __init__(self, cfg): super().__init__() 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) class MLA(nn.Module): def __init__(self, cfg): super().__init__() 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) class MoE(nn.Module): def __init__(self, cfg): super().__init__() 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, 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): super().__init__() self.cfg = cfg self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern) self._ready = False def _qw(self, mod, pack, i): pack[i] = mod.w_q.data_ptr() pack[i + 1] = mod.scales.data_ptr() pack[i + 2] = mod.zeros.data_ptr() def _ensure(self, device): if self._ready and self._dev == device: return self._dev = device self.buf_a = torch.empty(HIDDEN, device=device, dtype=torch.bfloat16) self.buf_b = torch.empty(HIDDEN, device=device, dtype=torch.bfloat16) self.ws_norm = torch.empty(HIDDEN, device=device, dtype=torch.float32) self.ws_q = torch.empty(MLA_Q, device=device, dtype=torch.float32) self.ws_k = torch.empty(KDA_C, device=device, dtype=torch.float32) self.ws_v = torch.empty(KDA_C, device=device, dtype=torch.float32) self.ws_g = torch.empty(KDA_C, device=device, dtype=torch.float32) self.ws_beta = torch.empty(KDA_H, device=device, dtype=torch.float32) self.ws_o = torch.empty(KDA_C, device=device, dtype=torch.float32) self.ws_kv = torch.empty(KV_A, device=device, dtype=torch.float32) self.ws_qabs = torch.empty(512 * 32, device=device, dtype=torch.float32) self.ws_qrope = torch.empty(32 * 64, device=device, dtype=torch.float32) self.ws_ctx = torch.empty(32 * 512, device=device, dtype=torch.float32) self.ws_scores = torch.empty(CACHE_CAP * 32, device=device, dtype=torch.float32) self.ws_gate = torch.empty(9 * MOE_INTER, device=device, dtype=torch.float32) self.ws_up = torch.empty(9 * MOE_INTER, device=device, dtype=torch.float32) self.ws_hid = torch.empty(9 * MOE_INTER, device=device, dtype=torch.float32) self.ws_moe = torch.empty(HIDDEN, device=device, dtype=torch.float32) self.exp_idx = torch.empty(8, device=device, dtype=torch.int32) self.exp_w = torch.empty(9, device=device, dtype=torch.float32) self.ws_partial = torch.empty(MAX_BLOCKS * 32 * 512, device=device, dtype=torch.float32) self.ckv_store = torch.empty(CACHE_CAP, 512, device=device, dtype=torch.bfloat16) self.krope_store = torch.empty(CACHE_CAP, 64, device=device, dtype=torch.bfloat16) pack = torch.zeros(PACK_N, dtype=torch.int64) pack[18] = self.ws_norm.data_ptr() pack[19] = self.ws_q.data_ptr() pack[20] = self.ws_k.data_ptr() pack[21] = self.ws_v.data_ptr() pack[22] = self.ws_g.data_ptr() pack[23] = self.ws_beta.data_ptr() pack[24] = self.ws_o.data_ptr() pack[25] = self.ws_kv.data_ptr() pack[26] = self.ws_qabs.data_ptr() pack[27] = self.ws_qrope.data_ptr() pack[28] = self.ws_ctx.data_ptr() pack[29] = self.ws_scores.data_ptr() pack[30] = self.ws_gate.data_ptr() pack[31] = self.ws_up.data_ptr() pack[32] = self.ws_hid.data_ptr() pack[33] = self.ws_moe.data_ptr() pack[34] = self.exp_idx.data_ptr() pack[35] = self.exp_w.data_ptr() pack[36] = self.ws_partial.data_ptr() kda_i = 0 for li, blk in enumerate(self.blocks): if blk.kind == "K": b = KDA_BASE + kda_i * KDA_STRIDE pack[b] = blk.attn_norm.data_ptr() pack[b + 1] = blk.moe_norm.data_ptr() attn = blk.attn self._qw(attn.q_proj, pack, b + 2) self._qw(attn.k_proj, pack, b + 5) self._qw(attn.v_proj, pack, b + 8) self._qw(attn.g_proj, pack, b + 11) self._qw(attn.o_proj, pack, b + 14) pack[b + 17] = attn.beta_proj.weight.data_ptr() pack[b + 18] = attn.conv_w.data_ptr() kda_i += 1 else: b = MLA_BASE pack[b] = blk.attn_norm.data_ptr() pack[b + 1] = blk.moe_norm.data_ptr() attn = blk.attn self._qw(attn.q_proj, pack, b + 2) self._qw(attn.kv_a, pack, b + 5) self._qw(attn.kv_b, pack, b + 8) self._qw(attn.o_proj, pack, b + 11) mb = MOE_BASE + li * MOE_STRIDE moe = blk.moe pack[mb] = moe.router.weight.data_ptr() self._qw(moe.gate, pack, mb + 1) self._qw(moe.up, pack, mb + 4) self._qw(moe.down, pack, mb + 7) self._qw(moe.s_gate, pack, mb + 10) self._qw(moe.s_up, pack, mb + 13) self._qw(moe.s_down, pack, mb + 16) self._pack = pack self._mla = self.cfg.pattern.index("M") self._ready = True _ext() def step(self, hidden, state): self._ensure(hidden.device) if hidden.data_ptr() == self.buf_a.data_ptr(): out = self.buf_b else: out = self.buf_a p = self._pack p[0] = hidden.data_ptr() p[1] = out.data_ptr() mla = self._mla c_kv = state[mla]["c_kv"] k_rope = state[mla]["k_rope"] pos = c_kv.shape[0] p[2] = c_kv.data_ptr() p[3] = k_rope.data_ptr() p[4] = self.ckv_store.data_ptr() p[5] = self.krope_store.data_ptr() kda_i = 0 for li, blk in enumerate(self.blocks): if blk.kind != "K": continue st = state[li] p[6 + kda_i] = st["S"].data_ptr() p[9 + kda_i] = st["cq"].data_ptr() p[12 + kda_i] = st["ck"].data_ptr() p[15 + kda_i] = st["cv"].data_ptr() kda_i += 1 _ext().launch(p, pos) new_len = pos + 1 state[mla]["c_kv"] = self.ckv_store[:new_len] state[mla]["k_rope"] = self.krope_store[:new_len] return out, state