// Direct group32 import diagnostic. Not a full-model performance benchmark. #include "ling3/w4_linear.h" #include #include #include #include #include #include #include #include #include #include #include template std::vector 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 out(count); if (!f.read(reinterpret_cast(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(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(path+"/codes.bin",size_t(k)*n); auto scales=Read(path+"/scales.bin",size_t(n)*groups); auto input=Read(path+"/input-"+std::to_string(rows)+".bin",size_t(rows)*k); auto expected=Read(path+"/oracle-"+std::to_string(rows)+".bin",size_t(rows)*n); std::vector> runners; std::vector> inputs(groups,std::vector(size_t(rows)*group)); size_t resident=0; for (int g=0;g packed(size_t(group)*n/2); std::vector scale(n); std::vector correction(n); // Avoid INT16 overflow in the single per-channel baseline matmul. int splits=1; for (;;) { bool safe=true; for (int col=0;col32767) 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( ling3::W4LinearConfig{group,n,splits,{0,1,2},3},packed,scale,correction)); runners.back()->PrepareBatch(rows); resident+=runners.back()->resident_weight_bytes(); } std::vector output(size_t(rows)*n), partial(output.size()); auto run=[&] { std::fill(output.begin(),output.end(),0.f); for (int g=0;gRunBatch(inputs[g],rows,partial); for (size_t j=0;j1e-5) throw std::runtime_error("NPU integer oracle mismatch: "+std::to_string(rmse)); std::vector ms; for (int i=0;i(std::chrono::steady_clock::now()-start).count()); } std::sort(ms.begin(),ms.end()); std::cout << "{\"case\":\""<