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;}