File size: 4,824 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
// 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; }