Ling-3.0-tiny-RKNN / src /linear.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
19.4 kB
#include "ling3/linear.h"
#include "ling3/precision_policy.h"
#include "ling3/quantization.h"
#include "ling3/cpu_kernels.h"
#include "core_workers.h"
#include <algorithm>
#include <array>
#include <chrono>
#include <cmath>
#include <cstdlib>
#include <cstring>
#include <map>
#include <set>
#include <sstream>
#include <iostream>
#include <stdexcept>
#if LING3_WITH_RKNN
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#include <arm_neon.h>
#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");}
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;
#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<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;
}
#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<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{
#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<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);
#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<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);
#if LING3_WITH_RKNN
return impl_->Batch(x,rows,*weight.impl_,y);
#else
throw std::runtime_error("RKNN disabled");
#endif
}
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);
#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::span<const float>x){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::span<float>y){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<rknn_core_mask>(1<<core)),"bridge core rebind");
impl_->weights[0]->core=core;impl_->config.cores[0]=core;
#endif
}
}