#include "mla_npu.h" #include "core_workers.h" #include #include #include #include #include #include #include #include #include #include #include #include #if LING3_WITH_RKNN && defined(__aarch64__) #include #include #endif 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"); } #if LING3_WITH_RKNN && defined(__aarch64__) 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(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(1U<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;r0.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(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 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 q,std::span k,std::span v, int rows,std::size_t history,int h,std::span out){ std::fill(maximum.begin(),maximum.end(),-std::numeric_limits::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(T,history-base); for(int j=0;j(Half(x)))*residual_scale); } qk.Run(false); for(int r=0;rbase?std::min(valid,causal_end-base):0; if(!count){factor[r]=1;for(int j=0;j(Half(p)))*residual_scale); } av.Run(false); for(int r=0;r,3> lanes; explicit Workspace(int rows){for(int i=0;i<3;++i)lanes[i]=std::make_unique(rows,i);} }; #endif } struct MlaNpu::Impl { std::string mode=std::getenv("LING3_MLA_BACKEND")?std::getenv("LING3_MLA_BACKEND"):"auto"; bool disabled=false; MlaBackendStats stats; #if LING3_WITH_RKNN && defined(__aarch64__) std::map> workspaces; #endif 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: "<(bucket));} catch(const std::exception & e){Failure(e);} #else disabled=true; #endif } }; MlaNpu::MlaNpu():impl_(std::make_unique()){} 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 q,std::span k,std::span v, std::size_t rows,std::size_t history,std::span output){ if(rows<1||rows>128||historymode=="cpu"||impl_->disabled||rows<16||(impl_->mode=="auto"&&history<256))return false; impl_->Prepare(rows); if(impl_->disabled)return false; #if LING3_WITH_RKNN && defined(__aarch64__) auto & work=*impl_->workspaces.at(Bucket(rows));const auto start=Clock::now(); try{ constexpr std::array cores{0,1,2}; CoreWorkers::Instance().Run(cores,[&](int c){for(int h=c;hHead(q,k,v,rows,history,h,output);}); ++impl_->stats.npu_calls; impl_->stats.npu_ms+=std::chrono::duration(Clock::now()-start).count(); return true; }catch(const std::exception & e){impl_->Failure(e);} #endif return false; } }