kernelbench.com

KernelBench mega · H100

Kimi-Linear Decode GLM-5.2

4.52×geomean speedup across shapes

manually audited: clean

harnesszai-claude
Kernel source (redacted)
"""Fused single-kernel W4A16 decode for the Kimi-Linear hybrid unit (batch=1).

The whole per-token forward (3 KDA + 1 MLA layer, each with a 64-expert MoE,
all int4 dequant-GEMVs, short causal conv, KDA recurrent-state update, MLA
latent-cache absorb attention, MoE router + expert GEMVs, both RMSNorms and
residual adds) is one custom CUDA __global__ kernel built with load_inline and
invoked exactly once in step(). No CUDA graph, no torch.compile, no per-op loop.

int4 weights are streamed once through a fused dequant-GEMV (warp-per-output,
out-major layout, per-group dequant in registers); the bf16 weight is never
materialized. MLA uses the absorb form so kv_b is never materialized across the
context. RMSNorm is folded into each GEMV's input load (sum-of-squares inline).
Cross-block sync uses a monotonic-generation atomic barrier; the grid is sized
to the GPU's max-resident block count so it is deadlock-free.
"""
from __future__ import annotations

import os
from dataclasses import dataclass, field

os.environ.setdefault("CUDA_HOME", "/usr/local/cuda")

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline

OP_TYPE = "kimi_linear_w4a16_decode"
HARDWARE_REQUIRED = ["RTX_PRO_6000"]
EPS = 1.0e-6
GROUP = 128

H = 32; DK = 128; C = H * DK; HID = 2304
QK_NOPE = 128; QK_ROPE = 64; VHD = 128; KV_LORA = 512
M = 1024; NSH = 1
LMAX_ALLOC = 32768  # max context + slack for MLA cache buffers


@dataclass(frozen=True)
class Config:
    hidden: int = HID
    kda_heads: int = H
    kda_head_dim: int = DK
    short_conv: int = 4
    mla_heads: int = H
    kv_lora: int = KV_LORA
    qk_nope: int = QK_NOPE
    qk_rope: int = QK_ROPE
    v_head: int = VHD
    rope_theta: float = 10000.0
    n_experts: int = 64
    n_active: int = 8
    n_shared: int = NSH
    moe_inter: int = M
    routed_scaling: float = 2.446
    group: int = GROUP
    pattern: tuple = ("K", "K", "K", "M")
    dtype: torch.dtype = field(default=torch.bfloat16)


def build_config(shape):
    return Config(n_experts=int(shape.get("n_experts", 64)))


# --------------------------------------------------------------------------- #
# Module structure -- mirrors reference.py exactly (names/shapes) so the
# reference state_dict loads with strict=True.
# --------------------------------------------------------------------------- #
class QuantLinear(nn.Module):
    def __init__(self, in_f, out_f, group=GROUP):
        super().__init__(); 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, in_f, out_f, group=GROUP):
        super().__init__(); 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__(); self.cfg = cfg
        self.q_proj = QuantLinear(HID, C); self.k_proj = QuantLinear(HID, C)
        self.v_proj = QuantLinear(HID, C); self.g_proj = QuantLinear(HID, C)
        self.beta_proj = nn.Linear(HID, H, bias=False, dtype=cfg.dtype)
        self.conv_w = nn.Parameter(torch.empty(3, C, cfg.short_conv, dtype=cfg.dtype))
        self.o_proj = QuantLinear(C, HID); self.scale = DK ** -0.5


class MLA(nn.Module):
    def __init__(self, cfg):
        super().__init__(); self.cfg = cfg
        self.q_proj = QuantLinear(HID, H * (QK_NOPE + QK_ROPE))
        self.kv_a = QuantLinear(HID, KV_LORA + QK_ROPE)
        self.kv_b = QuantLinear(KV_LORA, H * (QK_NOPE + VHD))
        self.o_proj = QuantLinear(H * VHD, HID); self.scale = (QK_NOPE + QK_ROPE) ** -0.5


class MoE(nn.Module):
    def __init__(self, cfg):
        super().__init__(); E = cfg.n_experts
        self.router = nn.Linear(HID, E, bias=False, dtype=cfg.dtype)
        self.gate = QuantExperts(E, HID, M); self.up = QuantExperts(E, HID, M)
        self.down = QuantExperts(E, M, HID)
        self.s_gate = QuantExperts(NSH, HID, M); self.s_up = QuantExperts(NSH, HID, M)
        self.s_down = QuantExperts(NSH, M, HID)


class Block(nn.Module):
    def __init__(self, cfg, kind):
        super().__init__(); self.kind = kind
        self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
        self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype))
        self.attn = KDA(cfg) if kind == "K" else MLA(cfg)
        self.moe = MoE(cfg)


