File size: 8,100 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 | // 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;}
|