diff --git a/Core_CPP/niyah_core.c b/Core_CPP/niyah_core.c index 5651edf..daca43e 100644 --- a/Core_CPP/niyah_core.c +++ b/Core_CPP/niyah_core.c @@ -84,6 +84,16 @@ static float dot_f32(const float * restrict a,const float * restrict b,size_t n) #endif } +static void axpy_f32(float * restrict dst,const float * restrict src,float scale,size_t n){ +#if defined(SIMD_AVX2) + __m256 s=_mm256_set1_ps(scale);size_t i=0;for(;i+15scratch+4u*(m)->cfg.embed_dim+2u*(m)->ffn_dim+(m)->kv_dim) #define SCR_ATT(m)((m)->scratch+4u*(m)->cfg.embed_dim+2u*(m)->ffn_dim+2u*(m)->kv_dim) -float*niyah_forward(NiyahModel*m,uint32_t token,uint32_t pos){if(!m||!m->_pool||token>=m->cfg.vocab_size||pos>=m->cfg.ctx_len)return NULL;const NiyahConfig*c=&m->cfg;uint32_t d=c->embed_dim,hd=m->head_dim,nh=c->n_heads,nkv=c->n_kv_heads,xctx=c->ctx_len;float*x=SCR_X(m),*xb=SCR_XB(m),*xb2=SCR_XB2(m),*hb=SCR_HB(m),*hb2=SCR_HB2(m),*q=SCR_Q(m),*k=SCR_K(m),*v=SCR_V(m),*att=SCR_ATT(m);memcpy(x,m->token_embed+(size_t)token*d,d*sizeof(float));for(uint32_t l=0;ln_layers;l++){const NiyahLayer*lw=&m->layers[l];rmsnorm(xb,x,lw->rms_att,d,c->rms_eps);matvec(q,lw->wq,xb,d,d);matvec(k,lw->wk,xb,m->kv_dim,d);matvec(v,lw->wv,xb,m->kv_dim,d);for(uint32_t h=0;hrope_theta);for(uint32_t h=0;hrope_theta);size_t L=(size_t)nkv*xctx*hd;float*kc=m->kv_k+(size_t)l*L,*vc=m->kv_v+(size_t)l*L;for(uint32_t h=0;h80.f)score=80.f;if(score<-80.f)score=-80.f;ah[t]=score;}float maxs=ah[0];for(uint32_t t=1;t<=pos;t++)if(ah[t]>maxs)maxs=ah[t];float den=0.f;for(uint32_t t=0;t<=pos;t++){ah[t]=expf(ah[t]-maxs);den+=ah[t];}den=den>0.f?den:1.f;for(uint32_t t=0;t<=pos;t++)ah[t]/=den;memset(xb2+h*hd,0,hd*sizeof(float));for(uint32_t t=0;t<=pos;t++){const float*vv=vc+(size_t)kvh*xctx*hd+(size_t)t*hd;for(uint32_t j=0;jwo,xb2,d,d);for(uint32_t i=0;irms_ffn,d,c->rms_eps);matvec(hb,lw->w_gate,xb,m->ffn_dim,d);matvec(hb2,lw->w_up,xb,m->ffn_dim,d);for(uint32_t i=0;iffn_dim;i++)hb[i]=silu(hb[i])*hb2[i];matvec(xb,lw->w_down,hb,d,m->ffn_dim);for(uint32_t i=0;irms_final,d,c->rms_eps);matvec(m->_logits,m->lm_head,xb,c->vocab_size,d);return m->_logits;} +float*niyah_forward(NiyahModel*m,uint32_t token,uint32_t pos){if(!m||!m->_pool||token>=m->cfg.vocab_size||pos>=m->cfg.ctx_len)return NULL;const NiyahConfig*c=&m->cfg;uint32_t d=c->embed_dim,hd=m->head_dim,nh=c->n_heads,nkv=c->n_kv_heads,xctx=c->ctx_len;float*x=SCR_X(m),*xb=SCR_XB(m),*xb2=SCR_XB2(m),*hb=SCR_HB(m),*hb2=SCR_HB2(m),*q=SCR_Q(m),*k=SCR_K(m),*v=SCR_V(m),*att=SCR_ATT(m);memcpy(x,m->token_embed+(size_t)token*d,d*sizeof(float));for(uint32_t l=0;ln_layers;l++){const NiyahLayer*lw=&m->layers[l];rmsnorm(xb,x,lw->rms_att,d,c->rms_eps);matvec(q,lw->wq,xb,d,d);matvec(k,lw->wk,xb,m->kv_dim,d);matvec(v,lw->wv,xb,m->kv_dim,d);for(uint32_t h=0;hrope_theta);for(uint32_t h=0;hrope_theta);size_t L=(size_t)nkv*xctx*hd;float*kc=m->kv_k+(size_t)l*L,*vc=m->kv_v+(size_t)l*L;for(uint32_t h=0;h80.f)score=80.f;if(score<-80.f)score=-80.f;ah[t]=score;}float maxs=ah[0];for(uint32_t t=1;t<=pos;t++)if(ah[t]>maxs)maxs=ah[t];float den=0.f;for(uint32_t t=0;t<=pos;t++){ah[t]=expf(ah[t]-maxs);den+=ah[t];}den=den>0.f?den:1.f;for(uint32_t t=0;t<=pos;t++)ah[t]/=den;memset(xb2+h*hd,0,hd*sizeof(float));for(uint32_t t=0;t<=pos;t++){const float*vv=vc+(size_t)kvh*xctx*hd+(size_t)t*hd;axpy_f32(xb2+h*hd,vv,ah[t],hd);}}matvec(xb,lw->wo,xb2,d,d);for(uint32_t i=0;irms_ffn,d,c->rms_eps);matvec(hb,lw->w_gate,xb,m->ffn_dim,d);matvec(hb2,lw->w_up,xb,m->ffn_dim,d);for(uint32_t i=0;iffn_dim;i++)hb[i]=silu(hb[i])*hb2[i];matvec(xb,lw->w_down,hb,d,m->ffn_dim);for(uint32_t i=0;irms_final,d,c->rms_eps);matvec(m->_logits,m->lm_head,xb,c->vocab_size,d);return m->_logits;} uint32_t niyah_sample(const float*logits,uint32_t vocab_size,NiyahSampler*s){if(!logits||!s||vocab_size==0u)return 0u;if(!(s->temperature>=0.f)||!isfinite(s->temperature)||!isfinite(s->top_p))return 0u;if(s->temperature<=0.f){uint32_t best=0u;for(uint32_t i=1;ilogits[best])best=i;return best;}float mx=logits[0];for(uint32_t i=1;imx)mx=logits[i];float sm=0.f;for(uint32_t i=0;itemperature; if(z<-80.f)z=-80.f;sm+=expf(z);}if(!(sm>0.f)||!isfinite(sm))return 0u;s->seed=s->seed*6364136223846793005ULL+1442695040888963407ULL;float r=(float)((s->seed>>11)&0x0FFFFFFU)/(float)0x0FFFFFFU;float top=(s->top_p>0.f&&s->top_p<1.f)?s->top_p:1.f;float target=r*sm*top;float cum=0.f;for(uint32_t i=0;itemperature;if(z<-80.f)z=-80.f;cum+=expf(z);if(cum>=target)return i;}return vocab_size-1u;}