Download tools/official_int4_group_probe.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_int4_group_probe.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/official_int4_group_probe.cpp
-
curl -L -o official_int4_group_probe.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/official_int4_group_probe.cpp
4.82 kB
| // Direct group32 import diagnostic. Not a full-model performance benchmark. | |
| 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; } | |