From 22191440e70004d3a9eb5e92d30dbdb1e9c4c23d Mon Sep 17 00:00:00 2001 From: Doug Sharp Date: Fri, 24 Jul 2026 14:24:44 +0100 Subject: [PATCH] Metal fmt=4 (grouped-int4) decode: fused-attention kv_b + routed-expert GEMV - a_deqrow/a_qabs/a_ctx grouped-scale aware; AttnW.kvb_gs; attn guard allows kv_b.fmt==4 - moe_gemv fmt==4 branch + qgs; moe_submit / moe_block(_begin) qgs param + host gate; colibri.c threads expert gs - metal-test: grouped-int4 (g128/g64) arms for run_moe and run_attn Closes #585 --- c/backend_metal.h | 30 +++--- c/backend_metal.mm | 170 +++++++++++++++++++++++----------- c/colibri.c | 46 ++++----- c/tests/test_backend_metal.mm | 82 ++++++++++------ 4 files changed, 212 insertions(+), 116 deletions(-) diff --git a/c/backend_metal.h b/c/backend_metal.h index 35c4931c..7ea3962f 100644 --- a/c/backend_metal.h +++ b/c/backend_metal.h @@ -68,14 +68,14 @@ void coli_metal_unregister(void *base); */ int coli_metal_layer_decode(float *x, const float *in_ln, const float *post_ln, - const void *qa_w, const float *qa_s, int qa_fmt, const float *qa_ln, - const void *qb_w, const float *qb_s, int qb_fmt, - const void *kva_w, const float *kva_s, int kva_fmt, const float *kva_ln, - const void *kvb_w, const float *kvb_s, int kvb_fmt, - const void *o_w, const float *o_s, int o_fmt, - const void *shg_w, const float *shg_s, int shg_fmt, - const void *shu_w, const float *shu_s, int shu_fmt, - const void *shd_w, const float *shd_s, int shd_fmt, + const void *qa_w, const float *qa_s, int qa_fmt, int qa_gs, const float *qa_ln, + const void *qb_w, const float *qb_s, int qb_fmt, int qb_gs, + const void *kva_w, const float *kva_s, int kva_fmt, int kva_gs, const float *kva_ln, + const void *kvb_w, const float *kvb_s, int kvb_fmt, int kvb_gs, + const void *o_w, const float *o_s, int o_fmt, int o_gs, + const void *shg_w, const float *shg_s, int shg_fmt, int shg_gs, + const void *shu_w, const float *shu_s, int shu_fmt, int shu_gs, + const void *shd_w, const float *shd_s, int shd_fmt, int shd_gs, const float *router_w, const float *router_bias, int E, int K, int Ksel, float topp, int normk, float rscale, float *Lc, float *Rc, int S, int pos_base, int st0, @@ -104,11 +104,11 @@ int coli_metal_rtop8(int par, const float *sig, const float *bias, int S, int E, void coli_metal_attn_counts(uint64_t *ok, double *wall, double *kernel); void coli_metal_attn_lat(double *ksched, double *gsched); int coli_metal_attn_decode(const float *x, - const void *qa_w, const float *qa_s, int qa_fmt, const float *qa_ln, - const void *qb_w, const float *qb_s, int qb_fmt, - const void *kva_w, const float *kva_s, int kva_fmt, const float *kva_ln, - const void *kvb_w, const float *kvb_s, int kvb_fmt, - const void *o_w, const float *o_s, int o_fmt, + const void *qa_w, const float *qa_s, int qa_fmt, int qa_gs, const float *qa_ln, + const void *qb_w, const float *qb_s, int qb_fmt, int qb_gs, + const void *kva_w, const float *kva_s, int kva_fmt, int kva_gs, const float *kva_ln, + const void *kvb_w, const float *kvb_s, int kvb_fmt, int kvb_gs, + const void *o_w, const float *o_s, int o_fmt, int o_gs, float *Lc, float *Rc, int S, int pos_base, int st0, float eps, float theta, float ascale, float *out); /* Diagnostics: GPU blocks executed, CPU-fallback blocks, experts run on GPU. */ @@ -135,7 +135,7 @@ int coli_metal_resset_stats(double *flush_s); * out = [S, D] accumulate target * Returns 1 on success, 0 to signal the caller to fall back to the CPU path. */ -int coli_metal_moe_block(int nb, int D, int Iinter, int fmt, +int coli_metal_moe_block(int nb, int D, int Iinter, int fmt, int qgs, const void *const *g, const void *const *u, const void *const *d, const float *const *gs, const float *const *us, const float *const *ds, const float *xg, const int *xoff, const int *nr, @@ -150,7 +150,7 @@ int coli_metal_moe_block(int nb, int D, int Iinter, int fmt, * end returns 0 on GPU fault (caller redoes those experts on CPU). */ typedef struct ColiMetalMoeHandle ColiMetalMoeHandle; -ColiMetalMoeHandle* coli_metal_moe_block_begin(int nb, int D, int Iinter, int fmt, +ColiMetalMoeHandle* coli_metal_moe_block_begin(int nb, int D, int Iinter, int fmt, int qgs, const void *const *g, const void *const *u, const void *const *d, const float *const *gs, const float *const *us, const float *const *ds, const float *xg, const int *xoff, const int *nr, diff --git a/c/backend_metal.mm b/c/backend_metal.mm index 2894d791..d406c210 100644 --- a/c/backend_metal.mm +++ b/c/backend_metal.mm @@ -14,12 +14,12 @@ using namespace metal; kernel void mm_gemv(device const uchar* w [[buffer(0)]], // raw weight bytes - device const float* scale [[buffer(1)]], // [O] + device const float* scale [[buffer(1)]], // [O] (fmt=2) or [O*ng] (fmt=4, ng=(I+gs-1)/gs) device const float* x [[buffer(2)]], // [S,I] device float* y [[buffer(3)]], // [S,O] constant int& S [[buffer(4)]], constant int& I [[buffer(5)]], constant int& O [[buffer(6)]], constant int& fmt [[buffer(7)]], - constant int& NT [[buffer(8)]], + constant int& NT [[buffer(8)]], constant int& gs [[buffer(9)]], uint tg [[threadgroup_position_in_grid]], uint slane [[thread_index_in_simdgroup]], uint sgid [[simdgroup_index_in_threadgroup]]) { @@ -35,6 +35,8 @@ kernel void mm_gemv(device const uchar* w [[buffer(0)]], // raw weight by device const char4* w4 = (device const char4*)wr; for (int c = slane; c < I8; c += 32) acc += dot(float4(w4[2*c]),x4[2*c]) + dot(float4(w4[2*c+1]),x4[2*c+1]); for (int i = I8*8 + slane; i < I; i += 32) acc += float(wr[i]) * xr[i]; + acc = simd_sum(acc); + if (slane == 0) y[row] = acc * scale[o]; } else if (fmt == 2) { // int4 packed, rb=(I+1)/2 int rb = (I+1)/2; device const uchar* wr = w + (long)o * rb; @@ -47,20 +49,42 @@ kernel void mm_gemv(device const uchar* w [[buffer(0)]], // raw weight by for (int i = I8*8 + slane; i < I; i += 32) { uchar b = wr[i>>1]; int v = (i&1) ? (b>>4) : (b&0xF); acc += float(v-8) * xr[i]; } + acc = simd_sum(acc); + if (slane == 0) y[row] = acc * scale[o]; } else if (fmt == 3) { // int2 packed, rb=(I+3)/4 int rb = (I+3)/4; device const uchar* wr = w + (long)o * rb; for (int i = slane; i < I; i += 32) { uchar b = wr[i>>2]; int v = (b >> (2*(i&3))) & 0x3; acc += float(v-2) * xr[i]; } + acc = simd_sum(acc); + if (slane == 0) y[row] = acc * scale[o]; + } else if (fmt == 4) { // grouped int4, scale[o*ng + i/gs] + int rb = (I+1)/2, ng = (I+gs-1)/gs; + device const uchar* wr = w + (long)o * rb; + device const float* sr = scale + (long)o * ng; + device const uchar4* w4 = (device const uchar4*)wr; + for (int c = slane; c < I8; c += 32) { uchar4 b = w4[c]; + float4 w0 = float4(float(int(b.x&0xF)-8), float(int(b.x>>4)-8), float(int(b.y&0xF)-8), float(int(b.y>>4)-8)); + float4 w1 = float4(float(int(b.z&0xF)-8), float(int(b.z>>4)-8), float(int(b.w&0xF)-8), float(int(b.w>>4)-8)); + int g0 = (2*c*4+0)/gs, g1 = (2*c*4+1)/gs, g2 = (2*c*4+2)/gs, g3 = (2*c*4+3)/gs; + int g4 = (2*c*4+4)/gs, g5 = (2*c*4+5)/gs, g6 = (2*c*4+6)/gs, g7 = (2*c*4+7)/gs; + acc += dot(w0 * float4(sr[g0],sr[g1],sr[g2],sr[g3]), x4[2*c]) + + dot(w1 * float4(sr[g4],sr[g5],sr[g6],sr[g7]), x4[2*c+1]); + } + for (int i = I8*8 + slane; i < I; i += 32) { + uchar b = wr[i>>1]; int v = (i&1) ? (b>>4) : (b&0xF); acc += float(v-8) * xr[i] * sr[i/gs]; + } + acc = simd_sum(acc); + if (slane == 0) y[row] = acc; } else { // f32 device const float* wr = (device const float*)(w) + (long)o * I; device const float4* w4 = (device const float4*)wr; for (int c = slane; c < I8; c += 32) acc += dot(w4[2*c],x4[2*c]) + dot(w4[2*c+1],x4[2*c+1]); for (int i = I8*8 + slane; i < I; i += 32) acc += wr[i] * xr[i]; + acc = simd_sum(acc); + if (slane == 0) y[row] = acc * scale[o]; } - acc = simd_sum(acc); - if (slane == 0) y[row] = acc * scale[o]; } // Batched bindless expert GEMV: each row gr belongs to expert erow[gr], whose weight and @@ -72,7 +96,7 @@ kernel void moe_gemv(device const ulong* waddr [[buffer(0)]], device const ulong device float* yout [[buffer(4)]], constant int& O [[buffer(5)]], constant int& K [[buffer(6)]], constant int& Kin [[buffer(7)]], constant int& fmt [[buffer(8)]], - constant int& NT [[buffer(9)]], + constant int& NT [[buffer(9)]], constant int& qgs [[buffer(10)]], uint tg [[threadgroup_position_in_grid]], uint slane [[thread_index_in_simdgroup]], uint sgid [[simdgroup_index_in_threadgroup]]) { @@ -89,13 +113,25 @@ kernel void moe_gemv(device const ulong* waddr [[buffer(0)]], device const ulong float4 w1=float4(float(int(b.z&0xF)-8),float(int(b.z>>4)-8),float(int(b.w&0xF)-8),float(int(b.w>>4)-8)); acc+=dot(w0,x4[2*c])+dot(w1,x4[2*c+1]); } for(int i=K8*8+slane;i>1]; int v=(i&1)?(b>>4):(b&0xF); acc+=float(v-8)*xr[i]; } + } else if (fmt == 4) { // grouped int4: per-expert scale [O][ng], ng=ceil(K/qgs) + int rb=(K+1)/2, ng=(K+qgs-1)/qgs; device const uchar* w=(device const uchar*)(waddr[e])+(long)o*rb; + device const float* sr=sc+(long)o*ng; // grouped scales for this output row + device const uchar4* w4=(device const uchar4*)w; + for(int c=slane;c>4)-8),float(int(b.y&0xF)-8),float(int(b.y>>4)-8)); + float4 w1=float4(float(int(b.z&0xF)-8),float(int(b.z>>4)-8),float(int(b.w&0xF)-8),float(int(b.w>>4)-8)); + int g0=(8*c+0)/qgs,g1=(8*c+1)/qgs,g2=(8*c+2)/qgs,g3=(8*c+3)/qgs; + int g4=(8*c+4)/qgs,g5=(8*c+5)/qgs,g6=(8*c+6)/qgs,g7=(8*c+7)/qgs; + acc+=dot(w0*float4(sr[g0],sr[g1],sr[g2],sr[g3]),x4[2*c]) + +dot(w1*float4(sr[g4],sr[g5],sr[g6],sr[g7]),x4[2*c+1]); } + for(int i=K8*8+slane;i>1]; int v=(i&1)?(b>>4):(b&0xF); acc+=float(v-8)*xr[i]*sr[i/qgs]; } } else { device const char* w=(device const char*)(waddr[e])+(long)o*K; device const char4* w4=(device const char4*)w; for(int c=slane;c>1]; int val=(i&1)?(b>>4):(b&0xF); return float(val-8)*sc[row]; } +// kv_b inline dequant of column i of output row `row`. fmt=2 -> one scale per row; +// fmt=4 -> grouped int4, one scale per gs-wide group along the A_KVL input dim +// (scale layout [O][ng], ng=ceil(A_KVL/gs)), matching QT fmt=4 / mm_gemv above. +inline float a_deqrow(device const uchar* base, int row, int i, device const float* sc, int fmt, int gs){ + device const uchar* w=base+(long)row*((A_KVL+1)/2); uchar b=w[i>>1]; int val=(i&1)?(b>>4):(b&0xF); + float s = (fmt==4) ? sc[(long)row*((A_KVL+gs-1)/gs) + i/gs] : sc[row]; + return float(val-8)*s; } kernel void a_qabs(device const uchar* kvb [[buffer(0)]], device const float* sc [[buffer(1)]], device const float* q [[buffer(2)]], device float* qabs [[buffer(3)]], + constant int& fmt [[buffer(4)]], constant int& gs [[buffer(5)]], uint gid [[thread_position_in_grid]]) { int s=gid/(A_H*A_KVL), r=gid%(A_H*A_KVL), h=r/A_KVL, i=r%A_KVL; int rbase=h*A_ROWSH; device const float* qp=q+(long)s*A_QHH+(long)h*A_QH; - float a=0; for(int d=0;d e, const void* w, const float* s, int fmt, int I, int O, +static bool bind_gemv(id e, const void* w, const float* s, int fmt, int gs, int I, int O, id xin, id yout, int S){ uint64_t wa=0,sa=0; id wb=resolve(w,&wa); id sb=resolve(s,&sa); if(!wb||!sb) return false; @@ -616,7 +660,7 @@ static bool bind_gemv(id e, const void* w, const float [e setBuffer:xin offset:0 atIndex:2]; [e setBuffer:yout offset:0 atIndex:3]; int NT=S*O; [e setBytes:&S length:4 atIndex:4]; [e setBytes:&I length:4 atIndex:5]; [e setBytes:&O length:4 atIndex:6]; [e setBytes:&fmt length:4 atIndex:7]; - [e setBytes:&NT length:4 atIndex:8]; + [e setBytes:&NT length:4 atIndex:8]; [e setBytes:&gs length:4 atIndex:9]; [e dispatchThreadgroups:MTLSizeMake(((size_t)NT+3)/4,1,1) threadsPerThreadgroup:MTLSizeMake(128,1,1)]; return true; } @@ -624,20 +668,20 @@ static bool bind_gemv(id e, const void* w, const float // Weight-pointer bundle for one layer's attention (+optional layer tail). All pointers // must be inside registered allocations. typedef struct { - const void *qa_w; const float *qa_s; int qa_fmt; const float *qa_ln; - const void *qb_w; const float *qb_s; int qb_fmt; - const void *kva_w; const float *kva_s; int kva_fmt; const float *kva_ln; - const void *kvb_w; const float *kvb_s; int kvb_fmt; - const void *o_w; const float *o_s; int o_fmt; + const void *qa_w; const float *qa_s; int qa_fmt; int qa_gs; const float *qa_ln; + const void *qb_w; const float *qb_s; int qb_fmt; int qb_gs; + const void *kva_w; const float *kva_s; int kva_fmt; int kva_gs; const float *kva_ln; + const void *kvb_w; const float *kvb_s; int kvb_fmt; int kvb_gs; + const void *o_w; const float *o_s; int o_fmt; int o_gs; } AttnW; -// Encode the fused attention chain into encoder e. Input: ax_ holds the NORMED x [S,AH]. -// Output: aout_ holds attention output [S,AH]. Returns false on unresolved weights. -static bool encode_attention(id e, const AttnW *W, +// Phase 1: projections (qa, kva, qb, RMS, RoPE, qabs) for all S rows. +// Reads ax_[S*AH], writes aqr_[S*AQLORA], acomp_[S*(AKVL+AROPE)], aqf_[S*AHQH], aqabs_[S*AHEADS*AKVL]. +// Also writes Lc (keys) and Rc (rope keys) into the KV cache at pos_base. +static bool encode_attn_projections(id e, const AttnW *W, id Lb, size_t loff, id Rb, size_t roff, id kvbW, size_t kvbwoff, id kvbS, size_t kvbsoff, - int S, int pos_base, float eps, float theta, float ascale) { - int T=pos_base+S; + int S, int pos_base, float eps, float theta) { memcpy([aqaln_ contents],W->qa_ln,AQLORA*4); memcpy([akvaln_ contents],W->kva_ln,AKVL*4); size_t Loff=loff+(size_t)pos_base*AKVL*4, Roff=roff+(size_t)pos_base*AROPE*4; auto BAR=[&]{ [e memoryBarrierWithScope:MTLBarrierScopeBuffers]; }; @@ -651,13 +695,24 @@ static bool encode_attention(id e, const AttnW *W, [e setBuffer:acomp_ offset:0 atIndex:0]; [e setBytes:&off length:4 atIndex:1]; [e setBytes:&ss length:4 atIndex:2]; [e setBuffer:dst offset:doff atIndex:3]; [e setBytes:&n length:4 atIndex:4]; [e setBytes:&n length:4 atIndex:5]; [e dispatchThreads:MTLSizeMake((size_t)S*n,1,1) threadsPerThreadgroup:MTLSizeMake(64,1,1)]; }; - bind_gemv(e,W->qa_w,W->qa_s,W->qa_fmt,AH,AQLORA,ax_,aqr_,S); - bind_gemv(e,W->kva_w,W->kva_s,W->kva_fmt,AH,AKVL+AROPE,ax_,acomp_,S); BAR(); + bind_gemv(e,W->qa_w,W->qa_s,W->qa_fmt,W->qa_gs,AH,AQLORA,ax_,aqr_,S); + bind_gemv(e,W->kva_w,W->kva_s,W->kva_fmt,W->kva_gs,AH,AKVL+AROPE,ax_,acomp_,S); BAR(); rms(aqr_,0,aqaln_,AQLORA,S); cpy(0,Lb,Loff,AKVL); cpy(AKVL,Rb,Roff,AROPE); BAR(); - bind_gemv(e,W->qb_w,W->qb_s,W->qb_fmt,AQLORA,AHQH,aqr_,aqf_,S); rms(Lb,Loff,akvaln_,AKVL,S); rope(Rb,Roff,0,AROPE,0,1); BAR(); + bind_gemv(e,W->qb_w,W->qb_s,W->qb_fmt,W->qb_gs,AQLORA,AHQH,aqr_,aqf_,S); rms(Lb,Loff,akvaln_,AKVL,S); rope(Rb,Roff,0,AROPE,0,1); BAR(); rope(aqf_,0,ANOPE,AHQH,AQH,AHEADS); BAR(); [e setComputePipelineState:g_a_qabs]; [e setBuffer:kvbW offset:kvbwoff atIndex:0]; [e setBuffer:kvbS offset:kvbsoff atIndex:1]; [e setBuffer:aqf_ offset:0 atIndex:2]; [e setBuffer:aqabs_ offset:0 atIndex:3]; + [e setBytes:&W->kvb_fmt length:4 atIndex:4]; [e setBytes:&W->kvb_gs length:4 atIndex:5]; [e dispatchThreads:MTLSizeMake((size_t)S*AHEADS*AKVL,1,1) threadsPerThreadgroup:MTLSizeMake(256,1,1)]; BAR(); + return true; +} + +// Phase 2: attention core (score, softmax, context). Reads aqabs_, Lc, Rc, kvbW/S, writes aclat_, actx_. +static bool encode_attn_core(id e, const AttnW *W, + id Lb, size_t loff, id Rb, size_t roff, + id kvbW, size_t kvbwoff, id kvbS, size_t kvbsoff, + int S, int pos_base, float ascale) { + int T=pos_base+S; + auto BAR=[&]{ [e memoryBarrierWithScope:MTLBarrierScopeBuffers]; }; [e setComputePipelineState:g_a_score]; [e setBuffer:aqabs_ offset:0 atIndex:0]; [e setBuffer:Lb offset:loff atIndex:1]; [e setBuffer:Rb offset:roff atIndex:2]; [e setBuffer:aqf_ offset:0 atIndex:3]; [e setBuffer:ascore_ offset:0 atIndex:4]; [e setBytes:&T length:4 atIndex:5]; [e setBytes:&ascale length:4 atIndex:6]; [e setBytes:&pos_base length:4 atIndex:7]; [e dispatchThreads:MTLSizeMake((size_t)S*AHEADS*T,1,1) threadsPerThreadgroup:MTLSizeMake(256,1,1)]; BAR(); @@ -666,8 +721,19 @@ static bool encode_attention(id e, const AttnW *W, [e setComputePipelineState:g_a_clat]; [e setBuffer:ascore_ offset:0 atIndex:0]; [e setBuffer:Lb offset:loff atIndex:1]; [e setBuffer:aclat_ offset:0 atIndex:2]; [e setBytes:&T length:4 atIndex:3]; [e dispatchThreads:MTLSizeMake((size_t)S*AHEADS*AKVL,1,1) threadsPerThreadgroup:MTLSizeMake(256,1,1)]; BAR(); [e setComputePipelineState:g_a_ctx]; [e setBuffer:kvbW offset:kvbwoff atIndex:0]; [e setBuffer:kvbS offset:kvbsoff atIndex:1]; [e setBuffer:aclat_ offset:0 atIndex:2]; [e setBuffer:actx_ offset:0 atIndex:3]; + [e setBytes:&W->kvb_fmt length:4 atIndex:4]; [e setBytes:&W->kvb_gs length:4 atIndex:5]; [e dispatchThreads:MTLSizeMake((size_t)S*AHEADS*AVH,1,1) threadsPerThreadgroup:MTLSizeMake(256,1,1)]; BAR(); - bind_gemv(e,W->o_w,W->o_s,W->o_fmt,AHVH,AH,actx_,aout_,S); + return true; +} + +// Full fused encode: projections + core + o_proj. Returns false on unresolved weights. +static bool encode_attention(id e, const AttnW *W, + id Lb, size_t loff, id Rb, size_t roff, + id kvbW, size_t kvbwoff, id kvbS, size_t kvbsoff, + int S, int pos_base, float eps, float theta, float ascale) { + if(!encode_attn_projections(e,W,Lb,loff,Rb,roff,kvbW,kvbwoff,kvbS,kvbsoff,S,pos_base,eps,theta)) return false; + if(!encode_attn_core(e,W,Lb,loff,Rb,roff,kvbW,kvbwoff,kvbS,kvbsoff,S,pos_base,ascale)) return false; + bind_gemv(e,W->o_w,W->o_s,W->o_fmt,W->o_gs,AHVH,AH,actx_,aout_,S); return true; } // Resolve Lc/Rc + kv_b (+pre-check the projection weights). Returns false -> CPU fallback. @@ -685,18 +751,18 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, } extern "C" int coli_metal_attn_decode(const float* x, - const void* qa_w,const float* qa_s,int qa_fmt,const float* qa_ln, - const void* qb_w,const float* qb_s,int qb_fmt, - const void* kva_w,const float* kva_s,int kva_fmt,const float* kva_ln, - const void* kvb_w,const float* kvb_s,int kvb_fmt, - const void* o_w,const float* o_s,int o_fmt, + const void* qa_w,const float* qa_s,int qa_fmt,int qa_gs,const float* qa_ln, + const void* qb_w,const float* qb_s,int qb_fmt,int qb_gs, + const void* kva_w,const float* kva_s,int kva_fmt,int kva_gs,const float* kva_ln, + const void* kvb_w,const float* kvb_s,int kvb_fmt,int kvb_gs, + const void* o_w,const float* o_s,int o_fmt,int o_gs, float* Lc,float* Rc,int S,int pos_base,int st0,float eps,float theta,float ascale,float* out){ if(!g_dev) return 0; if(st0!=0 || S<1 || S>AMAXS) return 0; // partial-KV / S>4 -> CPU int T=pos_base+S; @autoreleasepool { attn_scratch_init(); - AttnW W={qa_w,qa_s,qa_fmt,qa_ln,qb_w,qb_s,qb_fmt,kva_w,kva_s,kva_fmt,kva_ln,kvb_w,kvb_s,kvb_fmt,o_w,o_s,o_fmt}; + AttnW W={qa_w,qa_s,qa_fmt,qa_gs,qa_ln,qb_w,qb_s,qb_fmt,qb_gs,kva_w,kva_s,kva_fmt,kva_gs,kva_ln,kvb_w,kvb_s,kvb_fmt,kvb_gs,o_w,o_s,o_fmt,o_gs}; id Lb,Rb,kvbW,kvbS; size_t loff,roff,kvbwoff,kvbsoff; if(!resolve_attn(&W,Lc,Rc,&Lb,&loff,&Rb,&roff,&kvbW,&kvbwoff,&kvbS,&kvbsoff)) return 0; ascore_=ensure(ascore_,&ascore_cap,(size_t)S*AHEADS*T*4); @@ -723,14 +789,14 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, // idx/w/keff (routing). Returns 0 -> CPU fallback (whole layer falls back). extern "C" int coli_metal_layer_decode(float *x, const float *in_ln, const float *post_ln, - const void* qa_w,const float* qa_s,int qa_fmt,const float* qa_ln, - const void* qb_w,const float* qb_s,int qb_fmt, - const void* kva_w,const float* kva_s,int kva_fmt,const float* kva_ln, - const void* kvb_w,const float* kvb_s,int kvb_fmt, - const void* o_w,const float* o_s,int o_fmt, - const void* shg_w,const float* shg_s,int shg_fmt, - const void* shu_w,const float* shu_s,int shu_fmt, - const void* shd_w,const float* shd_s,int shd_fmt, + const void* qa_w,const float* qa_s,int qa_fmt,int qa_gs,const float* qa_ln, + const void* qb_w,const float* qb_s,int qb_fmt,int qb_gs, + const void* kva_w,const float* kva_s,int kva_fmt,int kva_gs,const float* kva_ln, + const void* kvb_w,const float* kvb_s,int kvb_fmt,int kvb_gs, + const void* o_w,const float* o_s,int o_fmt,int o_gs, + const void* shg_w,const float* shg_s,int shg_fmt,int shg_gs, + const void* shu_w,const float* shu_s,int shu_fmt,int shu_gs, + const void* shd_w,const float* shd_s,int shd_fmt,int shd_gs, const float *router_w, const float *router_bias, int E, int K, int Ksel, float topp, int normk, float rscale, float *Lc, float *Rc, int S, int pos_base, int st0, @@ -741,7 +807,7 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, int T=pos_base+S; const int SI=2048; @autoreleasepool { attn_scratch_init(); - AttnW W={qa_w,qa_s,qa_fmt,qa_ln,qb_w,qb_s,qb_fmt,kva_w,kva_s,kva_fmt,kva_ln,kvb_w,kvb_s,kvb_fmt,o_w,o_s,o_fmt}; + AttnW W={qa_w,qa_s,qa_fmt,qa_gs,qa_ln,qb_w,qb_s,qb_fmt,qb_gs,kva_w,kva_s,kva_fmt,kva_gs,kva_ln,kvb_w,kvb_s,kvb_fmt,kvb_gs,o_w,o_s,o_fmt,o_gs}; id Lb,Rb,kvbW,kvbS; size_t loff,roff,kvbwoff,kvbsoff; if(!resolve_attn(&W,Lc,Rc,&Lb,&loff,&Rb,&roff,&kvbW,&kvbwoff,&kvbS,&kvbsoff)) return 0; uint64_t ina=0,pna=0,rwa=0,rba=0,d; @@ -778,8 +844,8 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, [e dispatchThreads:MTLSizeMake((size_t)S*AH,1,1) threadsPerThreadgroup:MTLSizeMake(256,1,1)]; BAR(); copyrow(axr_,anrm_,AH); BAR(); rmsw(anrm_,pnB,pnoff,AH,S); BAR(); // 4) shared expert gate/up + router (all read anrm_, independent) - bind_gemv(e,shg_w,shg_s,shg_fmt,AH,SI,anrm_,ash1_,S); - bind_gemv(e,shu_w,shu_s,shu_fmt,AH,SI,anrm_,ash2_,S); + bind_gemv(e,shg_w,shg_s,shg_fmt,shg_gs,AH,SI,anrm_,ash1_,S); + bind_gemv(e,shu_w,shu_s,shu_fmt,shu_gs,AH,SI,anrm_,ash2_,S); { int NT=S*E, D=AH; [e setComputePipelineState:g_r_router]; [e setBuffer:rwB offset:rwoff atIndex:0]; [e setBuffer:anrm_ offset:0 atIndex:1]; [e setBuffer:asig_ offset:0 atIndex:2]; [e setBytes:&E length:4 atIndex:3]; [e setBytes:&D length:4 atIndex:4]; [e setBytes:&NT length:4 atIndex:5]; @@ -806,7 +872,7 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, else [e dispatchThreads:MTLSizeMake(S,1,1) threadsPerThreadgroup:MTLSizeMake(S,1,1)]; } BAR(); // 6) shared down - bind_gemv(e,shd_w,shd_s,shd_fmt,SI,AH,ash1_,ashout_,S); + bind_gemv(e,shd_w,shd_s,shd_fmt,shd_gs,SI,AH,ash1_,ashout_,S); double tc=mnow(); [e endEncoding]; [cb commit]; [cb waitUntilCompleted]; if(cb.status==MTLCommandBufferStatusError){ fprintf(stderr,"[metal] layer cmdbuf error: %s\n", cb.error?[[cb.error localizedDescription]UTF8String]:"?"); return 0; } @@ -926,12 +992,12 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, // if Metal is off or any expert pointer is not in a registered slab. // Encode + commit a MoE block (no wait). Writes hh[R,D] into hh_buf. Returns nil on // unresolved slab / bad fmt (caller falls back to CPU). -static id moe_submit(int nb, int D, int Iinter, int fmt, +static id moe_submit(int nb, int D, int Iinter, int fmt, int qgs, const void *const *g, const void *const *u, const void *const *d, const float *const *gs, const float *const *us, const float *const *ds, const float *xg, const int *xoff, const int *nr, int R, id xg_buf, id gg_buf, id uu_buf, id hh_buf) { - if (!g_dev || (fmt != 1 && fmt != 2)) return nil; + if (!g_dev || (fmt != 1 && fmt != 2 && fmt != 4)) return nil; if (g_resset_enabled) { // E5: commit any pending slab adds before we may skip useResource: double t0 = mnow(); resset_flush(); g_t_resset_flush += mnow() - t0; // METAL-RESSET line } @@ -974,7 +1040,7 @@ static bool resolve_attn(const AttnW *W, float *Lc, float *Rc, [e setBuffer:wa offset:0 atIndex:0];[e setBuffer:sa offset:0 atIndex:1];[e setBuffer:berow offset:0 atIndex:2]; [e setBuffer:xin offset:0 atIndex:3];[e setBuffer:y offset:0 atIndex:4]; [e setBytes:&O length:4 atIndex:5];[e setBytes:&K length:4 atIndex:6];[e setBytes:&Kin length:4 atIndex:7];[e setBytes:&fmt length:4 atIndex:8]; - [e setBytes:&NT length:4 atIndex:9]; + [e setBytes:&NT length:4 atIndex:9];[e setBytes:&qgs length:4 atIndex:10]; [e dispatchThreadgroups:MTLSizeMake(((size_t)NT+3)/4,1,1) threadsPerThreadgroup:MTLSizeMake(128,1,1)]; }; gemv(bag,bsg,xg_buf,gg_buf,Iinter,D,D); // gate gemv(bau,bsu,xg_buf,uu_buf,Iinter,D,D); // up @@ -1009,7 +1075,7 @@ static int moe_finish(id cb, id hh_buf, int nb, int return 1; } -extern "C" int coli_metal_moe_block(int nb, int D, int Iinter, int fmt, +extern "C" int coli_metal_moe_block(int nb, int D, int Iinter, int fmt, int qgs, const void *const *g, const void *const *u, const void *const *d, const float *const *gs, const float *const *us, const float *const *ds, const float *xg, const int *xoff, const int *nr, @@ -1022,7 +1088,7 @@ static int moe_finish(id cb, id hh_buf, int nb, int g_gg = ensure(g_gg,&g_gg_cap,(size_t)R*Iinter*4); g_uu = ensure(g_uu,&g_uu_cap,(size_t)R*Iinter*4); g_hh = ensure(g_hh,&g_hh_cap,(size_t)R*D*4); - id cb = moe_submit(nb,D,Iinter,fmt,g,u,d,gs,us,ds,xg,xoff,nr,R,g_xg,g_gg,g_uu,g_hh); + id cb = moe_submit(nb,D,Iinter,fmt,qgs,g,u,d,gs,us,ds,xg,xoff,nr,R,g_xg,g_gg,g_uu,g_hh); if (!cb) return 0; return moe_finish(cb,g_hh,nb,R,D,rows,rw,out); } @@ -1035,7 +1101,7 @@ static int moe_finish(id cb, id hh_buf, int nb, int std::vector rows; std::vector rwv; int nb, R, D; }; -extern "C" ColiMetalMoeHandle* coli_metal_moe_block_begin(int nb, int D, int Iinter, int fmt, +extern "C" ColiMetalMoeHandle* coli_metal_moe_block_begin(int nb, int D, int Iinter, int fmt, int qgs, const void *const *g, const void *const *u, const void *const *d, const float *const *gs, const float *const *us, const float *const *ds, const float *xg, const int *xoff, const int *nr, @@ -1047,7 +1113,7 @@ static int moe_finish(id cb, id hh_buf, int nb, int id bgg=[g_dev newBufferWithLength:(size_t)R*Iinter*4 options:g_res_opts]; id buu=[g_dev newBufferWithLength:(size_t)R*Iinter*4 options:g_res_opts]; id bhh=[g_dev newBufferWithLength:(size_t)R*D*4 options:g_res_opts]; - id cb = moe_submit(nb,D,Iinter,fmt,g,u,d,gs,us,ds,xg,xoff,nr,R,bxg,bgg,buu,bhh); + id cb = moe_submit(nb,D,Iinter,fmt,qgs,g,u,d,gs,us,ds,xg,xoff,nr,R,bxg,bgg,buu,bhh); if (!cb) return nullptr; ColiMetalMoeHandle *h = new ColiMetalMoeHandle(); h->cb=cb; h->hh=bhh; h->rows.assign(rows,rows+R); h->rwv.assign(rw,rw+R); diff --git a/c/colibri.c b/c/colibri.c index bb5f7aae..e9ce3684 100644 --- a/c/colibri.c +++ b/c/colibri.c @@ -547,7 +547,7 @@ static void matmul_i4_grouped_pair(float *yg, float *yu, const float *x, * (~+12% perplexity), measured. Every other prefill matmul keeps IDOT as before. */ static void matmul_qt_ex(float *y, const float *x, QT *w, int S, int allow_idot){ #ifdef COLI_METAL - if(g_metal_enabled && S>=g_metal_gemm_min && !spec_pinned() && (w->fmt==1||w->fmt==2) && !omp_in_parallel()){ + if(g_metal_enabled && S>=g_metal_gemm_min && !spec_pinned() && (w->fmt==1||w->fmt==2||w->fmt==4) && !omp_in_parallel()){ const void *wp = w->fmt==1 ? (const void*)w->q8 : (const void*)w->q4; if(coli_metal_gemm(y,x,wp,w->s,w->fmt,S,w->I,w->O)) return; } @@ -1039,7 +1039,7 @@ static void qt_from_disk(Model *m, const char *name, int O, int I, int bits, int int fmt = qt_resolve_fmt(name,O,I,nb,ns,&gs); if(fmt==1){ if(t->fmt!=1||!t->q8){ t->fmt=1; t->O=O; t->I=I; t->gs=0; t->q8=qalloc(nb); t->s=qsalloc(O); } st_read_raw(&m->S,name,t->q8,drop); } else if(fmt==4){ int ng=(I+gs-1)/gs; - if(t->fmt!=4||!t->q4){ t->fmt=4; t->O=O; t->I=I; t->gs=gs; t->q4=qalloc(nb); t->s=falloc((int64_t)O*ng); } + if(t->fmt!=4||!t->q4){ t->fmt=4; t->O=O; t->I=I; t->gs=gs; t->q4=qalloc(nb); t->s=(float*)qalloc((size_t)O*ng*4); } st_read_raw(&m->S,name,t->q4,drop); } else if(fmt==5){ int64_t ng=i3_groups(I); /* int3-g64: 24B/group weights + O*ng group scales */ if(t->fmt!=5||!t->q4){ t->fmt=5; t->O=O; t->I=I; t->gs=0; t->q4=qalloc(nb); t->s=falloc((int64_t)O*ng); } @@ -2363,7 +2363,7 @@ static void attention_rows(Model *m, Layer *l, int layer, float *x, int S, int p * Ragged rows take the CPU absorb path below, which reads kvs[s]/positions[s]. */ if(g_metal_enabled && !kvs && S<=4 && (g_absorb==1||(g_absorb<0&&S<=4)) && m->kv_start[layer]==0 && D==6144 && H==64 && c->q_lora==2048 && c->kv_lora==512 && c->qk_nope==192 - && c->qk_rope==64 && vh==256 && l->kv_b.fmt==2){ + && c->qk_rope==64 && vh==256 && (l->kv_b.fmt==2||l->kv_b.fmt==4)){ int sel_active = m->has_dsa && layern_layers && c->idx_type[layer] && (pos_base+S) > c->index_topk; if(!sel_active){ if(m->has_dsa && layern_layers && c->idx_type[layer]){ /* index keys for future selection */ @@ -2374,11 +2374,11 @@ static void attention_rows(Model *m, Layer *l, int layer, float *x, int S, int p } #define WP_(q) ((q).fmt==1?(const void*)(q).q8:(const void*)(q).q4) int ok = coli_metal_attn_decode(x, - WP_(l->q_a), l->q_a.s, l->q_a.fmt, l->q_a_ln, - WP_(l->q_b), l->q_b.s, l->q_b.fmt, - WP_(l->kv_a), l->kv_a.s, l->kv_a.fmt, l->kv_a_ln, - WP_(l->kv_b), l->kv_b.s, l->kv_b.fmt, - WP_(l->o), l->o.s, l->o.fmt, + WP_(l->q_a), l->q_a.s, l->q_a.fmt, l->q_a.gs, l->q_a_ln, + WP_(l->q_b), l->q_b.s, l->q_b.fmt, l->q_b.gs, + WP_(l->kv_a), l->kv_a.s, l->kv_a.fmt, l->kv_a.gs, l->kv_a_ln, + WP_(l->kv_b), l->kv_b.s, l->kv_b.fmt, l->kv_b.gs, + WP_(l->o), l->o.s, l->o.fmt, l->o.gs, m->Lc[layer], m->Rc[layer], S, pos_base, m->kv_start[layer], c->eps, c->theta, c->attn_scale, out); #undef WP_ if(ok){ m->t_attn += now_s()-ta0; return; } @@ -3163,19 +3163,19 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out, int * preads run while the GPU computes; the missed subset follows in a second submit. * Per-subset CPU fallback on unresolved slab / bad fmt / GPU fault. */ int is_miss[64]={0}; ColiMetalMoeHandle *mh=NULL; - int cpu_res=1, cpu_miss=1, mh_shared=0, nbb=0, Rtot=0, mfmt=-1, sh_in=0; + int cpu_res=1, cpu_miss=1, mh_shared=0, nbb=0, Rtot=0, mfmt=-1, mgs=0, sh_in=0; const void *MG[65],*MU[65],*MD[65]; const float *MGS[65],*MUS[65],*MDS[65]; int xoffb[65],nrb[65]; float *mxg=NULL; int *mrows=NULL; float *mrw=NULL; /* subset builder: experts with is_miss==WANTMISS (+ shared expert when TRY_SH) */ #define MB_BUILD(WANTMISS, TRY_SH) do{ \ - nbb=0; Rtot=0; mfmt=-1; sh_in=0; \ + nbb=0; Rtot=0; mfmt=-1; mgs=0; sh_in=0; \ for(int j=0;jg.fmt; \ + if(mfmt<0){ mfmt=e->g.fmt; mgs=e->g.gs; } \ MG[nbb]=e->g.fmt==1?(const void*)e->g.q8:(const void*)e->g.q4; \ MU[nbb]=e->u.fmt==1?(const void*)e->u.q8:(const void*)e->u.q4; \ MD[nbb]=e->d.fmt==1?(const void*)e->d.q8:(const void*)e->d.q4; \ @@ -3184,7 +3184,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out, int } \ if(TRY_SH){ int shf = mfmt<0 ? l->sh_gate.fmt : mfmt; \ if(c->n_shared==1 && sI==I && l->sh_gate.fmt==shf && l->sh_up.fmt==shf && l->sh_down.fmt==shf){ \ - if(mfmt<0) mfmt=shf; \ + if(mfmt<0){ mfmt=shf; mgs=l->sh_gate.gs; } \ MG[nbb]=shf==1?(const void*)l->sh_gate.q8:(const void*)l->sh_gate.q4; \ MU[nbb]=shf==1?(const void*)l->sh_up.q8 :(const void*)l->sh_up.q4; \ MD[nbb]=shf==1?(const void*)l->sh_down.q8:(const void*)l->sh_down.q4; \ @@ -3207,7 +3207,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out, int MB_BUILD(0, base==0 && !g_pre_sh); if(nbb>0){ double t0=now_s(); - mh=coli_metal_moe_block_begin(nbb,D,I,mfmt,MG,MU,MD,MGS,MUS,MDS,mxg,xoffb,nrb,mrows,mrw); + mh=coli_metal_moe_block_begin(nbb,D,I,mfmt,mgs,MG,MU,MD,MGS,MUS,MDS,mxg,xoffb,nrb,mrows,mrw); m->t_emm += now_s()-t0; if(mh){ cpu_res=0; mh_shared=sh_in; } } else cpu_res=0; @@ -3266,7 +3266,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out, int MB_BUILD(1, 0); /* missed experts, now loaded */ if(nbb>0){ double t0=now_s(); - if(coli_metal_moe_block(nbb,D,I,mfmt,MG,MU,MD,MGS,MUS,MDS,mxg,xoffb,nrb,mrows,mrw,out,S)) cpu_miss=0; + if(coli_metal_moe_block(nbb,D,I,mfmt,mgs,MG,MU,MD,MGS,MUS,MDS,mxg,xoffb,nrb,mrows,mrw,out,S)) cpu_miss=0; m->t_emm += now_s()-t0; } else cpu_miss=0; if(mh){ double t0=now_s(); @@ -4195,7 +4195,7 @@ static void layer_forward_rows(Model *m, Layer *l, int li, float *x, int S, int if(g_metal_enabled && !kvs && S<=4 && lin_layers && l->sparse && (g_absorb==1||(g_absorb<0&&S<=4)) && m->kv_start[li]==0 && D==6144 && c->n_heads==64 && c->q_lora==2048 && c->kv_lora==512 - && c->qk_nope==192 && c->qk_rope==64 && c->v_head==256 && l->kv_b.fmt==2 + && c->qk_nope==192 && c->qk_rope==64 && c->v_head==256 && (l->kv_b.fmt==2||l->kv_b.fmt==4) && c->n_experts==256 && c->topk==8 && c->n_shared==1 && c->moe_inter==2048){ int sel_active = m->has_dsa && c->idx_type[li] && (pos_base+S) > c->index_topk; if(!sel_active){ @@ -4207,14 +4207,14 @@ static void layer_forward_rows(Model *m, Layer *l, int li, float *x, int S, int double ta0=now_s(); #define WP_(q) ((q).fmt==1?(const void*)(q).q8:(const void*)(q).q4) int ok = coli_metal_layer_decode(x, l->in_ln, l->post_ln, - WP_(l->q_a), l->q_a.s, l->q_a.fmt, l->q_a_ln, - WP_(l->q_b), l->q_b.s, l->q_b.fmt, - WP_(l->kv_a), l->kv_a.s, l->kv_a.fmt, l->kv_a_ln, - WP_(l->kv_b), l->kv_b.s, l->kv_b.fmt, - WP_(l->o), l->o.s, l->o.fmt, - WP_(l->sh_gate), l->sh_gate.s, l->sh_gate.fmt, - WP_(l->sh_up), l->sh_up.s, l->sh_up.fmt, - WP_(l->sh_down), l->sh_down.s, l->sh_down.fmt, + WP_(l->q_a), l->q_a.s, l->q_a.fmt, l->q_a.gs, l->q_a_ln, + WP_(l->q_b), l->q_b.s, l->q_b.fmt, l->q_b.gs, + WP_(l->kv_a), l->kv_a.s, l->kv_a.fmt, l->kv_a.gs, l->kv_a_ln, + WP_(l->kv_b), l->kv_b.s, l->kv_b.fmt, l->kv_b.gs, + WP_(l->o), l->o.s, l->o.fmt, l->o.gs, + WP_(l->sh_gate), l->sh_gate.s, l->sh_gate.fmt, l->sh_gate.gs, + WP_(l->sh_up), l->sh_up.s, l->sh_up.fmt, l->sh_up.gs, + WP_(l->sh_down), l->sh_down.s, l->sh_down.fmt, l->sh_down.gs, l->router, l->router_bias, c->n_experts, c->topk, Ksel, tp, c->norm_topk, c->routed_scale, m->Lc[li], m->Rc[li], S, pos_base, m->kv_start[li], diff --git a/c/tests/test_backend_metal.mm b/c/tests/test_backend_metal.mm index bff0444b..bf9110a5 100644 --- a/c/tests/test_backend_metal.mm +++ b/c/tests/test_backend_metal.mm @@ -56,42 +56,50 @@ static int run(int fmt, int O, int I, int S, const char *name) { static size_t roundpg(size_t n){ size_t p=16384; return ((n+p-1)/p)*p; } // Validate coli_metal_moe_block against a CPU reference (gate/up/silu/down + weighted scatter-add). -static int run_moe(const std::vector& nrv, const char* name) { - const int D=6144, I=2048, fmt=2; int rbG=(D+1)/2, rbD=(I+1)/2, nb=(int)nrv.size(); +// qgs==0 -> fmt=2 (per-row scale). qgs>0 -> fmt=4 grouped int4: per-expert scale slab is +// [O][ng] (ng=ceil(K/qgs)) for each of gate/up (K=D) and down (K=Iinter). +static int run_moe(const std::vector& nrv, int qgs, const char* name) { + const int D=6144, I=2048; int fmt = qgs>0 ? 4 : 2; + int rbG=(D+1)/2, rbD=(I+1)/2, nb=(int)nrv.size(); + int ngG = qgs>0 ? (D+qgs-1)/qgs : 1, ngD = qgs>0 ? (I+qgs-1)/qgs : 1; // scales/row for gate-up / down int R=0; std::vector xoff(nb),nr(nrv); for(int e=0;e0 ? s[(size_t)o*ngG + k/qgs] : s[o]; }; + auto scaD=[&](const float* s,int o,int k){ return qgs>0 ? s[(size_t)o*ngD + k/qgs] : s[o]; }; // per-expert page-aligned slab [Wg|Wu|Wd] and fslab [Sg|Su|Sd]; register both. std::vector slab(nb), fslab(nb); std::vector g(nb),u(nb),d(nb); std::vector gs(nb),us(nb),ds(nb); - size_t wlen=roundpg((size_t)I*rbG*2 + (size_t)D*rbD), flen=roundpg(((size_t)I*2+D)*sizeof(float)); + size_t nsc=(size_t)I*ngG*2 + (size_t)D*ngD; // gate + up + down scale counts + size_t wlen=roundpg((size_t)I*rbG*2 + (size_t)D*rbD), flen=roundpg(nsc*sizeof(float)); for(int e=0;e xg((size_t)R*D); for(auto&v:xg) v=((rand()%2000)-1000)/1000.f; std::vector rows(R); std::vector rw(R); for(int gr=0;gr position 0 int S=1; - // CPU reference + // CPU reference (grouped scale folded per-term; for fmt=2 that reduces to a*s[o]) std::vector refout((size_t)S*D,0.f), gg(I),uu(I),hh(D); for(int e=0;e gout((size_t)S*D,0.f); - int ok = coli_metal_moe_block(nb,D,I,fmt,g.data(),u.data(),d.data(),gs.data(),us.data(),ds.data(), + int ok = coli_metal_moe_block(nb,D,I,fmt,qgs,g.data(),u.data(),d.data(),gs.data(),us.data(),ds.data(), xg.data(),xoff.data(),nr.data(),rows.data(),rw.data(),gout.data(),S); double maxabs=0,ymax=0; for(size_t i=0;i& nrv, const char* name) { for(size_t i=0;i<(size_t)O*rb;i++) t.w[i]=(uint8_t)(rand()&0xFF); for(int i=0;i kv_b as fmt=2 (per-row scale); kvb_gs>0 -> kv_b as fmt=4 grouped int4. +static int run_attn(int S, int pos_base, int kvb_gs, const char* name){ const float eps=1e-5f, theta=10000.f, ascale=1.f/16.f; srand(4242+S+pos_base); - TW qa=t_mkw(TQL,TH), qb=t_mkw(THH*TQH,TQL), kva=t_mkw(TKVL+TROPE,TH), kvb=t_mkw(THH*TROWSH,TKVL), o=t_mkw(TH,THH*TVH); + int kvb_fmt = kvb_gs>0 ? 4 : 2, kvng = kvb_gs>0 ? (TKVL+kvb_gs-1)/kvb_gs : 1; + TW qa=t_mkw(TQL,TH), qb=t_mkw(THH*TQH,TQL), kva=t_mkw(TKVL+TROPE,TH); + TW kvb = kvb_gs>0 ? t_mkw_g(THH*TROWSH,TKVL,kvb_gs) : t_mkw(THH*TROWSH,TKVL); + TW o=t_mkw(TH,THH*TVH); + // per-column kv_b scale: grouped (fmt=4) picks scale[row*ng + i/gs], else per-row. + auto kvb_sc=[&](int row,int i)->float{ return kvb_gs>0 ? kvb.s[(size_t)row*kvng + i/kvb_gs] : kvb.s[row]; }; std::vector qaln(TQL), kvaln(TKVL); for(auto&v:qaln) v=0.5f+(rand()%1000)/1000.f; for(auto&v:kvaln) v=0.5f+(rand()%1000)/1000.f; int T=pos_base+S; size_t lcb=(((size_t)T*TKVL*4)+16383)&~(size_t)16383, rcb=(((size_t)T*TROPE*4)+16383)&~(size_t)16383; @@ -146,22 +168,22 @@ static int run_attn(int S, int pos_base, const char* name){ for(int h=0;h qabs(TKVL,0); - for(int d=0;d>1]; int v=(i&1)?(b>>4):(b&0xF); qabs[i]+=qp[d]*(float)(v-8)*sc; } } + for(int d=0;d>1]; int v=(i&1)?(b>>4):(b&0xF); qabs[i]+=qp[d]*(float)(v-8)*kvb_sc(rbase+d,i); } } std::vector a(pos+1); for(int t=0;t<=pos;t++){ const float*Lt=&Lr[(size_t)t*TKVL]; const float*Rt=&Rr[(size_t)t*TROPE]; float v=0; for(int i=0;i cl(TKVL,0); for(int t=0;t<=pos;t++){ const float*Lt=&Lr[(size_t)t*TKVL]; for(int i=0;i>1]; int vv=(i&1)?(b>>4):(b&0xF); v+=cl[i]*(float)(vv-8)*sc; } + for(int j=0;j>1]; int vv=(i&1)?(b>>4):(b&0xF); v+=cl[i]*(float)(vv-8)*kvb_sc(rbase+TNOPE+j,i); } ctx[(size_t)h*TVH+j]=v; } } t_gemv4(&ref[(size_t)s*TH],ctx.data(),o.w,o.s,TH,THH*TVH); } std::vector got((size_t)S*TH); - int ok=coli_metal_attn_decode(x.data(), qa.w,qa.s,2,qaln.data(), qb.w,qb.s,2, - kva.w,kva.s,2,kvaln.data(), kvb.w,kvb.s,2, o.w,o.s,2, + int ok=coli_metal_attn_decode(x.data(), qa.w,qa.s,2,0,qaln.data(), qb.w,qb.s,2,0, + kva.w,kva.s,2,0,kvaln.data(), kvb.w,kvb.s,kvb_fmt,kvb_gs, o.w,o.s,2,0, Lc,Rc,S,pos_base,0,eps,theta,ascale,got.data()); double ma=0,ym=0; for(size_t i=0;i