Ling-3.0-tiny-RKNN / tools /official_int4_group_probe.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
4.82 kB
// Direct group32 import diagnostic. Not a full-model performance benchmark.
#include "ling3/w4_linear.h"
#include <algorithm>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <fstream>
#include <iostream>
#include <memory>
#include <numeric>
#include <stdexcept>
#include <string>
#include <vector>
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> out(count);
if (!f.read(reinterpret_cast<char*>(out.data()), out.size()*sizeof(T))) throw std::runtime_error("read: "+path);
return out;
}
int main(int argc, char **argv) try {
if (argc != 4) throw std::runtime_error("usage: probe DIRECTORY ROWS ITERATIONS");
unsetenv("LING3_PREFILL_W4A4");
const std::string path=argv[1];
const int rows=std::stoi(argv[2]), iterations=std::stoi(argv[3]);
auto shape=Read<int32_t>(path+"/shape.bin",3);
const int k=shape[0], n=shape[1], group=shape[2], groups=k/group;
if (rows<1 || rows>128 || iterations<1 || k<32 || n<64 || group<32 || k%group) throw std::runtime_error("bad shape");
auto codes=Read<int8_t>(path+"/codes.bin",size_t(k)*n);
auto scales=Read<float>(path+"/scales.bin",size_t(n)*groups);
auto input=Read<float>(path+"/input-"+std::to_string(rows)+".bin",size_t(rows)*k);
auto expected=Read<float>(path+"/oracle-"+std::to_string(rows)+".bin",size_t(rows)*n);
std::vector<std::unique_ptr<ling3::DynamicW4Linear>> runners;
std::vector<std::vector<float>> inputs(groups,std::vector<float>(size_t(rows)*group));
size_t resident=0;
for (int g=0;g<groups;++g) {
std::vector<std::byte> packed(size_t(group)*n/2);
std::vector<float> scale(n);
std::vector<int32_t> correction(n);
// Avoid INT16 overflow in the single per-channel baseline matmul.
int splits=1;
for (;;) {
bool safe=true;
for (int col=0;col<n && safe;++col) for (int part=0;part<splits;++part) {
int sum=0;
for (int j=part*group/splits;j<(part+1)*group/splits;++j) sum+=std::abs(int(codes[col*k+g*group+j]));
if (8*sum>32767) safe=false;
}
if (safe) break;
splits*=2;
if (group%splits || (group/splits)%32) throw std::runtime_error("no safe split");
}
for (int col=0;col<n;++col) {
scale[col]=scales[col*groups+g];
for (int j=0;j<group;++j) {
const int code=codes[col*k+g*group+j];
const size_t i=size_t(j)*n+col;
packed[i/2] |= std::byte((code & 15) << (4*(i%2)));
correction[col]+=8*code;
}
}
for (int r=0;r<rows;++r) std::copy_n(input.begin()+r*k+g*group,group,inputs[g].begin()+r*group);
runners.push_back(std::make_unique<ling3::DynamicW4Linear>(
ling3::W4LinearConfig{group,n,splits,{0,1,2},3},packed,scale,correction));
runners.back()->PrepareBatch(rows);
resident+=runners.back()->resident_weight_bytes();
}
std::vector<float> output(size_t(rows)*n), partial(output.size());
auto run=[&] {
std::fill(output.begin(),output.end(),0.f);
for (int g=0;g<groups;++g) {
runners[g]->RunBatch(inputs[g],rows,partial);
for (size_t j=0;j<output.size();++j) output[j]+=partial[j];
}
};
run(); run(); run();
double error2=0, ref2=0, max_error=0;
for (size_t i=0;i<output.size();++i) {
const double e=double(output[i])-expected[i];
if (!std::isfinite(output[i])) throw std::runtime_error("nonfinite output");
error2+=e*e; ref2+=double(expected[i])*expected[i]; max_error=std::max(max_error,std::abs(e));
}
const double rmse=std::sqrt(error2/std::max(ref2,1e-30));
if (rmse>1e-5) throw std::runtime_error("NPU integer oracle mismatch: "+std::to_string(rmse));
std::vector<double> ms;
for (int i=0;i<iterations;++i) {
auto start=std::chrono::steady_clock::now(); run();
ms.push_back(std::chrono::duration<double,std::milli>(std::chrono::steady_clock::now()-start).count());
}
std::sort(ms.begin(),ms.end());
std::cout << "{\"case\":\""<<path<<"\",\"rows\":"<<rows<<",\"groups\":"<<groups
<<",\"mean_ms\":"<<std::accumulate(ms.begin(),ms.end(),0.)/ms.size()<<",\"p50_ms\":"<<ms[ms.size()/2]
<<",\"p95_ms\":"<<ms[std::min(ms.size()-1,ms.size()*95/100)]<<",\"oracle_relative_rmse\":"<<rmse
<<",\"oracle_max_abs_error\":"<<max_error<<",\"resident_weight_bytes\":"<<resident<<"}\n";
} catch (const std::exception &e) { std::cerr<<e.what()<<"\n"; return 1; }