Download tools/official_linear_bridge_probe.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 8.1 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_linear_bridge_probe.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/official_linear_bridge_probe.cpp
-
curl -L -o official_linear_bridge_probe.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_linear_bridge_probe.cpp
8.1 kB
| // Test FP16 and W8 execution bridges for released group32 weights. | |
| using Clock=std::chrono::steady_clock; | |
| void Check(int rc,const char* where){if(rc)throw std::runtime_error(std::string(where)+": "+std::to_string(rc));} | |
| template<class T> std::vector<T> Read(const std::string& path,size_t count){ | |
| std::ifstream f(path,std::ios::binary|std::ios::ate); | |
| if(!f||f.tellg()!=std::streamoff(count*sizeof(T)))throw std::runtime_error("bad input "+path); | |
| f.seekg(0);std::vector<T> x(count); | |
| if(!f.read(reinterpret_cast<char*>(x.data()),count*sizeof(T)))throw std::runtime_error("read "+path); | |
| return x; | |
| } | |
| struct Matrix { | |
| int m,k,n,offset;bool w8; | |
| rknn_matmul_ctx ctx=0;rknn_matmul_io_attr io{}; | |
| rknn_tensor_mem *a=nullptr,*b=nullptr,*c=nullptr; | |
| std::vector<float> scales,as; | |
| Matrix(int rows,int inner,int cols,int begin,int core,bool quant,std::span<const float> weight):m(rows),k(inner),n(cols),offset(begin),w8(quant),scales(n,1.f),as(m,1.f){ | |
| try{ | |
| rknn_matmul_info info{};info.M=m;info.K=k;info.N=n;info.iommu_domain_id=3; | |
| 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; | |
| Check(rknn_matmul_create(&ctx,&info,&io),"create"); | |
| Check(rknn_matmul_set_core_mask(ctx,static_cast<rknn_core_mask>(1<<core)),"core"); | |
| const int subn=w8?32:16; | |
| if(io.B.n_dims!=4||io.B.dims[2]!=unsigned(subn)||io.B.dims[3]!=32|| | |
| io.B.size!=unsigned(k*n*(w8?1:2))||io.A.size!=unsigned(m*k*(w8?1:2))||io.C.size!=unsigned(m*n*4)) | |
| throw std::runtime_error("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("DMA allocation"); | |
| for(int col=0;col<n;++col){ | |
| auto row=weight.subspan(size_t(offset+col)*k,k); | |
| if(w8){float mx=0;for(float v:row)mx=std::max(mx,std::abs(v));scales[col]=mx>0?mx/127:1;} | |
| for(int j=0;j<k;++j){ | |
| const size_t idx=((size_t(col/subn)*(k/32)+j/32)*subn+col%subn)*32+j%32; | |
| if(w8)static_cast<int8_t*>(b->virt_addr)[idx]=std::clamp(int(std::nearbyint(row[j]/scales[col])),-127,127); | |
| else static_cast<__fp16*>(b->virt_addr)[idx]=static_cast<__fp16>(row[j]); | |
| } | |
| } | |
| Check(rknn_mem_sync(ctx,b,RKNN_MEMORY_SYNC_TO_DEVICE),"sync B"); | |
| Check(rknn_matmul_set_io_mem(ctx,a,&io.A),"bind A"); | |
| Check(rknn_matmul_set_io_mem(ctx,b,&io.B),"bind B"); | |
| Check(rknn_matmul_set_io_mem(ctx,c,&io.C),"bind C"); | |
| }catch(...){Release();throw;} | |
| } | |
| ~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 Run(std::span<const float> input,std::span<float> output,int out_stride){ | |
| const int subk=w8?16:8; | |
| for(int r=0;r<m;++r){ | |
| if(w8){float mx=0;for(int j=0;j<k;++j)mx=std::max(mx,std::abs(input[r*k+j]));as[r]=mx>0?mx/127:1;} | |
| for(int j=0;j<k;++j){ | |
| const size_t idx=(size_t(j/subk)*m+r)*subk+j%subk; | |
| if(w8)static_cast<int8_t*>(a->virt_addr)[idx]=std::clamp(int(std::nearbyint(input[r*k+j]/as[r])),-127,127); | |
| else static_cast<__fp16*>(a->virt_addr)[idx]=static_cast<__fp16>(input[r*k+j]); | |
| } | |
| } | |
| Check(rknn_mem_sync(ctx,a,RKNN_MEMORY_SYNC_TO_DEVICE),"sync A"); | |
| Check(rknn_matmul_run(ctx),"run"); | |
| Check(rknn_mem_sync(ctx,c,RKNN_MEMORY_SYNC_FROM_DEVICE),"sync C"); | |
| for(int r=0;r<m;++r)for(int col=0;col<n;++col){ | |
| const size_t i=(size_t(col/4)*m+r)*4+col%4; | |
| output[r*out_stride+offset+col]=w8?float(static_cast<int32_t*>(c->virt_addr)[i])*as[r]*scales[col]:static_cast<float*>(c->virt_addr)[i]; | |
| } | |
| } | |
| }; | |
| int main(int argc,char**argv)try{ | |
| if(argc!=6)throw std::runtime_error("usage: probe DIRECTORY ROWS fp16|w8 CORES ITERATIONS"); | |
| std::string path=argv[1],mode=argv[3];int m=std::stoi(argv[2]),cores=std::stoi(argv[4]),iters=std::stoi(argv[5]); | |
| if((m!=1&&m!=128)||(cores!=1&&cores!=3)||iters<1||(mode!="w8"&&mode!="fp16"))throw std::runtime_error("invalid arguments"); | |
| auto shape=Read<int32_t>(path+"/shape.bin",3);int k=shape[0],n=shape[1],group=shape[2]; | |
| auto codes=Read<int8_t>(path+"/codes.bin",size_t(k)*n); | |
| auto scales=Read<float>(path+"/scales.bin",size_t(n)*k/group); | |
| auto input=Read<float>(path+"/input-"+std::to_string(m)+".bin",size_t(m)*k); | |
| std::vector<float> weight(size_t(n)*k),output(size_t(m)*n); | |
| for(int col=0;col<n;++col)for(int j=0;j<k;++j)weight[col*k+j]=codes[col*k+j]*scales[col*(k/group)+j/group]; | |
| auto start=Clock::now();std::vector<std::unique_ptr<Matrix>> matrices; | |
| int offset=0;size_t bytes=0; | |
| for(int core=0;core<cores;++core){ | |
| int width=((n/32)/cores+(core<(n/32)%cores))*32; | |
| matrices.push_back(std::make_unique<Matrix>(m,k,width,offset,core,mode=="w8",weight));offset+=width; | |
| bytes+=matrices.back()->io.B.size; | |
| } | |
| double load_ms=std::chrono::duration<double,std::milli>(Clock::now()-start).count(); | |
| auto run=[&]{ | |
| if(cores==1){matrices[0]->Run(input,output,n);return;} | |
| // Three worker launches are included here; production uses persistent workers. | |
| std::vector<std::thread> threads;std::vector<std::exception_ptr> errors(cores); | |
| for(int c=0;c<cores;++c)threads.emplace_back([&,c]{try{matrices[c]->Run(input,output,n);}catch(...){errors[c]=std::current_exception();}}); | |
| for(auto&t:threads)t.join();for(auto&e:errors)if(e)std::rethrow_exception(e); | |
| }; | |
| run();run();run(); | |
| double error=0,refnorm=0,implerr=0,implnorm=0; | |
| // Validate a complete first row against scalar independent reference math. | |
| float as=1;for(int j=0;j<k;++j)as=std::max(as,std::abs(input[j]));as/=127; | |
| for(int col=0;col<n;++col){ | |
| double expected=0,implemented=0;float ws=0; | |
| for(int j=0;j<k;++j)ws=std::max(ws,std::abs(weight[col*k+j]));ws=ws>0?ws/127:1; | |
| int64_t integer=0; | |
| for(int j=0;j<k;++j){ | |
| expected+=double(input[j])*weight[col*k+j]; | |
| if(mode=="w8")integer+=int(std::nearbyint(input[j]/as))*int(std::nearbyint(weight[col*k+j]/ws)); | |
| else implemented+=float(static_cast<__fp16>(input[j]))*double(float(static_cast<__fp16>(weight[col*k+j]))); | |
| } | |
| if(mode=="w8")implemented=float(integer)*as*ws; | |
| error+=std::pow(output[col]-expected,2);refnorm+=expected*expected; | |
| implerr+=std::pow(output[col]-implemented,2);implnorm+=implemented*implemented; | |
| } | |
| const double implrmse=std::sqrt(implerr/implnorm); | |
| if(!std::isfinite(implrmse)||implrmse>1e-4)throw std::runtime_error("reference mismatch "+std::to_string(implrmse)); | |
| std::vector<double> times; | |
| for(int i=0;i<iters;++i){start=Clock::now();run();times.push_back(std::chrono::duration<double,std::milli>(Clock::now()-start).count());} | |
| std::sort(times.begin(),times.end()); | |
| std::cout<<"{\"path\":\""<<path<<"\",\"mode\":\""<<mode<<"\",\"rows\":"<<m<<",\"cores\":"<<cores | |
| <<",\"p50_ms\":"<<times[times.size()/2]<<",\"mean_ms\":"<<std::accumulate(times.begin(),times.end(),0.)/times.size() | |
| <<",\"weight_bytes\":"<<bytes<<",\"load_ms\":"<<load_ms<<",\"reference_relative_rmse\":"<<implrmse | |
| <<",\"extra_output_relative_rmse\":"<<std::sqrt(error/refnorm)<<"}\n"; | |
| }catch(const std::exception&e){std::cerr<<e.what()<<"\n";return 1;} | |