Ling-3.0-tiny-RKNN / tools /official_linear_bridge_probe.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
8.1 kB
// Test FP16 and W8 execution bridges for released group32 weights.
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#include <algorithm>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <fstream>
#include <iostream>
#include <memory>
#include <numeric>
#include <span>
#include <stdexcept>
#include <string>
#include <thread>
#include <vector>
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;}