# --------------------------------------------------------------------------- #
# Weight layout codegen (single source of truth).
# --------------------------------------------------------------------------- #
def _build_layout():
    ql = []
    for b in range(3):
        ql += [(f"K{b}Q", HID, C, 1), (f"K{b}K", HID, C, 1), (f"K{b}V", HID, C, 1),
               (f"K{b}G", HID, C, 1), (f"K{b}O", C, HID, 1)]
    ql += [("MQ", HID, H * (QK_NOPE + QK_ROPE), 1),
           ("MVA", HID, KV_LORA + QK_ROPE, 1), ("MO", H * VHD, HID, 1)]
    for b in range(4):
        ql += [(f"B{b}GATE", HID, M, 64), (f"B{b}UP", HID, M, 64), (f"B{b}DOWN", M, HID, 64),
               (f"B{b}SGATE", HID, M, NSH), (f"B{b}SUP", HID, M, NSH), (f"B{b}SDOWN", M, HID, NSH)]
    wq_off = sc_off = zs_off = 0
    off_wq = {}; off_sc = {}; off_zs = {}; meta = {}
    lines = []
    for name, in_f, out_f, e in ql:
        ng = in_f // GROUP
        off_wq[name] = wq_off; off_sc[name] = sc_off; off_zs[name] = zs_off
        meta[name] = (in_f, out_f, e)
        lines.append(f"#define WQ_{name} {wq_off}")
        lines.append(f"#define SC_{name} {sc_off}")
        lines.append(f"#define ZS_{name} {zs_off}")
        wq_off += e * (in_f // 2) * out_f
        sc_off += e * ng * out_f
        zs_off += e * ng * out_f
    return ql, off_wq, off_sc, off_zs, meta, "\n".join(lines), (wq_off, sc_off, zs_off)


_QLIST, _OFFWQ, _OFFSC, _OFFZS, _META, _DEFINES, _WQSZ = _build_layout()


def _bf_layout():
    off = {}; lines = []; cur = 0
    def emit(name, n):
        nonlocal cur
        off[name] = cur; lines.append(f"#define BF_{name} {cur}"); cur += n
    for b in range(4):
        emit(f"AN{b}", HID); emit(f"MN{b}", HID)
    for b in range(3):
        emit(f"BETA{b}", H * HID); emit(f"CONV{b}", 3 * C * 4)
    for b in range(4):
        emit(f"ROUT{b}", 64 * HID)
    return off, "\n".join(lines), cur


_OFFBF, _BF_DEFINES, _BFSZ = _bf_layout()

# --------------------------------------------------------------------------- #
# CUDA helpers + main kernel
# --------------------------------------------------------------------------- #
_CONST = """
#define H 32
#define DK 128
#define C 4096
#define HID 2304
#define QK_NOPE 128
#define QK_ROPE 64
#define VHD 128
#define KV_LORA 512
#define M 1024
#define NSH 1
#define GROUP 128
#define WARPS 8
#define THREADS 256
#define KDA_SMEM 4368
#define EPSF 1.0e-6f
#define ROUTED_SCALING 2.446f
typedef unsigned int u32; typedef __nv_bfloat16 bf16;
"""

_HELPERS = r"""
__device__ inline float b2f(bf16 v){return __bfloat162float(v);}
__device__ inline bf16 f2b(float v){return __float2bfloat16(v);}
__device__ inline float wsum(float v){for(int o=16;o>0;o>>=1)v+=__shfl_xor_sync(0xffffffffu,v,o);return v;}
__device__ inline float bsum(float v,volatile float* red){
    int tid=threadIdx.x; v=wsum(v); if((tid&31)==0)red[tid>>5]=v; __syncthreads();
    v=(tid<WARPS)?red[tid]:0.0f; if(tid<32){v=wsum(v);if(tid==0)red[0]=v;} __syncthreads(); return red[0];
}
__device__ inline void gbar(u32* cnt,int phase,int gen){
    __syncthreads(); __threadfence();
    if(threadIdx.x==0){ u32 want=(u32)((gen+1)*gridDim.x); u32* c=cnt+phase;
        (void)atomicAdd(c,1u); while(*((volatile u32*)c)<want){} }
    __syncthreads();
}
// load fp32 x[K] -> smem sx; fold rmsnorm if wnorm!=0 (sumsq inline)
__device__ inline void load_fold(float* sx,const float* x,const bf16* wnorm,int K,volatile float* red){
    int tid=threadIdx.x;
    if(wnorm){ float part=0.0f;
        for(int i=tid;i<K;i+=THREADS){float v=x[i]; sx[i]=v; part+=v*v;}
        __syncthreads(); float ss=bsum(part,red); float rss=rsqrtf(ss/(float)HID+EPSF);
        for(int i=tid;i<K;i+=THREADS) sx[i]=sx[i]*b2f(wnorm[i])*rss;
    } else { for(int i=tid;i<K;i+=THREADS) sx[i]=x[i]; }
    __syncthreads();
}
// dequant contribution of one uint32 (8 int4) given x base, scale, zero
__device__ inline float dq_u32(uint32_t u,const float* xb,float ss,float zz){
    float a=0;
    #pragma unroll
    for(int b=0;b<4;b++){ uint32_t by=(u>>(b<<3))&0xFFu;
        a += xb[b*2]*(((float)(by&0xF)-zz)*ss) + xb[b*2+1]*(((float)((by>>4)&0xF)-zz)*ss);
    }
    return a;
}
// vectorized dequant dot over one weight row (uint32 loads). 2 accumulators for ILP.
__device__ inline float gemv_dot_v(const float* x,const uint8_t* wq_row,const bf16* sc,const bf16* zr,int K,int lane){
    const uint32_t* r=(const uint32_t*)wq_row; int K8=K>>3; float acc0=0,acc1=0; int j=lane;
    for(; j+32<K8; j+=64){
        acc0 += dq_u32(r[j],   x+(j<<3),     b2f(sc[j>>4]),    b2f(zr[j>>4]));
        int j2=j+32; acc1 += dq_u32(r[j2], x+(j2<<3), b2f(sc[j2>>4]), b2f(zr[j2>>4]));
    }
    for(; j<K8; j+=32) acc0 += dq_u32(r[j], x+(j<<3), b2f(sc[j>>4]), b2f(zr[j>>4]));
    return wsum(acc0+acc1);
}
// out-major dequant GEMV, warp-per-output. wq[N,K/2] sc/zr[N,ng].
__device__ inline void gemv_om(const float* sx,const uint8_t* wq,const bf16* sc,const bf16* zr,int K,int N,float* y,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,ng=K/GROUP,gw=bid*WARPS+warp,TW=gridDim.x*WARPS;
    for(int out=gw;out<N;out+=TW){
        float acc=gemv_dot_v(sx,wq+(size_t)out*(K/2),sc+(size_t)out*ng,zr+(size_t)out*ng,K,lane);
        if(lane==0) y[out]=acc;
    }
}
__device__ inline void gemv_conv(const float* sx,const uint8_t* wq,const bf16* sc,const bf16* zr,int K,int N,float* y,int bid,bf16* cwin,const bf16* convw){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,ng=K/GROUP,gw=bid*WARPS+warp,TW=gridDim.x*WARPS;
    for(int out=gw;out<N;out+=TW){
        float acc=gemv_dot_v(sx,wq+(size_t)out*(K/2),sc+(size_t)out*ng,zr+(size_t)out*ng,K,lane);
        if(lane==0){
            float w0=b2f(convw[out*4]),w1=b2f(convw[out*4+1]),w2=b2f(convw[out*4+2]),w3=b2f(convw[out*4+3]);
            float p0=b2f(cwin[0*N+out]),p1=b2f(cwin[1*N+out]),p2=b2f(cwin[2*N+out]);
            float co=acc*w3+p0*w0+p1*w1+p2*w2; co=co/(1.0f+__expf(-co));
            y[out]=co; cwin[0*N+out]=f2b(p1); cwin[1*N+out]=f2b(p2); cwin[2*N+out]=f2b(acc);
        }
    }
}
__device__ inline void gemv_res(const float* sx,const uint8_t* wq,const bf16* sc,const bf16* zr,int K,int N,float* y,int bid,const float* xres){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,ng=K/GROUP,gw=bid*WARPS+warp,TW=gridDim.x*WARPS;
    for(int out=gw;out<N;out+=TW){
        float acc=gemv_dot_v(sx,wq+(size_t)out*(K/2),sc+(size_t)out*ng,zr+(size_t)out*ng,K,lane);
        if(lane==0) y[out]=acc+xres[out];
    }
}
__device__ inline void gemv_bf(const float* sx,const bf16* w,int K,int N,float* y,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,gw=bid*WARPS+warp,TW=gridDim.x*WARPS;
    for(int out=gw;out<N;out+=TW){ const bf16* row=w+(size_t)out*K; float acc=0;
        for(int k=lane;k<K;k+=32) acc+=sx[k]*b2f(row[k]);
        acc=wsum(acc); if(lane==0) y[out]=acc; }
}
// y[m]=silu(gate_dot)*up_dot  (single expert)
__device__ inline void gemv_gateup1(const float* sx,const uint8_t* gwq,const bf16* gsc,const bf16* gzr,
        const uint8_t* uwq,const bf16* usc,const bf16* uzr,int K,int N,float* y,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,ng=K/GROUP,gw=bid*WARPS+warp,TW=gridDim.x*WARPS;
    for(int out=gw;out<N;out+=TW){
        float ga=gemv_dot_v(sx,gwq+(size_t)out*(K/2),gsc+(size_t)out*ng,gzr+(size_t)out*ng,K,lane);
        float ua=gemv_dot_v(sx,uwq+(size_t)out*(K/2),usc+(size_t)out*ng,uzr+(size_t)out*ng,K,lane);
        if(lane==0){float s=ga/(1.0f+__expf(-ga)); y[out]=s*ua;}
    }
}
// warp-level dequant dot (vectorized). x may be global.
__device__ inline float dq_dot(const float* x,int K,const uint8_t* wq_row,const bf16* sc_row,const bf16* zr_row){
    return gemv_dot_v(x,wq_row,sc_row,zr_row,K,threadIdx.x&31);
}
// KDA recurrent for head h=bid (if bid<H)
__device__ inline void kda_recur(float* S,float* q,float* k,float* v,float* g,float* o,float* beta,
                                 float* smem,volatile float* red,int bid){
    if(bid>=H) return;
    int h=bid,tid=threadIdx.x; float scalef=rsqrtf((float)DK);
    float* shk=smem+KDA_SMEM; float* shv=shk+DK; float* shq=shv+DK; float* shd=shq+DK; float* shp=shd+DK; float* sho=shp+DK;
    for(int d=tid;d<DK;d+=THREADS){int idx=h*DK+d; shk[d]=k[idx]; shv[d]=v[idx]; shq[d]=q[idx]*scalef; shd[d]=1.0f/(1.0f+__expf(g[idx]));}
    __syncthreads();
    float bh=1.0f/(1.0f+__expf(-beta[h]));
    for(int jd=tid;jd<DK*DK;jd+=THREADS){int j=jd>>7,d=jd&127; S[h*DK*DK+j*DK+d]*=shd[j];}
    __threadfence_block(); __syncthreads();
    for(int d=tid;d<DK;d+=THREADS) shp[d]=0; __syncthreads();
    for(int j=0;j<DK;j++){float kj=shk[j]; for(int d=tid;d<DK;d+=THREADS) shp[d]+=S[h*DK*DK+j*DK+d]*kj; __syncthreads();}
    for(int jd=tid;jd<DK*DK;jd+=THREADS){int j=jd>>7,d=jd&127; S[h*DK*DK+j*DK+d]+=bh*shk[j]*(shv[d]-shp[d]);}
    __threadfence_block(); __syncthreads();
    for(int d=tid;d<DK;d+=THREADS) sho[d]=0; __syncthreads();
    for(int j=0;j<DK;j++){float qj=shq[j]; for(int d=tid;d<DK;d+=THREADS) sho[d]+=S[h*DK*DK+j*DK+d]*qj; __syncthreads();}
    for(int d=tid;d<DK;d+=THREADS) o[h*DK+d]=sho[d];
    __syncthreads();
}
// MLA rope (32 pairs) applied to vec[64]; pos known to caller via L
__device__ inline void rope_apply(float* vec,int pos){
    int tid=threadIdx.x;
    for(int i=tid;i<QK_ROPE/2;i+=THREADS){
        float inv=1.0f/powf(10000.0f,(float)(2*i)/(float)QK_ROPE);
        float ang=(float)pos*inv; float cs=cosf(ang),sn=sinf(ang);
        int a0=2*i,a1=2*i+1; float e=vec[a0],od=vec[a1];
        vec[a0]=e*cs-od*sn; vec[a1]=od*cs+e*sn;
    }
    __syncthreads();
}
// MLA phase B: rope q_rope+k_rope, append cache (with prime copy), then q_absorb
__device__ inline void mla_rope_append_qabs(float* qfull,float* kvfull,bf16* ckv,bf16* krc,
        const bf16* ckvin,const bf16* krcin,int L,int prime,float* qa,
        const uint8_t* ukwq,const bf16* uksc,const bf16* ukzr,int bid){
    int tid=threadIdx.x;
    // rope on q_rope [H,64] (per head contig at h*192+128) -- treat as H*64 vector? no, each head 64 with shared cos/sin
    // We rope the whole q_rope region as [H][64]; cos/sin same for all heads at given pos.
    for(int i=tid;i<H*(QK_ROPE/2);i+=THREADS){
        int hh=i/(QK_ROPE/2), pi=i%(QK_ROPE/2);
        float inv=1.0f/powf(10000.0f,(float)(2*pi)/(float)QK_ROPE); float ang=(float)L*inv;
        float cs=cosf(ang),sn=sinf(ang); int a0=hh*(QK_NOPE+QK_ROPE)+QK_NOPE+2*pi, a1=a0+1;
        float e=qfull[a0],od=qfull[a1]; qfull[a0]=e*cs-od*sn; qfull[a1]=od*cs+e*sn;
    }
    // rope on k_rope_new kvfull[512..575]
    for(int i=tid;i<QK_ROPE/2;i+=THREADS){
        float inv=1.0f/powf(10000.0f,(float)(2*i)/(float)QK_ROPE); float ang=(float)L*inv;
        float cs=cosf(ang),sn=sinf(ang); int a0=KV_LORA+2*i,a1=a0+1; float e=kvfull[a0],od=kvfull[a1];
        kvfull[a0]=e*cs-od*sn; kvfull[a1]=od*cs+e*sn;
    }
    __syncthreads();
    // prime copy incoming cache -> bigbuf
    if(prime){
        for(int i=tid;i<L*KV_LORA;i+=THREADS){ ((bf16*)ckv)[i]=ckvin[i]; }
        for(int i=tid;i<L*QK_ROPE;i+=THREADS){ ((bf16*)krc)[i]=krcin[i]; }
        __syncthreads();
    }
    // append row L
    for(int i=tid;i<KV_LORA;i+=THREADS) ckv[(size_t)L*KV_LORA+i]=f2b(kvfull[i]);
    for(int i=tid;i<QK_ROPE;i+=THREADS) krc[(size_t)L*QK_ROPE+i]=f2b(kvfull[KV_LORA+i]);
    __syncthreads();
    // q_absorb: qa[i,h]=sum_d uk[h,i,d]*q_nope[h,d]; uk in-major [H,256,128]
    {
      int warp=tid>>5,lane=tid&31,gw=bid*WARPS+warp,TW=gridDim.x*WARPS; int N=KV_LORA*H;
      for(int out=gw;out<N;out+=TW){
        int i=out/H, h=out%H; int byteb=i/2, grp=i/128; int odd=i&1;
        const uint8_t* row=ukwq+(size_t)h*256*128+(size_t)byteb*128;
        const bf16* s=uksc+(size_t)h*4*128+(size_t)grp*128; const bf16* z=ukzr+(size_t)h*4*128+(size_t)grp*128;
        const float* qn=qfull+h*(QK_NOPE+QK_ROPE);
        float acc=0;
        for(int d=lane;d<QK_NOPE;d+=32){ u32 by=row[d]; float nib=(odd?((float)((by>>4)&0xF)):((float)(by&0xF)));
            acc+=qn[d]*((nib-b2f(z[d]))*b2f(s[d])); }
        acc=wsum(acc); if(lane==0) qa[out]=acc;
      }
    }
}
// MLA scores: scores[l,h]=scale*( c_kv[l]@qa[:,h] + k_rope[l]@qrope[h] )
__device__ inline void mla_scores(const float* qfull,const float* qa,const bf16* ckv,const bf16* krc,float* scores,int L,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,gw=bid*WARPS+warp,TW=gridDim.x*WARPS; int Lp=L+1;
    float scalef=rsqrtf((float)(QK_NOPE+QK_ROPE));
    for(int l=gw;l<Lp;l+=TW){
        const bf16* ckl=ckv+(size_t)l*KV_LORA; const bf16* krl=krc+(size_t)l*QK_ROPE;
        // lane = head h
        float nope=0;
        for(int i=0;i<KV_LORA;i++){ float cv=b2f(ckl[i]); nope += cv*qa[i*H+lane]; }
        // rope part for this head
        const float* qr=qfull+lane*(QK_NOPE+QK_ROPE)+QK_NOPE;
        float rp=0; for(int d=0;d<QK_ROPE;d++) rp+=b2f(krl[d])*qr[d];
        float sc=(nope+rp)*scalef;
        if(lane<H) scores[l*H+lane]=sc;
    }
}
// softmax over l (per head): one block per head (bid<H). 2-pass, block-internal reduce.
__device__ inline void mla_softmax(const float* scores,float* p,int L,int bid,volatile float* red){
    if(bid>=H) return;
    int h=bid, tid=threadIdx.x, Lp=L+1;
    float mx=-1e30f;
    for(int l=tid;l<Lp;l+=THREADS) mx=fmaxf(mx,scores[l*H+h]);
    mx=bsum(mx,red);
    float sm=0;
    for(int l=tid;l<Lp;l+=THREADS){ float e=__expf(scores[l*H+h]-mx); p[l*H+h]=e; sm+=e; }
    sm=bsum(sm,red);
    for(int l=tid;l<Lp;l+=THREADS) p[l*H+h]=p[l*H+h]/sm;
}
// v_accum[h,i]=sum_l p[l,h]*c_kv[l,i]; split over i-tiles AND l-chunks, atomicAdd reduce.
__device__ inline void mla_vaccum(const float* p,const bf16* ckv,float* vacc,int L,int bid){
    const int IW=32, NT_I=KV_LORA/IW;  // 16
    const int NT_L=8;
    if(bid >= NT_I*NT_L) return;
    int it=bid%NT_I, lc=bid/NT_I, ibase=it*IW;
    int tid=threadIdx.x; int il=tid&31; int h0=tid>>5;
    int Lp=L+1, lper=(Lp+NT_L-1)/NT_L, l0=lc*lper, l1=l0+lper; if(l1>Lp)l1=Lp;
    float a0=0,a1=0,a2=0,a3=0;
    for(int l=l0;l<l1;l++){
        float cv=b2f(ckv[(size_t)l*KV_LORA+ibase+il]);
        const float* pl=p+(size_t)l*H;
        a0+=pl[h0]*cv; a1+=pl[h0+8]*cv; a2+=pl[h0+16]*cv; a3+=pl[h0+24]*cv;
    }
    atomicAdd(&vacc[(size_t)h0*KV_LORA+ibase+il], a0);
    atomicAdd(&vacc[(size_t)(h0+8)*KV_LORA+ibase+il], a1);
    atomicAdd(&vacc[(size_t)(h0+16)*KV_LORA+ibase+il], a2);
    atomicAdd(&vacc[(size_t)(h0+24)*KV_LORA+ibase+il], a3);
}
// MLA o: o[h,d]=sum_i vacc[h,i]*kvb_UV[h,d,i]; kvb_UV_T out-major [H,128,256]
__device__ inline void mla_o(const float* vacc,const uint8_t* wq,const bf16* sc,const bf16* zr,float* ofull,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,gw=bid*WARPS+warp,TW=gridDim.x*WARPS; int N=H*VHD;
    int ng=KV_LORA/GROUP;
    for(int out=gw;out<N;out+=TW){
        int h=out/VHD, d=out%VHD;
        const uint8_t* row=wq+(size_t)h*VHD*(KV_LORA/2)+(size_t)d*(KV_LORA/2);
        const bf16* s=sc+(size_t)h*VHD*ng+(size_t)d*ng; const bf16* z=zr+(size_t)h*VHD*ng+(size_t)d*ng;
        float acc=gemv_dot_v(vacc+h*KV_LORA,row,s,z,KV_LORA,lane);
        if(lane==0) ofull[out]=acc;
    }
}
// MoE topk: softmax router[64] then top-8 -> idxw[0..7]=idx, idxw[8..15]=w (norm*scaling)
__device__ inline void moe_topk(const float* rout,float* idxw_g,float* smem,int bid){
    int tid=threadIdx.x; float* sh=smem+KDA_SMEM; // [64] probs + [16] out
    float* pr=sh; float* out=sh+64;
    if(tid<64){ float mx=-1e30f; for(int j=0;j<64;j++) mx=fmaxf(mx,rout[j]);
        // softmax: need sum -- do in two passes within thread0? do per-thread reduce
        pr[tid]=__expf(rout[tid]-mx); }
    __syncthreads();
    if(tid==0){ float sm=0; for(int j=0;j<64;j++) sm+=pr[j];
        for(int j=0;j<64;j++) pr[j]=pr[j]/sm;
        // top-8
        for(int k=0;k<8;k++){ int bi=-1; float bv=-1e30f;
            for(int j=0;j<64;j++) if(pr[j]>bv){bv=pr[j];bi=j;}
            out[k]=(float)bi; out[8+k]=bv; pr[bi]=-1e30f; }
        float wsumw=0; for(int k=0;k<8;k++) wsumw+=out[8+k];
        for(int k=0;k<8;k++){ out[8+k]=out[8+k]/(wsumw+1e-9f)*ROUTED_SCALING;
            idxw_g[k]=out[k]; idxw_g[8+k]=out[8+k]; }
    }
    __syncthreads();
}
// routed gateup for 8 experts: warp-per-(j,m). hh[j*M+m]=silu(gate)*up
__device__ inline void moe_routed_gateup(const float* sx,const uint8_t* gwq,const bf16* gsc,const bf16* gzr,
        const uint8_t* uwq,const bf16* usc,const bf16* uzr,const float* idxw,float* hh,int bid){
    int tid=threadIdx.x,warp=tid>>5,lane=tid&31,ng=HID/GROUP,gw=bid*WARPS+warp,TW=gridDim.x*WARPS; int NA=8;
    int strideg = M*(HID/2); // per expert out-major
    for(int out=gw;out<NA*M;out+=TW){
        int j=out/M, m=out%M; int e=(int)idxw[j];
        float ga=gemv_dot_v(sx,gwq+(size_t)e*strideg+(size_t)m*(HID/2),gsc+(size_t)e*M*ng+(size_t)m*ng,gzr+(size_t)e*M*ng+(size_t)m*ng,HID,lane);
        float ua=gemv_dot_v(sx,uwq+(size_t)e*strideg+(size_t)m*(HID/2),usc+(size_t)e*M*ng+(size_t)m*ng,uzr+(size_t)e*M*ng+(size_t)m*ng,HID,lane);
        if(lane==0){float s=ga/(1.0f+__expf(-ga)); hh[j*M+m]=s*ua;}
    }
}
// down + combine + residual: warp-per-output c. hout[c]=hin[c]+sum_j w[j]*down[idxj][hh[j]] + shared_down[hh8]
__device__ inline void moe_down_resid(const uint8_t* dwq,const bf16* dsc,const bf16* dzr,
        const uint8_t* sdwq,const bf16* sdsc,const bf16* sdzr,const float* hh,const float* idxw,
        const float* hin,float* hout,int bid){
    int tid=threadIdx.x,warp=tid>>5,gw=bid*WARPS+warp,TW=gridDim.x*WARPS; int ng=M/GROUP;
    int stride_e=HID*(M/2);
    for(int c=gw;c<HID;c+=TW){
        float acc=0;
        for(int j=0;j<8;j++){ int e=(int)idxw[j]; float wj=idxw[8+j];
            const uint8_t* row=dwq+(size_t)e*stride_e+(size_t)c*(M/2);
            const bf16* sr=dsc+(size_t)e*HID*ng+(size_t)c*ng; const bf16* zr=dzr+(size_t)e*HID*ng+(size_t)c*ng;
            acc+=wj*dq_dot(hh+j*M,M,row,sr,zr);
        }
        { const uint8_t* row=sdwq+(size_t)c*(M/2);
          const bf16* sr=sdsc+(size_t)c*ng; const bf16* zr=sdzr+(size_t)c*ng;
          acc+=dq_dot(hh+8*M,M,row,sr,zr); }
        if((tid&31)==0) hout[c]=hin[c]+acc;
    }
}
"""


def _kda(b, hin, hout, SP, CQP, CKP, CVP):
    return f"""
    // ===== KDA block {b} (hin=P[{hin}] hout=P[{hout}]) =====
    load_fold(sx,(float*)P[{hin}],BF+BF_AN{b},HID,red);
    gemv_conv(sx,WQ+WQ_K{b}Q,SC+SC_K{b}Q,ZS+ZS_K{b}Q,HID,C,(float*)P[P_Q],bid,(bf16*)P[{CQP}],BF+BF_CONV{b}+0*C*4);
    __syncthreads();
    gemv_conv(sx,WQ+WQ_K{b}K,SC+SC_K{b}K,ZS+ZS_K{b}K,HID,C,(float*)P[P_K],bid,(bf16*)P[{CKP}],BF+BF_CONV{b}+1*C*4);
    __syncthreads();
    gemv_conv(sx,WQ+WQ_K{b}V,SC+SC_K{b}V,ZS+ZS_K{b}V,HID,C,(float*)P[P_V],bid,(bf16*)P[{CVP}],BF+BF_CONV{b}+2*C*4);
    __syncthreads();
    gemv_om(sx,WQ+WQ_K{b}G,SC+SC_K{b}G,ZS+ZS_K{b}G,HID,C,(float*)P[P_G],bid);
    __syncthreads();
    gemv_bf(sx,BF+BF_BETA{b},HID,H,(float*)P[P_BETA],bid);
    gbar(cnt,phase++,gen);
    kda_recur((float*)P[{SP}],(float*)P[P_Q],(float*)P[P_K],(float*)P[P_V],(float*)P[P_G],(float*)P[P_O],(float*)P[P_BETA],smem,red,bid);
    gbar(cnt,phase++,gen);
    load_fold(sx,(float*)P[P_O],(bf16*)0,C,red);
    gemv_res(sx,WQ+WQ_K{b}O,SC+SC_K{b}O,ZS+ZS_K{b}O,C,HID,(float*)P[{hout}],bid,(float*)P[{hin}]);
    gbar(cnt,phase++,gen);
    if(stop==100+{b}){{ float* hp=(float*)P[{hout}]; bf16* o2=(bf16*)P[P_OUT]; int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) o2[i]=f2b(hp[i]); return; }}
    load_fold(sx,(float*)P[{hout}],BF+BF_MN{b},HID,red);
    gemv_bf(sx,BF+BF_ROUT{b},HID,64,(float*)P[P_ROUT],bid);
    __syncthreads();
    gemv_gateup1(sx,WQ+WQ_B{b}SGATE,SC+SC_B{b}SGATE,ZS+ZS_B{b}SGATE,WQ+WQ_B{b}SUP,SC+SC_B{b}SUP,ZS+ZS_B{b}SUP,HID,M,(float*)P[P_HH]+8*M,bid);
    gbar(cnt,phase++,gen);
    moe_topk((float*)P[P_ROUT],(float*)P[P_IDXW],smem,bid);
    __syncthreads();
    moe_routed_gateup(sx,WQ+WQ_B{b}GATE,SC+SC_B{b}GATE,ZS+ZS_B{b}GATE,WQ+WQ_B{b}UP,SC+SC_B{b}UP,ZS+ZS_B{b}UP,(float*)P[P_IDXW],(float*)P[P_HH],bid);
    gbar(cnt,phase++,gen);
    moe_down_resid(WQ+WQ_B{b}DOWN,SC+SC_B{b}DOWN,ZS+ZS_B{b}DOWN,WQ+WQ_B{b}SDOWN,SC+SC_B{b}SDOWN,ZS+ZS_B{b}SDOWN,(float*)P[P_HH],(float*)P[P_IDXW],(float*)P[{hout}],(float*)P[{hout}],bid);
    gbar(cnt,phase++,gen);
    if(stop=={b+1}){{ float* hp=(float*)P[{hout}]; bf16* o2=(bf16*)P[P_OUT]; int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) o2[i]=f2b(hp[i]); return; }}
    """


def _mla(hin, hout):
    return f"""
    // ===== MLA block 3 (hin=P[{hin}] hout=P[{hout}]) =====
    load_fold(sx,(float*)P[{hin}],BF+BF_AN3,HID,red);
    gemv_om(sx,WQ+WQ_MQ,SC+SC_MQ,ZS+ZS_MQ,HID,H*(QK_NOPE+QK_ROPE),(float*)P[P_QFULL],bid);
    __syncthreads();
    gemv_om(sx,WQ+WQ_MVA,SC+SC_MVA,ZS+SC_MVA,HID,KV_LORA+QK_ROPE,(float*)P[P_KVFULL],bid);
    gbar(cnt,phase++,gen);
    mla_rope_append_qabs((float*)P[P_QFULL],(float*)P[P_KVFULL],(bf16*)P[P_CKV],(bf16*)P[P_KRC],(bf16*)P[P_CKVIN],(bf16*)P[P_KRCIN],L,prime,(float*)P[P_QABS],(uint8_t*)P[P_KVBUK_WQ],(bf16*)P[P_KVBUK_SC],(bf16*)P[P_KVBUK_ZR],bid);
    gbar(cnt,phase++,gen);
    mla_scores((float*)P[P_QFULL],(float*)P[P_QABS],(bf16*)P[P_CKV],(bf16*)P[P_KRC],(float*)P[P_SCORES],L,bid);
    gbar(cnt,phase++,gen);
    if(stop==31){{ int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) ((bf16*)P[P_OUT])[i]=f2b(((float*)P[P_H0])[i]); return; }}
    mla_softmax((float*)P[P_SCORES],(float*)P[P_P],L,bid,red);
    gbar(cnt,phase++,gen);
    if(stop==32){{ int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) ((bf16*)P[P_OUT])[i]=f2b(((float*)P[P_H0])[i]); return; }}
    {{ int t2=threadIdx.x; float* v=(float*)P[P_VACC]; for(int i=t2;i<H*KV_LORA;i+=THREADS) v[i]=0.0f; }}
    gbar(cnt,phase++,gen);
    mla_vaccum((float*)P[P_P],(bf16*)P[P_CKV],(float*)P[P_VACC],L,bid);
    gbar(cnt,phase++,gen);
    if(stop==33){{ int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) ((bf16*)P[P_OUT])[i]=f2b(((float*)P[P_H0])[i]); return; }}
    mla_o((float*)P[P_VACC],(uint8_t*)P[P_KVBUV_WQ],(bf16*)P[P_KVBUV_SC],(bf16*)P[P_KVBUV_ZR],(float*)P[P_OFULL],bid);
    gbar(cnt,phase++,gen);
    load_fold(sx,(float*)P[P_OFULL],(bf16*)0,H*VHD,red);
    gemv_res(sx,WQ+WQ_MO,SC+SC_MO,ZS+ZS_MO,H*VHD,HID,(float*)P[{hout}],bid,(float*)P[{hin}]);
    gbar(cnt,phase++,gen);
    load_fold(sx,(float*)P[{hout}],BF+BF_MN3,HID,red);
    gemv_bf(sx,BF+BF_ROUT3,HID,64,(float*)P[P_ROUT],bid);
    __syncthreads();
    gemv_gateup1(sx,WQ+WQ_B3SGATE,SC+SC_B3SGATE,ZS+ZS_B3SGATE,WQ+WQ_B3SUP,SC+SC_B3SUP,ZS+ZS_B3SUP,HID,M,(float*)P[P_HH]+8*M,bid);
    gbar(cnt,phase++,gen);
    moe_topk((float*)P[P_ROUT],(float*)P[P_IDXW],smem,bid);
    __syncthreads();
    moe_routed_gateup(sx,WQ+WQ_B3GATE,SC+SC_B3GATE,ZS+ZS_B3GATE,WQ+WQ_B3UP,SC+SC_B3UP,ZS+ZS_B3UP,(float*)P[P_IDXW],(float*)P[P_HH],bid);
    gbar(cnt,phase++,gen);
    moe_down_resid(WQ+WQ_B3DOWN,SC+SC_B3DOWN,ZS+ZS_B3DOWN,WQ+WQ_B3SDOWN,SC+SC_B3SDOWN,ZS+ZS_B3SDOWN,(float*)P[P_HH],(float*)P[P_IDXW],(float*)P[{hout}],(float*)P[{hout}],bid);
    gbar(cnt,phase++,gen);
    if(stop==4){{ float* hp=(float*)P[{hout}]; bf16* o2=(bf16*)P[P_OUT]; int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) o2[i]=f2b(hp[i]); return; }}
    """

# I spotted a typo (SC_MVA used as ZS). Fix:
_MLA_FIX = _mla("P_H1", "P_H0").replace("ZS+SC_MVA", "ZS+ZS_MVA")

_KERN = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cooperative_groups.h>
#include <ATen/cuda/CUDAContext.h>
namespace cg = cooperative_groups;
@CONST@
@DEFINES@
@BF_DEFINES@
enum { P_WQ,P_SC,P_ZS,P_BF, P_H0,P_H1, P_CNT, P_HIDDEN,P_OUT,
  P_Q,P_K,P_V,P_G,P_O,P_BETA,
  P_QFULL,P_KVFULL,P_OFULL,P_QABS,P_SCORES,P_P,P_VACC,
  P_HH,P_ROUT,P_IDXW,
  P_S0,P_CQ0,P_CK0,P_CV0, P_S1,P_CQ1,P_CK1,P_CV1, P_S2,P_CQ2,P_CK2,P_CV2,
  P_CKV,P_KRC,P_CKVIN,P_KRCIN,
  P_KVBUK_WQ,P_KVBUK_SC,P_KVBUK_ZR, P_KVBUV_WQ,P_KVBUV_SC,P_KVBUV_ZR, P_NPTRS };
@HELPERS@
extern "C" __global__ __launch_bounds__(THREADS, 4) void mega(void** P, int L, int gen, int prime, int stop){
    int bid=blockIdx.x;
    extern __shared__ float smem[];
    float* sx=smem; volatile float* red=smem+4096;
    uint8_t* WQ=(uint8_t*)P[P_WQ]; bf16* SC=(bf16*)P[P_SC]; bf16* ZS=(bf16*)P[P_ZS]; bf16* BF=(bf16*)P[P_BF];
    u32* cnt=(u32*)P[P_CNT]; int phase=0;
    { bf16* hidden=(bf16*)P[P_HIDDEN]; float* h0=(float*)P[P_H0]; int tid=threadIdx.x;
      for(int i=tid;i<HID;i+=THREADS) h0[i]=b2f(hidden[i]); }
    gbar(cnt,phase++,gen);
    if(stop==0){ float* hp=(float*)P[P_H0]; bf16* o2=(bf16*)P[P_OUT]; int t2=threadIdx.x; for(int i=t2;i<HID;i+=THREADS) o2[i]=f2b(hp[i]); return; }
    @BODY@
    { float* hp=(float*)P[P_H0]; bf16* out=(bf16*)P[P_OUT]; int tid=threadIdx.x;
      for(int i=tid;i<HID;i+=THREADS) out[i]=f2b(hp[i]); }
}
void launch_mega(torch::Tensor ptrs, int64_t L, int64_t gen, int64_t prime, int64_t stop){
    int smem=(4096+256+2048)*sizeof(float);
    int dev=0; cudaDeviceProp prop; cudaGetDeviceProperties(&prop,dev); int numSM=prop.multiProcessorCount;
    int mb=1; cudaOccupancyMaxActiveBlocksPerMultiprocessor(&mb,mega,THREADS,smem);
    int G=numSM*mb; if(G<1)G=1;
    void* p=ptrs.data_ptr<int64_t>();
    mega<<<G,THREADS,smem,at::cuda::getCurrentCUDAStream()>>>((void**)p,(int)L,(int)gen,(int)prime,(int)stop);
}
"""

_CUDA_SRC = (_KERN
    .replace("@CONST@", _CONST)
    .replace("@DEFINES@", _DEFINES)
    .replace("@BF_DEFINES@", _BF_DEFINES)
    .replace("@HELPERS@", _HELPERS)
    .replace("@BODY@", _kda(0, "P_H0", "P_H1", "P_S0", "P_CQ0", "P_CK0", "P_CV0")
              + _kda(1, "P_H1", "P_H0", "P_S1", "P_CQ1", "P_CK1", "P_CV1")
              + _kda(2, "P_H0", "P_H1", "P_S2", "P_CQ2", "P_CK2", "P_CV2")
              + _MLA_FIX))

_cpp = "void launch_mega(torch::Tensor ptrs, int64_t L, int64_t gen, int64_t prime, int64_t stop);"
_ext = load_inline("kimi_mega", cpp_sources=[_cpp], cuda_sources=[_CUDA_SRC],
                   functions=["launch_mega"],
                   extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"],
                   verbose=False)


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._prepared = False
        self._gen = 0

    def _prepare(self):
        dev = next(self.parameters()).device
        wq = torch.zeros(_WQSZ[0], dtype=torch.uint8, device=dev)
        sc = torch.zeros(_WQSZ[1], dtype=torch.bfloat16, device=dev)
        zs = torch.zeros(_WQSZ[2], dtype=torch.bfloat16, device=dev)
        bf = torch.zeros(_BFSZ, dtype=torch.bfloat16, device=dev)

        def put(name):
            wqt, sct, zrt = self._wt[name]
            if wqt.dim() == 2:
                WqT = wqt.t().contiguous(); ScT = sct.t().contiguous(); ZrT = zrt.t().contiguous()
            else:
                WqT = wqt.transpose(1, 2).contiguous()
                ScT = sct.transpose(1, 2).contiguous()
                ZrT = zrt.transpose(1, 2).contiguous()
            ow = _OFFWQ[name]; wq[ow:ow + WqT.numel()].copy_(WqT.view(-1))
            os = _OFFSC[name]; sc[os:os + ScT.numel()].copy_(ScT.view(-1))
            oz = _OFFZS[name]; zs[oz:oz + ZrT.numel()].copy_(ZrT.view(-1))

        for name, *_ in _QLIST:
            put(name)

        # bf16 params
        def putbf(name, tensor):
            o = _OFFBF[name]; n = tensor.numel()
            bf[o:o + n].copy_(tensor.reshape(-1).to(torch.bfloat16))
        for b in range(4):
            putbf(f"AN{b}", self.blocks[b].attn_norm); putbf(f"MN{b}", self.blocks[b].moe_norm)
        for b in range(3):
            putbf(f"BETA{b}", self.blocks[b].attn.beta_proj.weight)       # [H, HID] row-major
            putbf(f"CONV{b}", self.blocks[b].attn.conv_w)                 # [3, C, 4]
        for b in range(4):
            putbf(f"ROUT{b}", self.blocks[b].moe.router.weight)           # [E, HID]

        # kvb_UK (in-major [H,256,128]) and kvb_UV_T (out-major [H,128,256]) from MLA kv_b.
        # Vectorized gather (avoid per-column GPU syncs).
        kvb = self.blocks[3].attn.kv_b
        wqb = kvb.w_q      # [256, 8192]  (in//2, out)
        scb = kvb.scales   # [4, 8192]
        zrb = kvb.zeros
        hidx = torch.arange(H, device=dev)
        didx = torch.arange(VHD, device=dev)
        uk_cols = (hidx[:, None] * 256 + torch.arange(QK_NOPE, device=dev)[None, :]).reshape(-1)        # [H*128]
        uv_cols = (hidx[:, None] * 256 + 128 + didx[None, :]).reshape(-1)                                # [H*128]
        uk_wq = wqb.index_select(1, uk_cols).reshape(256, H, QK_NOPE).permute(1, 0, 2).contiguous()
        uk_sc = scb.index_select(1, uk_cols).reshape(4, H, QK_NOPE).permute(1, 0, 2).contiguous()
        uk_zr = zrb.index_select(1, uk_cols).reshape(4, H, QK_NOPE).permute(1, 0, 2).contiguous()
        uv_wq = wqb.index_select(1, uv_cols).reshape(256, H, VHD).permute(1, 2, 0).contiguous()
        uv_sc = scb.index_select(1, uv_cols).reshape(4, H, VHD).permute(1, 2, 0).contiguous()
        uv_zr = zrb.index_select(1, uv_cols).reshape(4, H, VHD).permute(1, 2, 0).contiguous()

        self._wq = wq; self._sc = sc; self._zs = zs; self._bf = bf
        self._uk_wq = uk_wq; self._uk_sc = uk_sc; self._uk_zr = uk_zr
        self._uv_wq = uv_wq; self._uv_sc = uv_sc; self._uv_zr = uv_zr

        # scratch (fp32 except state)
        self._h0 = torch.zeros(HID, device=dev, dtype=torch.float32)
        self._h1 = torch.zeros(HID, device=dev, dtype=torch.float32)
        self._out = torch.zeros(HID, device=dev, dtype=torch.bfloat16)
        self._cnt = torch.zeros(128, device=dev, dtype=torch.int32)
        self._q = torch.zeros(C, device=dev, dtype=torch.float32)
        self._k = torch.zeros(C, device=dev, dtype=torch.float32)
        self._v = torch.zeros(C, device=dev, dtype=torch.float32)
        self._g = torch.zeros(C, device=dev, dtype=torch.float32)
        self._o = torch.zeros(C, device=dev, dtype=torch.float32)
        self._beta = torch.zeros(H, device=dev, dtype=torch.float32)
        self._qfull = torch.zeros(H * (QK_NOPE + QK_ROPE), device=dev, dtype=torch.float32)
        self._kvfull = torch.zeros(KV_LORA + QK_ROPE, device=dev, dtype=torch.float32)
        self._ofull = torch.zeros(H * VHD, device=dev, dtype=torch.float32)
        self._qabs = torch.zeros(KV_LORA * H, device=dev, dtype=torch.float32)
        self._scores = torch.zeros((LMAX_ALLOC + 64) * H, device=dev, dtype=torch.float32)
        self._p = torch.zeros((LMAX_ALLOC + 64) * H, device=dev, dtype=torch.float32)
        self._vacc = torch.zeros(H * KV_LORA, device=dev, dtype=torch.float32)
        self._hh = torch.zeros(9 * M, device=dev, dtype=torch.float32)
        self._rout = torch.zeros(64, device=dev, dtype=torch.float32)
        self._idxw = torch.zeros(16, device=dev, dtype=torch.float32)
        self._ckv_buf = None  # MLA bigbuf, allocated on first step
        self._prepared = True

    def _build_wt(self):
        # build name -> (wq, sc, zr) accessor map
        wt = {}
        for b in range(3):
            a = self.blocks[b].attn
            wt[f"K{b}Q"] = (a.q_proj.w_q, a.q_proj.scales, a.q_proj.zeros)
            wt[f"K{b}K"] = (a.k_proj.w_q, a.k_proj.scales, a.k_proj.zeros)
            wt[f"K{b}V"] = (a.v_proj.w_q, a.v_proj.scales, a.v_proj.zeros)
            wt[f"K{b}G"] = (a.g_proj.w_q, a.g_proj.scales, a.g_proj.zeros)
            wt[f"K{b}O"] = (a.o_proj.w_q, a.o_proj.scales, a.o_proj.zeros)
        a = self.blocks[3].attn
        wt["MQ"] = (a.q_proj.w_q, a.q_proj.scales, a.q_proj.zeros)
        wt["MVA"] = (a.kv_a.w_q, a.kv_a.scales, a.kv_a.zeros)
        wt["MO"] = (a.o_proj.w_q, a.o_proj.scales, a.o_proj.zeros)
        for b in range(4):
            m = self.blocks[b].moe
            wt[f"B{b}GATE"] = (m.gate.w_q, m.gate.scales, m.gate.zeros)
            wt[f"B{b}UP"] = (m.up.w_q, m.up.scales, m.up.zeros)
            wt[f"B{b}DOWN"] = (m.down.w_q, m.down.scales, m.down.zeros)
            wt[f"B{b}SGATE"] = (m.s_gate.w_q, m.s_gate.scales, m.s_gate.zeros)
            wt[f"B{b}SUP"] = (m.s_up.w_q, m.s_up.scales, m.s_up.zeros)
            wt[f"B{b}SDOWN"] = (m.s_down.w_q, m.s_down.scales, m.s_down.zeros)
        self._wt = wt

    def step(self, hidden, state):
        if not self._prepared:
            self._build_wt(); self._prepare()
        dev = self._h0.device
        # MLA cache: detect whether the incoming cache is our own bigbuf (a
        # continued AR chain within one run) or a fresh state (new run/seed).
        # Only a fresh state needs the prime copy of the existing cache in.
        ckv_in = state[3]["c_kv"]      # [L, 512]
        krc_in = state[3]["k_rope"]    # [L, 64]
        L = ckv_in.shape[0]
        is_my_buf = (self._ckv_buf is not None
                     and ckv_in.data_ptr() == self._ckv_buf.data_ptr())
        if is_my_buf:
            prime = 0
        else:
            prime = 1
            cap = L + 64
            if self._ckv_buf is None or self._ckv_buf.shape[0] < cap:
                self._ckv_buf = torch.zeros(cap, KV_LORA, device=dev, dtype=torch.bfloat16)
                self._krc_buf = torch.zeros(cap, QK_ROPE, device=dev, dtype=torch.bfloat16)
        # pointer table in fixed order (must match enum P_*)
        ptrs = torch.tensor([
            self._wq.data_ptr(), self._sc.data_ptr(), self._zs.data_ptr(), self._bf.data_ptr(),
            self._h0.data_ptr(), self._h1.data_ptr(),
            self._cnt.data_ptr(),
            hidden.data_ptr(), self._out.data_ptr(),
            self._q.data_ptr(), self._k.data_ptr(), self._v.data_ptr(),
            self._g.data_ptr(), self._o.data_ptr(), self._beta.data_ptr(),
            self._qfull.data_ptr(), self._kvfull.data_ptr(), self._ofull.data_ptr(),
            self._qabs.data_ptr(), self._scores.data_ptr(), self._p.data_ptr(), self._vacc.data_ptr(),
            self._hh.data_ptr(), self._rout.data_ptr(), self._idxw.data_ptr(),
            state[0]["S"].data_ptr(), state[0]["cq"].data_ptr(), state[0]["ck"].data_ptr(), state[0]["cv"].data_ptr(),
            state[1]["S"].data_ptr(), state[1]["cq"].data_ptr(), state[1]["ck"].data_ptr(), state[1]["cv"].data_ptr(),
            state[2]["S"].data_ptr(), state[2]["cq"].data_ptr(), state[2]["ck"].data_ptr(), state[2]["cv"].data_ptr(),
            self._ckv_buf.data_ptr(), self._krc_buf.data_ptr(),
            ckv_in.data_ptr(), krc_in.data_ptr(),
            self._uk_wq.data_ptr(), self._uk_sc.data_ptr(), self._uk_zr.data_ptr(),
            self._uv_wq.data_ptr(), self._uv_sc.data_ptr(), self._uv_zr.data_ptr(),
        ], device=dev, dtype=torch.int64)
        _ext.launch_mega(ptrs, L, self._gen, 1 if prime else 0, getattr(self, "_stop", 4))
        self._gen += 1
        # reflect MLA cache into state (view of bigbuf)
        state[3]["c_kv"] = self._ckv_buf[:L + 1]
        state[3]["k_rope"] = self._krc_buf[:L + 1]
        return self._out, state

20260721_125907_zai-claude_glm-5.2_02_kimi_linear_decode