Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion Core_CPP/niyah_core.c
Original file line number Diff line number Diff line change
Expand Up @@ -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+15<n;i+=16){__m256 d0=_mm256_loadu_ps(dst+i),d1=_mm256_loadu_ps(dst+i+8);d0=_mm256_fmadd_ps(s,_mm256_loadu_ps(src+i),d0);d1=_mm256_fmadd_ps(s,_mm256_loadu_ps(src+i+8),d1);_mm256_storeu_ps(dst+i,d0);_mm256_storeu_ps(dst+i+8,d1);}for(;i<n;i++)dst[i]+=scale*src[i];
#elif defined(SIMD_NEON)
float32x4_t s=vdupq_n_f32(scale);size_t i=0;for(;i+7<n;i+=8){float32x4_t d0=vld1q_f32(dst+i),d1=vld1q_f32(dst+i+4);d0=vfmaq_f32(d0,s,vld1q_f32(src+i));d1=vfmaq_f32(d1,s,vld1q_f32(src+i+4));vst1q_f32(dst+i,d0);vst1q_f32(dst+i+4,d1);}for(;i<n;i++)dst[i]+=scale*src[i];
#else
for(size_t i=0;i<n;i++)dst[i]+=scale*src[i];
#endif
}

static void rmsnorm(float * restrict out,const float * restrict x,const float * restrict w,size_t n,float eps){
if (n == 0u) return;
const float ss = dot_f32(x, x, n);
Expand Down Expand Up @@ -119,7 +129,7 @@ int niyah_load(NiyahModel**out,const char*path){if(!out||!path)return -1;*out=NU
#define SCR_V(m)((m)->scratch+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;l<c->n_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;h<nh;h++)rope(q+h*hd,pos,hd,c->rope_theta);for(uint32_t h=0;h<nkv;h++)rope(k+h*hd,pos,hd,c->rope_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;h<nkv;h++){float*dk=kc+(size_t)h*xctx*hd+(size_t)pos*hd;float*dv=vc+(size_t)h*xctx*hd+(size_t)pos*hd;memcpy(dk,k+h*hd,hd*sizeof(float));memcpy(dv,v+h*hd,hd*sizeof(float));}memset(att,0,(size_t)nh*xctx*sizeof(float));for(uint32_t h=0;h<nh;h++){uint32_t kvh=(h*nkv)/nh;float*ah=att+(size_t)h*xctx;for(uint32_t t=0;t<=pos;t++){float score=dot_f32(q+h*hd,kc+(size_t)kvh*xctx*hd+(size_t)t*hd,hd)/(sqrtf((float)hd));if(score>80.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;j<hd;j++)xb2[h*hd+j]+=ah[t]*vv[j];}}matvec(xb,lw->wo,xb2,d,d);for(uint32_t i=0;i<d;i++)x[i]+=xb[i];rmsnorm(xb,x,lw->rms_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;i<m->ffn_dim;i++)hb[i]=silu(hb[i])*hb2[i];matvec(xb,lw->w_down,hb,d,m->ffn_dim);for(uint32_t i=0;i<d;i++)x[i]+=xb[i];}rmsnorm(xb,x,m->rms_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;l<c->n_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;h<nh;h++)rope(q+h*hd,pos,hd,c->rope_theta);for(uint32_t h=0;h<nkv;h++)rope(k+h*hd,pos,hd,c->rope_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;h<nkv;h++){float*dk=kc+(size_t)h*xctx*hd+(size_t)pos*hd;float*dv=vc+(size_t)h*xctx*hd+(size_t)pos*hd;memcpy(dk,k+h*hd,hd*sizeof(float));memcpy(dv,v+h*hd,hd*sizeof(float));}memset(att,0,(size_t)nh*xctx*sizeof(float));for(uint32_t h=0;h<nh;h++){uint32_t kvh=(h*nkv)/nh;float*ah=att+(size_t)h*xctx;for(uint32_t t=0;t<=pos;t++){float score=dot_f32(q+h*hd,kc+(size_t)kvh*xctx*hd+(size_t)t*hd,hd)/(sqrtf((float)hd));if(score>80.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;i<d;i++)x[i]+=xb[i];rmsnorm(xb,x,lw->rms_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;i<m->ffn_dim;i++)hb[i]=silu(hb[i])*hb2[i];matvec(xb,lw->w_down,hb,d,m->ffn_dim);for(uint32_t i=0;i<d;i++)x[i]+=xb[i];}rmsnorm(xb,x,m->rms_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;i<vocab_size;i++)if(logits[i]>logits[best])best=i;return best;}float mx=logits[0];for(uint32_t i=1;i<vocab_size;i++)if(logits[i]>mx)mx=logits[i];float sm=0.f;for(uint32_t i=0;i<vocab_size;i++){float z=(logits[i]-mx)/s->temperature; 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;i<vocab_size;i++){float z=(logits[i]-mx)/s->temperature;if(z<-80.f)z=-80.f;cum+=expf(z);if(cum>=target)return i;}return vocab_size-1u;}

Expand Down
Loading