Download src/mla_npu.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 9.89 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/mla_npu.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/mla_npu.cpp
-
curl -L -o mla_npu.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/mla_npu.cpp
9.89 kB
| namespace ling3 { | |
| namespace { | |
| using Clock=std::chrono::steady_clock; | |
| constexpr int H=16,D=192,V=128,T=512; | |
| [[maybe_unused]] std::size_t Bucket(std::size_t rows) { | |
| for(std::size_t n:{16,32,64,128}) if(rows<=n)return n; | |
| throw std::invalid_argument("MLA NPU rows exceed 128"); | |
| } | |
| void Check(int rc,const char * op) { | |
| if(rc)throw std::runtime_error(std::string(op)+": "+std::to_string(rc)); | |
| } | |
| float Float(std::uint16_t x){return std::bit_cast<float>(std::uint32_t(x)<<16);} | |
| __fp16 Half(float x) { | |
| if(!std::isfinite(x)||std::abs(x)>65504.0F) | |
| throw std::runtime_error("MLA activation outside finite FP16 range"); | |
| return static_cast<__fp16>(x); | |
| } | |
| struct Matrix { | |
| rknn_matmul_ctx ctx=0; | |
| rknn_tensor_mem *a=nullptr,*b=nullptr,*c=nullptr; | |
| int m,k,n; | |
| Matrix(int rows,int inner,int cols,int core):m(rows),k(inner),n(cols) { | |
| try { | |
| rknn_matmul_info info{};rknn_matmul_io_attr io{}; | |
| info.M=m;info.K=k;info.N=n;info.type=RKNN_FLOAT16_MM_FLOAT16_TO_FLOAT32; | |
| info.B_layout=RKNN_MM_LAYOUT_NATIVE;info.AC_layout=RKNN_MM_LAYOUT_NATIVE; | |
| info.iommu_domain_id=2; | |
| Check(rknn_matmul_create(&ctx,&info,&io),"MLA create"); | |
| Check(rknn_matmul_set_core_mask(ctx,static_cast<rknn_core_mask>(1U<<core)),"MLA core mask"); | |
| if(io.A.size!=std::size_t(m)*k*2||io.B.size!=std::size_t(k)*n*2||io.C.size!=std::size_t(m)*n*4|| | |
| io.B.n_dims!=4||io.B.dims[2]!=16||io.B.dims[3]!=32) | |
| throw std::runtime_error("MLA unsupported native layout"); | |
| a=rknn_create_mem2(ctx,io.A.size,RKNN_FLAG_MEMORY_CACHEABLE); | |
| b=rknn_create_mem2(ctx,io.B.size,RKNN_FLAG_MEMORY_CACHEABLE); | |
| c=rknn_create_mem2(ctx,io.C.size,RKNN_FLAG_MEMORY_CACHEABLE); | |
| if(!a||!b||!c)throw std::runtime_error("MLA workspace allocation failed"); | |
| std::memset(a->virt_addr,0,a->size);std::memset(b->virt_addr,0,b->size); | |
| Check(rknn_matmul_set_io_mem(ctx,a,&io.A),"MLA bind A"); | |
| Check(rknn_matmul_set_io_mem(ctx,b,&io.B),"MLA bind B"); | |
| Check(rknn_matmul_set_io_mem(ctx,c,&io.C),"MLA bind C"); | |
| // Warm each native matmul without modifying decoder KV or recurrent state. | |
| Run(); | |
| for(int r=0;r<m;++r)for(int i=0;i<k;++i)A(r,i,0.125F); | |
| for(int i=0;i<k;++i)for(int col=0;col<n;++col)B(i,col,0.25F); | |
| Run(); | |
| for(int r=0;r<m;++r)for(int col=0;col<n;++col) | |
| if(!std::isfinite(C(r,col))||std::abs(C(r,col)-k/32.0F)>0.001F) | |
| throw std::runtime_error("MLA native FP16 capability probe failed"); | |
| } catch(...) {Release();throw;} | |
| } | |
| Matrix(const Matrix &)=delete; | |
| ~Matrix(){Release();} | |
| void Release(){if(c)rknn_destroy_mem(ctx,c);if(b)rknn_destroy_mem(ctx,b);if(a)rknn_destroy_mem(ctx,a);if(ctx)rknn_matmul_destroy(ctx);} | |
| void A(int row,int kk,float x){static_cast<__fp16 *>(a->virt_addr)[(kk/8*m+row)*8+kk%8]=Half(x);} | |
| void B(int kk,int col,float x){static_cast<__fp16 *>(b->virt_addr)[((col/16*(k/32)+kk/32)*16+col%16)*32+kk%32]=Half(x);} | |
| float C(int row,int col)const{return static_cast<const float *>(c->virt_addr)[(col/4*m+row)*4+col%4];} | |
| void Run(bool sync_b=true){ | |
| Check(rknn_mem_sync(ctx,a,RKNN_MEMORY_SYNC_TO_DEVICE),"MLA sync A"); | |
| if(sync_b)Check(rknn_mem_sync(ctx,b,RKNN_MEMORY_SYNC_TO_DEVICE),"MLA sync B"); | |
| Check(rknn_matmul_run(ctx),"MLA matmul"); | |
| Check(rknn_mem_sync(ctx,c,RKNN_MEMORY_SYNC_FROM_DEVICE),"MLA sync C"); | |
| } | |
| }; | |
| struct Lane { | |
| Matrix qk,av; | |
| std::vector<float> maximum,denominator,factor,accumulator,scores,av_high; | |
| Lane(int m,int core):qk(m,D,T,core),av(m,T,V,core),maximum(m),denominator(m),factor(m),accumulator(m*V),scores(m*T),av_high(m*V){} | |
| void Head(std::span<const float> q,std::span<const std::uint16_t> k,std::span<const std::uint16_t> v, | |
| int rows,std::size_t history,int h,std::span<float> out){ | |
| std::fill(maximum.begin(),maximum.end(),-std::numeric_limits<float>::infinity()); | |
| std::fill(denominator.begin(),denominator.end(),0);std::fill(accumulator.begin(),accumulator.end(),0); | |
| // Clear padded rows as workspaces are reused across partial requests. | |
| const float scale=1.0F/std::sqrt(float(D)); | |
| constexpr float residual_scale=4096.0F; | |
| for(std::size_t base=0;base<history;base+=T){ | |
| const int valid=std::min<std::size_t>(T,history-base); | |
| for(int j=0;j<T;++j){ | |
| for(int i=0;i<D;++i)qk.B(i,j,j<valid?Float(k[((base+j)*H+h)*D+i]):0); | |
| for(int i=0;i<V;++i)av.B(j,i,j<valid?Float(v[((base+j)*H+h)*V+i]):0); | |
| } | |
| // BF16 K/V are exact in FP16 over their representable range. Preserve | |
| // FP32 Q and probabilities with a scaled residual pass; plain FP16 Q/P | |
| // errors can otherwise amplify through quantization and MoE routing. | |
| for(int r=0;r<qk.m;++r)for(int i=0;i<D;++i)qk.A(r,i,r<rows?q[(r*H+h)*D+i]:0); | |
| qk.Run(); | |
| for(int r=0;r<rows;++r)for(int j=0;j<T;++j)scores[r*T+j]=qk.C(r,j); | |
| for(int r=0;r<qk.m;++r)for(int i=0;i<D;++i){ | |
| const float x=r<rows?q[(r*H+h)*D+i]:0; | |
| qk.A(r,i,(x-static_cast<float>(Half(x)))*residual_scale); | |
| } | |
| qk.Run(false); | |
| for(int r=0;r<rows;++r)for(int j=0;j<T;++j) | |
| scores[r*T+j]=(scores[r*T+j]+qk.C(r,j)/residual_scale)*scale; | |
| for(int r=0;r<qk.m;++r){ | |
| const auto causal_end=history-rows+std::min(r,rows-1)+1; | |
| const int count=r<rows&&causal_end>base?std::min<std::size_t>(valid,causal_end-base):0; | |
| if(!count){factor[r]=1;for(int j=0;j<T;++j){av.A(r,j,0);scores[r*T+j]=0;}continue;} | |
| float mx=maximum[r]; | |
| for(int j=0;j<count;++j)mx=std::max(mx,scores[r*T+j]); | |
| factor[r]=std::exp(maximum[r]-mx);float sum=0; | |
| for(int j=0;j<count;++j){const float p=std::exp(scores[r*T+j]-mx);sum+=p;av.A(r,j,p);scores[r*T+j]=p;} | |
| for(int j=count;j<T;++j){av.A(r,j,0);scores[r*T+j]=0;} | |
| denominator[r]=denominator[r]*factor[r]+sum;maximum[r]=mx; | |
| } | |
| av.Run(); | |
| for(int r=0;r<rows;++r)for(int i=0;i<V;++i)av_high[r*V+i]=av.C(r,i); | |
| for(int r=0;r<av.m;++r)for(int j=0;j<T;++j){ | |
| const float p=scores[r*T+j];av.A(r,j,(p-static_cast<float>(Half(p)))*residual_scale); | |
| } | |
| av.Run(false); | |
| for(int r=0;r<rows;++r)for(int i=0;i<V;++i) | |
| accumulator[r*V+i]=accumulator[r*V+i]*factor[r]+(av_high[r*V+i]+av.C(r,i)/residual_scale); | |
| } | |
| for(int r=0;r<rows;++r)for(int i=0;i<V;++i){ | |
| const float x=accumulator[r*V+i]/denominator[r]; | |
| if(!std::isfinite(x))throw std::runtime_error("non-finite MLA NPU output"); | |
| out[(r*H+h)*V+i]=x; | |
| } | |
| } | |
| }; | |
| struct Workspace { | |
| std::array<std::unique_ptr<Lane>,3> lanes; | |
| explicit Workspace(int rows){for(int i=0;i<3;++i)lanes[i]=std::make_unique<Lane>(rows,i);} | |
| }; | |
| } | |
| struct MlaNpu::Impl { | |
| std::string mode=std::getenv("LING3_MLA_BACKEND")?std::getenv("LING3_MLA_BACKEND"):"auto"; | |
| bool disabled=false; | |
| MlaBackendStats stats; | |
| std::map<std::size_t,std::unique_ptr<Workspace>> workspaces; | |
| Impl(){if(mode!="auto"&&mode!="cpu"&&mode!="npu")throw std::invalid_argument("MLA backend must be auto, cpu or npu");} | |
| void Failure(const std::exception & e){disabled=true;++stats.fallbacks;std::cerr<<"mla_warning: NPU attention unavailable, using CPU: "<<e.what()<<'\n';} | |
| void Prepare(std::size_t rows){ | |
| if(mode=="cpu"||disabled||rows<16)return; | |
| const auto bucket=Bucket(rows); | |
| try{if(!workspaces.contains(bucket))workspaces.emplace(bucket,std::make_unique<Workspace>(bucket));} | |
| catch(const std::exception & e){Failure(e);} | |
| disabled=true; | |
| } | |
| }; | |
| MlaNpu::MlaNpu():impl_(std::make_unique<Impl>()){} | |
| MlaNpu::~MlaNpu()=default; | |
| void MlaNpu::Prepare(std::size_t rows){impl_->Prepare(rows);} | |
| void MlaNpu::CpuCall(){++impl_->stats.cpu_calls;} | |
| MlaBackendStats MlaNpu::Stats()const{return impl_->stats;} | |
| bool MlaNpu::Run(std::span<const float> q,std::span<const std::uint16_t> k,std::span<const std::uint16_t> v, | |
| std::size_t rows,std::size_t history,std::span<float> output){ | |
| if(rows<1||rows>128||history<rows||q.size()!=rows*H*D||k.size()!=history*H*D||v.size()!=history*H*V||output.size()!=rows*H*V) | |
| throw std::invalid_argument("MLA backend shape mismatch"); | |
| if(impl_->mode=="cpu"||impl_->disabled||rows<16||(impl_->mode=="auto"&&history<256))return false; | |
| impl_->Prepare(rows); | |
| if(impl_->disabled)return false; | |
| auto & work=*impl_->workspaces.at(Bucket(rows));const auto start=Clock::now(); | |
| try{ | |
| constexpr std::array<int,3> cores{0,1,2}; | |
| CoreWorkers::Instance().Run(cores,[&](int c){for(int h=c;h<H;h+=3)work.lanes[c]->Head(q,k,v,rows,history,h,output);}); | |
| ++impl_->stats.npu_calls; | |
| impl_->stats.npu_ms+=std::chrono::duration<double,std::milli>(Clock::now()-start).count(); | |
| return true; | |
| }catch(const std::exception & e){impl_->Failure(e);} | |
| return false; | |
| } | |
| } | |