#include "ling3/linear.h" #include "ling3/precision_policy.h" #include "ling3/quantization.h" #include "ling3/cpu_kernels.h" #include "core_workers.h" #include #include #include #include #include #include #include #include #include #include #include #if LING3_WITH_RKNN #include #include #include #endif namespace ling3 { namespace { using Clock=std::chrono::steady_clock; [[maybe_unused]] void Check(int rc,const char* op){if(rc)throw std::runtime_error(std::string(op)+": "+std::to_string(rc));} [[maybe_unused]] int Bucket(std::size_t rows){for(int b:{1,2,4,8,16,32,64,128})if(rows<=std::size_t(b))return b;throw std::invalid_argument("linear batch >128");} templatestd::span View(const TensorView&t,DataType dtype){ if(t.entry->dtype!=unsigned(dtype)||t.entry->data_bytes%sizeof(T))throw std::runtime_error("linear tensor dtype mismatch"); return {reinterpret_cast(t.data),std::size_t(t.entry->data_bytes/sizeof(T))}; } const std::set& CalibratedW4Families(){ static const auto families=[] { const char* value=std::getenv("LING3_CALIBRATED_W4_FAMILIES"); return ParseW8Families(value?value:""); }(); return families; } const ModelPackage& CalibratedW4Source(const ModelPackage& original,const std::string& base){ if(!SelectW8(CalibratedW4Families(),base))return original; static const ModelPackage source([]{ const char* path=std::getenv("LING3_CALIBRATED_W4_SOURCE"); if(!path||!*path)throw std::invalid_argument("LING3_CALIBRATED_W4_SOURCE is required"); return std::filesystem::path(path); }()); const auto&t=source.tensor(base+".weight");const auto&o=original.tensor(base+".weight"); if((source.header().flags&0x100)||(original.header().flags&0x100)|| t.entry->dtype!=unsigned(DataType::kInt4Low)||t.entry->rank!=2|| t.entry->dims[0]!=o.entry->dims[0]||t.entry->dims[1]!=o.entry->dims[1]|| t.entry->layout!=o.entry->layout||t.entry->quant!=o.entry->quant||t.entry->flags!=o.entry->flags) throw std::runtime_error("calibrated W4 source geometry/partition mismatch: "+base); std::cerr<<"precision_calibration family="<dtype!=unsigned(DataType::kBFloat16)||t.entry->rank!=2|| t.entry->dims[0]!=o.entry->dims[0]||t.entry->dims[1]!=o.entry->dims[1]) throw std::runtime_error("mixed BF16 source mismatch: "+base); std::cerr<<"precision_upgrade family="< old; bool w8=true; bool shared_stage=std::getenv("LING3_BRIDGE_SHARED_STAGE") && std::string(std::getenv("LING3_BRIDGE_SHARED_STAGE"))=="1"; Impl* input_owner=nullptr; #if LING3_WITH_RKNN struct Weight { rknn_matmul_ctx ctx=0;rknn_tensor_mem* b=nullptr; int offset=0,n=0,core=0; rknn_matmul_io_attr io{}; std::vector scales; ~Weight(){if(b)rknn_destroy_mem(ctx,b);if(ctx)rknn_matmul_destroy(ctx);} }; struct Work { rknn_matmul_ctx ctx=0;rknn_tensor_mem *a=nullptr,*c=nullptr; rknn_matmul_io_attr io{}; const rknn_tensor_mem* bound_b=nullptr; const rknn_tensor_mem* bound_a=nullptr; std::vector scales; std::vector q; ~Work(){if(c)rknn_destroy_mem(ctx,c);if(a)rknn_destroy_mem(ctx,a);if(ctx)rknn_matmul_destroy(ctx);} }; std::vector> weights; std::map>> workspaces; rknn_matmul_info Info(int m,int n)const{ rknn_matmul_info info{};info.M=m;info.K=config.k;info.N=n;info.iommu_domain_id=config.iommu_domain_id; info.type=w8?RKNN_INT8_MM_INT8_TO_INT32:RKNN_FLOAT16_MM_FLOAT16_TO_FLOAT32; info.B_layout=info.AC_layout=RKNN_MM_LAYOUT_NATIVE;return info; } void CheckLayout(const rknn_matmul_io_attr&io,int rows,int n)const{ if(io.B.n_dims!=4||io.B.dims[2]!=unsigned(w8?32:16)||io.B.dims[3]!=32|| io.A.size!=unsigned(rows*config.k*(w8?1:2))||io.B.size!=unsigned(config.k*n*(w8?1:2))||io.C.size!=unsigned(rows*n*4)) throw std::runtime_error("unsupported bridge native layout"); } void Prepare(std::size_t rows){ const int m=Bucket(rows); if(workspaces.contains(m))return; std::vector> workspace; for(const auto&weight:weights){ auto w=std::make_unique();auto info=Info(m,weight->n); Check(rknn_matmul_create(&w->ctx,&info,&w->io),"bridge create workspace"); CheckLayout(w->io,m,weight->n); Check(rknn_matmul_set_core_mask(w->ctx,static_cast(1<core)),"bridge workspace core"); w->a=rknn_create_mem2(w->ctx,w->io.A.size,RKNN_FLAG_MEMORY_CACHEABLE); w->c=rknn_create_mem2(w->ctx,w->io.C.size,RKNN_FLAG_MEMORY_CACHEABLE); if(!w->a||!w->c)throw std::runtime_error("bridge allocate workspace"); std::memset(w->a->virt_addr,0,w->io.A.size);w->scales.resize(m,1.f);w->q.resize(config.k); Check(rknn_matmul_set_io_mem(w->ctx,w->c,&w->io.C),"bridge bind C"); workspace.push_back(std::move(w)); } workspaces.emplace(m,std::move(workspace)); } void Stage(Work&w,int m,std::size_t rows,std::spaninput, std::span qinput={},std::span qscales={},std::span indices={}){ const int k=config.k,sub=w8?16:8; std::memset(w.a->virt_addr,0,w.io.A.size); for(std::size_t r=0;r(w.a->virt_addr)+target,q+j,16); else{ const float* source=input.data()+r*k+j; auto* dest=reinterpret_cast(w.a->virt_addr)+target; vst1q_f16(dest,vcombine_f16(vcvt_f16_f32(vld1q_f32(source)),vcvt_f16_f32(vld1q_f32(source+4)))); } } } Check(rknn_mem_sync(w.ctx,w.a,RKNN_MEMORY_SYNC_TO_DEVICE),"bridge sync A"); } void LaneRun(int lane,int m,std::size_t rows,const Impl& source,std::span output,Work* shared=nullptr){ auto&w=*workspaces.at(m)[lane];auto&sw=shared?*shared:w;auto&weight=*source.weights[lane]; if(w.bound_a!=sw.a){Check(rknn_matmul_set_io_mem(w.ctx,sw.a,&w.io.A),"bridge bind A");w.bound_a=sw.a;} if(w.bound_b!=weight.b){Check(rknn_matmul_set_io_mem(w.ctx,weight.b,&w.io.B),"bridge bind B");w.bound_b=weight.b;} Check(rknn_matmul_run(w.ctx),"bridge matmul"); Check(rknn_mem_sync(w.ctx,w.c,RKNN_MEMORY_SYNC_FROM_DEVICE),"bridge sync C"); for(std::size_t r=0;r(w.c->virt_addr)+idx)); value=vmulq_f32(vmulq_n_f32(value,sw.scales[r]),vld1q_f32(weight.scales.data()+col));} else value=vld1q_f32(static_cast(w.c->virt_addr)+idx); vst1q_f32(output.data()+r*config.n+weight.offset+col,value); } } W4RunTimings Batch(std::span input,std::size_t rows,const Impl&source,std::spanoutput, std::span qi={},std::span qs={},std::span idx={},bool prepared=false){ if(rows<1||rows>128||output.size()!=rows*config.n||config.k!=source.config.k||config.n!=source.config.n||w8!=source.w8||weights.size()!=source.weights.size()) throw std::invalid_argument("bridge batch shape mismatch"); if(!prepared && idx.empty() && input.size()!=rows*config.k)throw std::invalid_argument("bridge input size"); if(!idx.empty()){ if(!w8||idx.size()!=rows||qi.size()!=qs.size()*config.k)throw std::invalid_argument("bridge indexed input"); for(auto id:idx)if(id>=qs.size()||!std::isfinite(qs[id])||qs[id]<=0)throw std::invalid_argument("bridge input index"); } auto start=Clock::now();Prepare(rows);const int m=Bucket(rows); Work* staged=nullptr; if(shared_stage && !prepared && weights.size()>1){ staged=workspaces.at(m)[0].get(); Stage(*staged,m,rows,input,qi,qs,idx); } auto run=[&](int lane){ Work* shared=staged; if(prepared){auto* owner=input_owner?input_owner:this;shared=owner->workspaces.at(1)[0].get();} else if(!staged)Stage(*workspaces.at(m)[lane],m,rows,input,qi,qs,idx); LaneRun(lane,m,rows,source,output,shared); }; if(weights.size()==1)run(0); else CoreWorkers::Instance().Run(config.cores,[&](int core){auto it=std::find(config.cores.begin(),config.cores.end(),core);run(int(it-config.cores.begin()));}); W4RunTimings t;t.total_ms=std::chrono::duration(Clock::now()-start).count();return t; } #endif Impl(const ModelPackage& original,const std::string&base,W4LinearConfig c):config(std::move(c)){ const auto& package=PrecisionSource(original,base); const auto& tensor=package.tensor(base+".weight"); const bool mixed=package.header().flags & kPackageMixedW4W8; if(mixed && !std::getenv("LING3_BRIDGE_SHARED_STAGE"))shared_stage=true; if(mixed && std::getenv("LING3_OFFICIAL_EXECUTION") && std::string(std::getenv("LING3_OFFICIAL_EXECUTION"))!="w8") throw std::invalid_argument("self-contained mixed package requires W8 bridge execution"); if(!(package.header().flags & kPackageOfficialInt4) && (!mixed || tensor.entry->dtype==unsigned(DataType::kInt4Low))){ if(tensor.entry->quant!=unsigned(QuantType::kPerOutputChannel) || tensor.entry->layout!=unsigned(TensorLayout::kPackedInt4Low)) throw std::runtime_error("invalid per-channel W4 entry: "+base); old=std::make_unique(config,View(tensor,DataType::kInt4Low), View(package.tensor(base+".scales"),DataType::kFloat32),View(package.tensor(base+".correction"),DataType::kInt32)); }else{ #if !LING3_WITH_RKNN throw std::runtime_error("official weight bridge requires RKNN"); #else const std::string mode=std::getenv("LING3_OFFICIAL_EXECUTION")?std::getenv("LING3_OFFICIAL_EXECUTION"):"w8"; if(mode!="w8"&&mode!="fp16")throw std::invalid_argument("LING3_OFFICIAL_EXECUTION must be w8 or fp16"); w8=mode=="w8"; if(config.k%32||config.n%32||config.k<=0||config.n<=0||config.cores.empty())throw std::runtime_error("bridge weight shape"); const bool group=tensor.entry->dtype==unsigned(DataType::kInt4Low); std::span packed;std::span raw,groups; if(group){ if(tensor.entry->quant!=4||tensor.entry->flags!=32||tensor.entry->layout!=unsigned(TensorLayout::kPackedInt4Low))throw std::runtime_error("bridge group32 metadata"); packed=View(tensor,DataType::kInt4Low);groups=View(package.tensor(base+".scales"),DataType::kBFloat16); if(packed.size()!=size_t(config.k)*config.n/2||groups.size()!=size_t(config.n)*(config.k/32))throw std::runtime_error("bridge group32 shape"); }else{ raw=View(tensor,DataType::kBFloat16);if(raw.size()!=size_t(config.n)*config.k)throw std::runtime_error("bridge BF16 shape"); } std::vector row(config.k);int offset=0; for(size_t lane=0;lane();weight->offset=offset;weight->core=config.cores[lane]; weight->n=((config.n/32)/config.cores.size()+(lanen<32)throw std::runtime_error("bridge empty output slice"); auto info=Info(1,weight->n);Check(rknn_matmul_create(&weight->ctx,&info,&weight->io),"bridge create weight"); CheckLayout(weight->io,1,weight->n);weight->scales.resize(weight->n,1.f); weight->b=rknn_create_mem2(weight->ctx,weight->io.B.size,RKNN_FLAG_MEMORY_CACHEABLE); if(!weight->b)throw std::runtime_error("bridge allocate B"); const int subn=w8?32:16; for(int col=0;coln;++col){ const int global=offset+col;float maximum=0; for(int j=0;j65504.f)throw std::runtime_error("nonfinite/out-of-range bridge weight"); row[j]=v;maximum=std::max(maximum,std::abs(v)); } const float s=maximum>0?maximum/127:1;weight->scales[col]=s; for(int j=0;j(weight->b->virt_addr)[idx]=std::clamp(int(std::nearbyint(row[j]/s)),-127,127); else static_cast<__fp16*>(weight->b->virt_addr)[idx]=static_cast<__fp16>(row[j]); } } Check(rknn_mem_sync(weight->ctx,weight->b,RKNN_MEMORY_SYNC_TO_DEVICE),"bridge sync B"); offset+=weight->n;weights.push_back(std::move(weight)); } Prepare(1); #endif } if(!std::getenv("LING3_KEEP_SOURCE_WEIGHTS"))package.DiscardCopiedLinearWeight(tensor); } }; Linear::Linear(const ModelPackage&p,const std::string&b,W4LinearConfig c):impl_(std::make_unique(p,b,std::move(c))){} Linear::~Linear()=default; W4RunTimings Linear::Run(std::spanx,std::spany){if(impl_->old)return impl_->old->Run(x,y);return RunBatch(x,1,y);} W4RunTimings Linear::RunBatch(std::spanx,std::size_t rows,std::spany){return RunBatchWithWeights(x,rows,*this,y);} W4RunTimings Linear::RunBatchWithWeights(std::spanx,std::size_t rows,const Linear&weight,std::spany){ if(bool(impl_->old)!=bool(weight.impl_->old))throw std::invalid_argument("linear weight precision mismatch"); if(impl_->old)return impl_->old->RunBatchWithWeights(x,rows,*weight.impl_->old,y); #if LING3_WITH_RKNN return impl_->Batch(x,rows,*weight.impl_,y); #else throw std::runtime_error("RKNN disabled"); #endif } W4RunTimings Linear::RunBatchQuantizedRows(std::spanx,std::spans,std::spanrows,const Linear&weight,std::spany){ if(bool(impl_->old)!=bool(weight.impl_->old))throw std::invalid_argument("indexed weight precision mismatch"); if(impl_->old)return impl_->old->RunBatchQuantizedRows(x,s,rows,*weight.impl_->old,y); #if LING3_WITH_RKNN return impl_->Batch({},rows.size(),*weight.impl_,y,x,s,rows); #else throw std::runtime_error("RKNN disabled"); #endif } void Linear::PrepareBatch(std::size_t rows,bool indexed){if(impl_->old){impl_->old->PrepareBatch(rows,indexed);return;} #if LING3_WITH_RKNN impl_->Prepare(rows); #endif } float Linear::PrepareInput(std::spanx){if(impl_->old)return impl_->old->PrepareInput(x); #if LING3_WITH_RKNN if(x.size()!=size_t(impl_->config.k)||impl_->weights.size()!=1)throw std::invalid_argument("bridge prepared input"); auto&w=*impl_->workspaces.at(1)[0];impl_->Stage(w,1,1,x);return w.scales[0]; #else throw std::runtime_error("RKNN disabled"); #endif } W4RunTimings Linear::RunPrepared(float scale,std::spany){if(impl_->old)return impl_->old->RunPrepared(scale,y); #if LING3_WITH_RKNN return impl_->Batch({},1,*impl_,y,{},{},{},true); #else throw std::runtime_error("RKNN disabled"); #endif } void Linear::ShareInputFrom(Linear&owner){ if(bool(impl_->old)!=bool(owner.impl_->old))throw std::invalid_argument("shared input precision mismatch"); if(impl_->old){impl_->old->ShareInputFrom(*owner.impl_->old);return;} if(impl_->config.k!=owner.impl_->config.k||impl_->w8!=owner.impl_->w8)throw std::invalid_argument("bridge input sharing mismatch"); impl_->input_owner=owner.impl_.get(); } void Linear::SetSingleCore(int core){if(impl_->old){impl_->old->SetSingleCore(core);return;} #if LING3_WITH_RKNN if(impl_->weights.size()!=1||core<0||core>2)throw std::invalid_argument("bridge single core"); if(core==impl_->config.cores[0])return; for(auto&[rows,work]:impl_->workspaces)Check(rknn_matmul_set_core_mask(work[0]->ctx,static_cast(1<weights[0]->core=core;impl_->config.cores[0]=core; #endif } }