KernelBench mega · RTX PRO 6000

Kimi-Linear Decode Grok 4.7

6.38×geomean speedup across shapes

manually audited: clean

Grok 4.7 built a genuine single-launch cooperative megakernel for the Kimi-Linear decode block in 63 minutes: in-kernel int4 dequant-GEMVs, short conv, gated-delta recurrence, absorbed MLA with a grid-parallel softmax, and the full 64-expert MoE, all behind one `cudaLaunchCooperativeKernel`. torch.profiler confirms 1.00 launches per step. Correctness is perfect (1.0000 on output, KDA state and MLA cache across all six check cases) and the templates are untouched. It is simply slow for this problem - 5.65x contended - and it stopped after one official benchmark.py run.

harnessgrok
Kernel source (redacted)
"""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 <cuda_runtime.h>
#include <cooperative_groups.h>
#include <cstdint>
#include <cstdio>
#include <algorithm>

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<float*>(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<float*>(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<float*>(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<float*>(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<float*>(smem);
    uint16_t* cchunk = reinterpret_cast<uint16_t*>(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<float*>(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<const uint32_t*>(es + ((int64_t)kp * N + tile * 128));
            uint32_t* d2 = reinterpret_cast<uint32_t*>(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<const uint32_t*>(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<float*>(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

20260917_005735_grok_grok-4.7_02_kimi_linear_decode