Download src/linear.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 19.4 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/linear.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/linear.cpp
-
curl -L -o linear.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/linear.cpp
19.4 kB
| 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");} | |
| template<class T>std::span<const T> 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<const T*>(t.data),std::size_t(t.entry->data_bytes/sizeof(T))}; | |
| } | |
| const std::set<std::string>& 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="<<W8MatrixFamily(base)<<" matrix="<<base<<" execution=W4A8\n"; | |
| return source; | |
| } | |
| // Process-scoped ablation policy. Only nonexpert BF16 source matrices may | |
| // replace the original W4 matrices; routed experts are never re-quantized here. | |
| const ModelPackage& PrecisionSource(const ModelPackage& original,const std::string& base){ | |
| static const auto families=[]{ | |
| const char* value=std::getenv("LING3_W8_FAMILIES"); | |
| return ParseW8Families(value?value:""); | |
| }(); | |
| if(original.header().flags & kPackageMixedW4W8){ | |
| if(!families.empty() || !CalibratedW4Families().empty()) | |
| throw std::invalid_argument("self-contained mixed package does not accept source overrides"); | |
| return original; | |
| } | |
| if(families.empty())return CalibratedW4Source(original,base); | |
| if(original.header().flags&0x100)throw std::invalid_argument("mixed ablation requires original W4 base package"); | |
| const char* execution=std::getenv("LING3_OFFICIAL_EXECUTION"); | |
| if(execution && std::string(execution)!="w8")throw std::invalid_argument("mixed ablation requires W8 execution"); | |
| if(!SelectW8(families,base))return CalibratedW4Source(original,base); | |
| if(SelectW8(CalibratedW4Families(),base))throw std::invalid_argument("conflicting W8 and calibrated W4 selectors"); | |
| const auto family=W8MatrixFamily(base); | |
| static const ModelPackage source([]{ | |
| const char* path=std::getenv("LING3_W8_SOURCE"); | |
| if(!path||!*path)throw std::invalid_argument("LING3_W8_SOURCE is required for mixed ablation"); | |
| return std::filesystem::path(path); | |
| }()); | |
| const auto&t=source.tensor(base+".weight");const auto&o=original.tensor(base+".weight"); | |
| if(!(source.header().flags&0x100)||t.entry->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="<<family<<" matrix="<<base<<" execution=W8A8 source=BF16\n"; | |
| return source; | |
| } | |
| } | |
| struct Linear::Impl { | |
| W4LinearConfig config; | |
| std::unique_ptr<DynamicW4Linear> 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; | |
| 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<float> 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<float> scales; | |
| std::vector<int8_t> q; | |
| ~Work(){if(c)rknn_destroy_mem(ctx,c);if(a)rknn_destroy_mem(ctx,a);if(ctx)rknn_matmul_destroy(ctx);} | |
| }; | |
| std::vector<std::unique_ptr<Weight>> weights; | |
| std::map<int,std::vector<std::unique_ptr<Work>>> 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<std::unique_ptr<Work>> workspace; | |
| for(const auto&weight:weights){ | |
| auto w=std::make_unique<Work>();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<rknn_core_mask>(1<<weight->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::span<const float>input, | |
| std::span<const int8_t> qinput={},std::span<const float> qscales={},std::span<const std::size_t> 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<rows;++r){ | |
| const int8_t* q=nullptr; | |
| if(w8){ | |
| if(!indices.empty()){q=qinput.data()+indices[r]*k;w.scales[r]=qscales[indices[r]];} | |
| else {w.scales[r]=QuantizeSymmetricInt8(input.subspan(r*k,k),w.q).scale;q=w.q.data();} | |
| } | |
| for(int j=0;j<k;j+=sub){ | |
| const size_t target=(size_t(j/sub)*m+r)*sub; | |
| if(w8)std::memcpy(static_cast<int8_t*>(w.a->virt_addr)+target,q+j,16); | |
| else{ | |
| const float* source=input.data()+r*k+j; | |
| auto* dest=reinterpret_cast<float16_t*>(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<float> 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<rows;++r)for(int col=0;col<weight.n;col+=4){ | |
| const size_t idx=(size_t(col/4)*m+r)*4; | |
| float32x4_t value; | |
| if(w8){value=vcvtq_f32_s32(vld1q_s32(static_cast<const int32_t*>(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<const float*>(w.c->virt_addr)+idx); | |
| vst1q_f32(output.data()+r*config.n+weight.offset+col,value); | |
| } | |
| } | |
| W4RunTimings Batch(std::span<const float> input,std::size_t rows,const Impl&source,std::span<float>output, | |
| std::span<const int8_t> qi={},std::span<const float> qs={},std::span<const std::size_t> 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<double,std::milli>(Clock::now()-start).count();return t; | |
| } | |
| 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<DynamicW4Linear>(config,View<std::byte>(tensor,DataType::kInt4Low), | |
| View<float>(package.tensor(base+".scales"),DataType::kFloat32),View<int32_t>(package.tensor(base+".correction"),DataType::kInt32)); | |
| }else{ | |
| throw std::runtime_error("official weight bridge requires RKNN"); | |
| 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<const std::byte> packed;std::span<const uint16_t> 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<std::byte>(tensor,DataType::kInt4Low);groups=View<uint16_t>(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<uint16_t>(tensor,DataType::kBFloat16);if(raw.size()!=size_t(config.n)*config.k)throw std::runtime_error("bridge BF16 shape"); | |
| } | |
| std::vector<float> row(config.k);int offset=0; | |
| for(size_t lane=0;lane<config.cores.size();++lane){ | |
| auto weight=std::make_unique<Weight>();weight->offset=offset;weight->core=config.cores[lane]; | |
| weight->n=((config.n/32)/config.cores.size()+(lane<size_t(config.n/32)%config.cores.size()))*32; | |
| if(weight->n<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;col<weight->n;++col){ | |
| const int global=offset+col;float maximum=0; | |
| for(int j=0;j<config.k;++j){ | |
| float v=group?float(DecodeInt4LowFirst(packed,size_t(j)*config.n+global))*BFloat16ToFloat(groups[size_t(global)*(config.k/32)+j/32]):BFloat16ToFloat(raw[size_t(global)*config.k+j]); | |
| if(!std::isfinite(v)||std::abs(v)>65504.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<config.k;++j){ | |
| const size_t idx=((size_t(col/subn)*(config.k/32)+j/32)*subn+col%subn)*32+j%32; | |
| if(w8)static_cast<int8_t*>(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); | |
| } | |
| 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<Impl>(p,b,std::move(c))){} | |
| Linear::~Linear()=default; | |
| W4RunTimings Linear::Run(std::span<const float>x,std::span<float>y){if(impl_->old)return impl_->old->Run(x,y);return RunBatch(x,1,y);} | |
| W4RunTimings Linear::RunBatch(std::span<const float>x,std::size_t rows,std::span<float>y){return RunBatchWithWeights(x,rows,*this,y);} | |
| W4RunTimings Linear::RunBatchWithWeights(std::span<const float>x,std::size_t rows,const Linear&weight,std::span<float>y){ | |
| 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); | |
| return impl_->Batch(x,rows,*weight.impl_,y); | |
| throw std::runtime_error("RKNN disabled"); | |
| } | |
| W4RunTimings Linear::RunBatchQuantizedRows(std::span<const int8_t>x,std::span<const float>s,std::span<const size_t>rows,const Linear&weight,std::span<float>y){ | |
| 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); | |
| return impl_->Batch({},rows.size(),*weight.impl_,y,x,s,rows); | |
| throw std::runtime_error("RKNN disabled"); | |
| } | |
| void Linear::PrepareBatch(std::size_t rows,bool indexed){if(impl_->old){impl_->old->PrepareBatch(rows,indexed);return;} | |
| impl_->Prepare(rows); | |
| } | |
| float Linear::PrepareInput(std::span<const float>x){if(impl_->old)return impl_->old->PrepareInput(x); | |
| 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]; | |
| throw std::runtime_error("RKNN disabled"); | |
| } | |
| W4RunTimings Linear::RunPrepared(float scale,std::span<float>y){if(impl_->old)return impl_->old->RunPrepared(scale,y); | |
| return impl_->Batch({},1,*impl_,y,{},{},{},true); | |
| throw std::runtime_error("RKNN disabled"); | |
| } | |
| 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(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<rknn_core_mask>(1<<core)),"bridge core rebind"); | |
| impl_->weights[0]->core=core;impl_->config.cores[0]=core; | |
| } | |
| } | |