diff --git a/c/backend_cuda.cu b/c/backend_cuda.cu index 188330d..413e53b 100644 --- a/c/backend_cuda.cu +++ b/c/backend_cuda.cu @@ -8,12 +8,21 @@ #include #include +struct RaggedKVEntry { + const void *key; + const float *host_l,*host_r; + float *latent,*rope; + int length,capacity,K,R; +}; + struct ColiCudaTensor { void *weights; float *scales; size_t weight_bytes; int fmt, I, O, device; int tracked; + RaggedKVEntry ragged[512]; + int ragged_count; }; typedef struct { @@ -23,7 +32,7 @@ typedef struct { size_t x_cap, y_cap, gate_cap, up_cap; uint8_t *qx; float *qscale; size_t qx_cap, qscale_cap; - float *host_x,*host_y; size_t host_x_cap,host_y_cap; + float *host_x,*host_y,*host_kv; size_t host_x_cap,host_y_cap,host_kv_cap; float *aq,*al,*ar,*ac; size_t aq_cap,al_cap,ar_cap,ac_cap; float *pipe_buf[24]; size_t pipe_cap[24]; /* scratch persistenti del resident pipeline */ cudaStream_t stream; @@ -338,17 +347,17 @@ __global__ static void attention_absorb_batch_kernel(float *ctx,const float *q, ctx[((size_t)s*H+h)*V+v]=a*(fmt?wscale[row]:1.f);} } -/* Independent KV sequence per row. latent/rope are packed as [S,T,*], while - * lengths selects the valid prefix for each row. */ +/* Independent device-resident KV sequence per row. lengths selects the valid + * prefix; latent/rope point at paged caches updated by the host wrapper. */ __global__ static void attention_absorb_ragged_kernel(float *ctx,const float *q, - const float *latent,const float *rope,const int *lengths, + const float *const *latent,const float *const *rope,const int *lengths, const void *weights,const float *wscale,int fmt,int S,int H,int Q,int R, int V,int K,int T,float scale){ int s=blockIdx.y,h=blockIdx.x,tid=threadIdx.x,nt=lengths[s],rbase=h*(Q+V); if(s>=S||nt<1||nt>T)return; extern __shared__ float sm[];float *qa=sm,*cl=qa+K,*scores=cl+K,*red=scores+T; const float *qs=q+((size_t)s*H+h)*(Q+R); - const float *ls=latent+(size_t)s*T*K,*rs=rope+(size_t)s*T*R; + const float *ls=latent[s],*rs=rope[s]; for(int k=tid;k= bytes) return 1; if (*ptr) cudaFree(*ptr); @@ -440,16 +458,17 @@ extern "C" void coli_cuda_shutdown(void) { for(int b=0;b<24;b++) if(ctx->pipe_buf[b]) cudaFree(ctx->pipe_buf[b]); if (ctx->host_x) cudaFreeHost(ctx->host_x); if (ctx->host_y) cudaFreeHost(ctx->host_y); + if (ctx->host_kv) cudaFreeHost(ctx->host_kv); if (ctx->stream) cudaStreamDestroy(ctx->stream); if (ctx->group_desc) cudaFree(ctx->group_desc); ctx->x = ctx->y = ctx->gate = ctx->up = nullptr; ctx->qx=nullptr; ctx->qscale=nullptr; ctx->aq=ctx->al=ctx->ar=ctx->ac=nullptr; - ctx->host_x=ctx->host_y=nullptr;ctx->stream=nullptr; + ctx->host_x=ctx->host_y=ctx->host_kv=nullptr;ctx->stream=nullptr; ctx->x_cap = ctx->y_cap = ctx->gate_cap = ctx->up_cap = 0; ctx->qx_cap=ctx->qscale_cap=0; ctx->aq_cap=ctx->al_cap=ctx->ar_cap=ctx->ac_cap=0; - ctx->host_x_cap=ctx->host_y_cap=0; + ctx->host_x_cap=ctx->host_y_cap=ctx->host_kv_cap=0; ctx->group_desc=nullptr; ctx->group_desc_cap=0; } g_nctx = 0; @@ -805,35 +824,81 @@ extern "C" int coli_cuda_attention_project_batch(ColiCudaTensor *w,ColiCudaTenso } extern "C" int coli_cuda_attention_project_ragged(ColiCudaTensor *w,ColiCudaTensor *proj, - float *out,const float *q,const float *const *latent,const float *const *rope, + float *out,const float *q,const void *const *keys, + const float *const *latent,const float *const *rope, const int *lengths,int S,int H,int Q,int R,int V,int K,int T,float scale){ - if(!w||!proj||!out||!q||!latent||!rope||!lengths||S<1||S>512||T<1||T>512|| + if(!w||!proj||!out||!q||!keys||!latent||!rope||!lengths||S<1||S>512||T<1||T>8192|| H<1||Q<1||R<1||V<1||K<1||K>512||w->I!=K||w->O!=H*(Q+V)|| proj->device!=w->device||proj->I!=H*V)return 0; - size_t ln=(size_t)S*T*K,rn=(size_t)S*T*R; - float *lh=(float*)std::calloc(ln,sizeof(float)),*rh=(float*)std::calloc(rn,sizeof(float)); - if(!lh||!rh){std::free(lh);std::free(rh);return 0;} - for(int s=0;sT){std::free(lh);std::free(rh);return 0;} - std::memcpy(lh+(size_t)s*T*K,latent[s],(size_t)lengths[s]*K*sizeof(float)); - std::memcpy(rh+(size_t)s*T*R,rope[s],(size_t)lengths[s]*R*sizeof(float)); - } DeviceContext *dc=find_ctx(w->device); - if(!select_ctx(dc)){std::free(lh);std::free(rh);return 0;} - size_t qb=(size_t)S*H*(Q+R)*sizeof(float),lb=ln*sizeof(float),rb=rn*sizeof(float); + if(!select_ctx(dc))return 0; + float **dl=(float**)std::malloc((size_t)S*sizeof(*dl)); + float **dr=(float**)std::malloc((size_t)S*sizeof(*dr)); + int *old=(int*)std::malloc((size_t)S*sizeof(*old)); + int *add=(int*)std::malloc((size_t)S*sizeof(*add)); + int *off=(int*)std::malloc((size_t)S*sizeof(*off));int packed_n=0; + if(!dl||!dr||!old||!add||!off){std::free(dl);std::free(dr);std::free(old);std::free(add);std::free(off);return 0;} + for(int s=0;sT){std::free(dl);std::free(dr);std::free(old);std::free(add);std::free(off);return 0;} + RaggedKVEntry *e=nullptr; + for(int i=0;iragged_count;i++)if(w->ragged[i].key==keys[s]){e=&w->ragged[i];break;} + if(!e){ + if(w->ragged_count>=512){std::free(dl);std::free(dr);std::free(old);std::free(add);std::free(off);return 0;} + e=&w->ragged[w->ragged_count++];std::memset(e,0,sizeof(*e));e->key=keys[s]; + } + if(e->K!=K||e->R!=R||e->host_l!=latent[s]||e->host_r!=rope[s]||lengths[s]length){ + if(e->latent)cudaFree(e->latent);if(e->rope)cudaFree(e->rope); + e->latent=e->rope=nullptr;e->length=e->capacity=0; + e->K=K;e->R=R;e->host_l=latent[s];e->host_r=rope[s]; + } + if(lengths[s]>e->capacity){ + int cap=(lengths[s]+63)&~63;float *nl=nullptr,*nr=nullptr; + if(!cuda_ok(cudaMalloc(&nl,(size_t)cap*K*sizeof(float)),"ragged KV latent page")|| + !cuda_ok(cudaMalloc(&nr,(size_t)cap*R*sizeof(float)),"ragged KV rope page")){ + if(nl)cudaFree(nl);if(nr)cudaFree(nr);std::free(dl);std::free(dr);std::free(old);std::free(add);std::free(off);return 0; + } + if(e->length){ + cudaMemcpyAsync(nl,e->latent,(size_t)e->length*K*sizeof(float),cudaMemcpyDeviceToDevice,dc->stream); + cudaMemcpyAsync(nr,e->rope,(size_t)e->length*R*sizeof(float),cudaMemcpyDeviceToDevice,dc->stream); + } + if(e->latent)cudaFree(e->latent);if(e->rope)cudaFree(e->rope); + e->latent=nl;e->rope=nr;e->capacity=cap; + } + dl[s]=e->latent;dr[s]=e->rope;old[s]=e->length;add[s]=lengths[s]-e->length; + off[s]=packed_n;packed_n+=add[s]*(K+R); + } + size_t qb=(size_t)S*H*(Q+R)*sizeof(float); size_t cb=(size_t)S*H*V*sizeof(float),ob=(size_t)S*proj->O*sizeof(float); - int ok=reserve(&dc->aq,&dc->aq_cap,qb)&&reserve(&dc->al,&dc->al_cap,lb)&& - reserve(&dc->ar,&dc->ar_cap,rb)&&reserve(&dc->ac,&dc->ac_cap,cb)&& - reserve(&dc->y,&dc->y_cap,ob)&& - reserve_bytes(&dc->group_desc,&dc->group_desc_cap,(size_t)S*sizeof(int)); + size_t pb=(size_t)packed_n*sizeof(float); + size_t desc=(size_t)S*(2*sizeof(float*)+4*sizeof(int)); + int ok=reserve(&dc->aq,&dc->aq_cap,qb)&&reserve(&dc->ac,&dc->ac_cap,cb)&& + reserve(&dc->y,&dc->y_cap,ob)&&reserve_bytes(&dc->group_desc,&dc->group_desc_cap,desc)&& + (!pb||(reserve(&dc->al,&dc->al_cap,pb)&&reserve_pinned(&dc->host_kv,&dc->host_kv_cap,pb))); + char *db=(char*)dc->group_desc;float **ddl=(float**)db,**ddr=ddl+S; + int *dn=(int*)(ddr+S),*dold=dn+S,*dadd=dold+S,*doff=dadd+S; + if(ok&&pb){ + for(int s=0;shost_kv+off[s]; + std::memcpy(p,latent[s]+(size_t)old[s]*K,(size_t)add[s]*K*sizeof(float)); + std::memcpy(p+(size_t)add[s]*K,rope[s]+(size_t)old[s]*R,(size_t)add[s]*R*sizeof(float)); + } + ok=cuda_ok(cudaMemcpyAsync(dc->al,dc->host_kv,pb,cudaMemcpyHostToDevice,dc->stream),"ragged KV append upload"); + } if(ok)ok=cuda_ok(cudaMemcpyAsync(dc->aq,q,qb,cudaMemcpyHostToDevice,dc->stream),"ragged q upload")&& - cuda_ok(cudaMemcpyAsync(dc->al,lh,lb,cudaMemcpyHostToDevice,dc->stream),"ragged latent upload")&& - cuda_ok(cudaMemcpyAsync(dc->ar,rh,rb,cudaMemcpyHostToDevice,dc->stream),"ragged rope upload")&& - cuda_ok(cudaMemcpyAsync(dc->group_desc,lengths,(size_t)S*sizeof(int),cudaMemcpyHostToDevice,dc->stream),"ragged lengths upload"); - std::free(lh);std::free(rh);if(!ok)return 0; + cuda_ok(cudaMemcpyAsync(ddl,dl,(size_t)S*sizeof(float*),cudaMemcpyHostToDevice,dc->stream),"ragged latent pointers")&& + cuda_ok(cudaMemcpyAsync(ddr,dr,(size_t)S*sizeof(float*),cudaMemcpyHostToDevice,dc->stream),"ragged rope pointers")&& + cuda_ok(cudaMemcpyAsync(dn,lengths,(size_t)S*sizeof(int),cudaMemcpyHostToDevice,dc->stream),"ragged lengths upload")&& + cuda_ok(cudaMemcpyAsync(dold,old,(size_t)S*sizeof(int),cudaMemcpyHostToDevice,dc->stream),"ragged old lengths")&& + cuda_ok(cudaMemcpyAsync(dadd,add,(size_t)S*sizeof(int),cudaMemcpyHostToDevice,dc->stream),"ragged append lengths")&& + cuda_ok(cudaMemcpyAsync(doff,off,(size_t)S*sizeof(int),cudaMemcpyHostToDevice,dc->stream),"ragged append offsets"); + if(ok&&pb)ragged_kv_append<<stream>>>(ddl,ddr,dc->al,dold,dadd,doff,K,R); + if(ok)for(int s=0;sragged_count;i++)if(w->ragged[i].key==keys[s]){w->ragged[i].length=lengths[s];break;} + } + std::free(dl);std::free(dr);std::free(old);std::free(add);std::free(off);if(!ok)return 0; size_t shared=(size_t)(2*K+T+256)*sizeof(float); - attention_absorb_ragged_kernel<<stream>>>(dc->ac,dc->aq,dc->al,dc->ar, - (const int*)dc->group_desc,w->weights,w->scales,w->fmt,S,H,Q,R,V,K,T,scale); + attention_absorb_ragged_kernel<<stream>>>(dc->ac,dc->aq,ddl,ddr, + dn,w->weights,w->scales,w->fmt,S,H,Q,R,V,K,T,scale); quant_matmul<<O,S),256,0,dc->stream>>>(dc->y,dc->ac,proj->weights, proj->scales,proj->fmt,S,proj->I,proj->O,row_bytes(proj->fmt,proj->I)); return cuda_ok(cudaGetLastError(),"ragged attention launch")&& @@ -852,6 +917,10 @@ extern "C" void coli_cuda_tensor_free(ColiCudaTensor *tensor) { } if (tensor->weights) cudaFree(tensor->weights); if (tensor->scales) cudaFree(tensor->scales); + for(int i=0;iragged_count;i++){ + if(tensor->ragged[i].latent)cudaFree(tensor->ragged[i].latent); + if(tensor->ragged[i].rope)cudaFree(tensor->ragged[i].rope); + } std::free(tensor); } diff --git a/c/backend_cuda.h b/c/backend_cuda.h index edddbfe..63c1f11 100644 --- a/c/backend_cuda.h +++ b/c/backend_cuda.h @@ -94,7 +94,8 @@ COLI_CUDA_DLLEXPORT int coli_cuda_attention_project_batch(ColiCudaTensor *kv_b,C int V,int K,int T,float attention_scale); COLI_CUDA_DLLEXPORT int coli_cuda_attention_project_ragged(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, - float *out,const float *q,const float *const *latent,const float *const *rope, + float *out,const float *q,const void *const *keys, + const float *const *latent,const float *const *rope, const int *lengths,int S,int H,int Q,int R,int V,int K,int max_t,float attention_scale); COLI_CUDA_DLLEXPORT void coli_cuda_tensor_free(ColiCudaTensor *tensor); diff --git a/c/backend_loader.c b/c/backend_loader.c index bbcb8ca..eedbd50 100644 --- a/c/backend_loader.c +++ b/c/backend_loader.c @@ -62,7 +62,8 @@ typedef int (*fn_attention_absorb_batch_dev)(ColiCudaTensor *kv_b_shard,float *c typedef int (*fn_attention_absorb_kvdev)(ColiCudaTensor *kv_b,float *ctx,const float *q, const float *latent_dev,const float *rope_dev,int H,int Q,int R,int V,int K,int T, float scale); typedef int (*fn_attention_project_batch)(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, float *out,const float *q,const float *latent, const float *rope,int S,int H,int Q,int R, int V,int K,int T,float attention_scale); typedef int (*fn_attention_project_ragged)(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, - float *out,const float *q,const float *const *latent,const float *const *rope, + float *out,const float *q,const void *const *keys, + const float *const *latent,const float *const *rope, const int *lengths,int S,int H,int Q,int R,int V,int K,int max_t,float attention_scale); typedef int (*fn_attention_project_batch_dev)(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, float *out,const float *q_dev,const float *latent_dev,const float *rope_dev, int S,int H,int Q,int R,int V,int K,int T,float scale); typedef int (*fn_attention_project_batch_dev_out)(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, float *out_dev,const float *q_dev,const float *latent_dev,const float *rope_dev, int S,int H,int Q,int R,int V,int K,int T,float scale); @@ -348,10 +349,11 @@ int coli_cuda_attention_project_batch(ColiCudaTensor *kv_b,ColiCudaTensor *o_pro } int coli_cuda_attention_project_ragged(ColiCudaTensor *kv_b,ColiCudaTensor *o_proj, - float *out,const float *q,const float *const *latent,const float *const *rope, + float *out,const float *q,const void *const *keys, + const float *const *latent,const float *const *rope, const int *lengths,int S,int H,int Q,int R,int V,int K,int max_t,float attention_scale){ if(!coli_cuda_load()) return 0; - return g_cuda.attention_project_ragged(kv_b,o_proj,out,q,latent,rope,lengths, + return g_cuda.attention_project_ragged(kv_b,o_proj,out,q,keys,latent,rope,lengths, S,H,Q,R,V,K,max_t,attention_scale); } diff --git a/c/glm.c b/c/glm.c index 70da320..a45fbc9 100644 --- a/c/glm.c +++ b/c/glm.c @@ -2717,18 +2717,20 @@ static void attention_rows(Model *m, Layer *l, int layer, float *x, int S, int p !dnsel&&l->kv_b.cuda_eligible&&l->o.cuda_eligible&& qt_cuda_upload(&l->kv_b)&&qt_cuda_upload(&l->o)){ const float **rl=malloc((size_t)S*sizeof(*rl)),**rr=malloc((size_t)S*sizeof(*rr)); + const void **rk=malloc((size_t)S*sizeof(*rk)); int *rn=malloc((size_t)S*sizeof(*rn)); int mt=0; - if(rl&&rr&&rn){ + if(rk&&rl&&rr&&rn){ for(int s=0;skv_start[layer]; rn[s]=pos+1-st0; + rk[s]=kvs[s]; rl[s]=coli_kv_row(kvs[s]->Lc[layer],st0,kvl); rr[s]=coli_kv_row(kvs[s]->Rc[layer],st0,c->qk_rope); if(rn[s]>mt)mt=rn[s]; } cuda_core=cuda_projected=coli_cuda_attention_project_ragged(l->kv_b.cuda,l->o.cuda, - out,Q,rl,rr,rn,S,H,c->qk_nope,c->qk_rope,vh,kvl,mt,c->attn_scale); + out,Q,rk,rl,rr,rn,S,H,c->qk_nope,c->qk_rope,vh,kvl,mt,c->attn_scale); } - free(rl);free(rr);free(rn); + free(rk);free(rl);free(rr);free(rn); } else if(cuda_absorb&&l->n_kv_b_shard>1){ int n=l->n_kv_b_shard,st0=m->kv_start[layer],nt=pos_base+S-st0,ok=1; float *qs=falloc((int64_t)S*H*qh),*cs=falloc((int64_t)S*H*vh); @@ -5267,7 +5269,7 @@ static void run_serve_mux(Model *m, const char *snap){ g_draft=0; /* one scheduler owns every forward; MTP/speculation is not ragged-safe */ int maxctx=getenv("CTX")?atoi(getenv("CTX")):4096; int nctx=getenv("KV_SLOTS")?atoi(getenv("KV_SLOTS")):1; - if(nctx<1||nctx>16){fprintf(stderr,"KV_SLOTS deve essere tra 1 e 16\n");exit(2);} + if(nctx<1||nctx>512){fprintf(stderr,"KV_SLOTS must be between 1 and 512\n");exit(2);} g_kvsave=getenv("KVSAVE")?atoi(getenv("KVSAVE")):1; KVState *initial=m->kv; free(initial->kv_start); free(initial); ServeCtx *ctx=calloc(nctx,sizeof(*ctx)); ServeReq *req=calloc(nctx,sizeof(*req)); @@ -5323,7 +5325,7 @@ static void run_serve_mux(Model *m, const char *snap){ } active=0; for(int i=0;i1?atoi(argv[1]):64; int ebits= argc>2?atoi(argv[2]):8; int dbits= argc>3?atoi(argv[3]):ebits; - if(getenv("SERVE") && (kv_slot_count()<1 || kv_slot_count()>16)){ - fprintf(stderr,"KV_SLOTS must be between 1 and 16\n"); return 2; + int kv_limit=(getenv("SERVE_BATCH")&&atoi(getenv("SERVE_BATCH")))?512:16; + if(getenv("SERVE") && (kv_slot_count()<1 || kv_slot_count()>kv_limit)){ + fprintf(stderr,"KV_SLOTS must be between 1 and %d\n",kv_limit); return 2; } #ifdef COLI_CUDA if(getenv("COLI_CUDA") && atoi(getenv("COLI_CUDA"))){ diff --git a/c/tests/test_ragged_attention.cu b/c/tests/test_ragged_attention.cu index c00a8eb..dd1e2f8 100644 --- a/c/tests/test_ragged_attention.cu +++ b/c/tests/test_ragged_attention.cu @@ -16,14 +16,18 @@ int main(){ !coli_cuda_tensor_upload(&tp,p.data(),nullptr,0,D,O,dev))return 1; int n[S]={1,2,3};std::vector> l(S),r(S); const float *lp[S],*rp[S]; + const void *keys[S]; for(int s=0;s