File size: 5,854 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
#include "ling3/decoder.h"
#include "ling3/model_package.h"
#include "ling3/tokenizer.h"
#include <nlohmann/json.hpp>
#include <algorithm>
#include <chrono>
#include <cmath>
#include <cstdlib>
#include <filesystem>
#include <fstream>
#include <iostream>
#include <sched.h>
#include <sys/resource.h>

// Diagnostic entry only: exact official token IDs, one warmed decoder for all cases.
int main(int argc, char **argv) {
    using Json = nlohmann::json;
    using Clock = std::chrono::steady_clock;
    try {
        if (argc != 4) throw std::invalid_argument("usage: quant-loss-probe MODEL.l3r SUITE.json OUTPUT_DIR");
        for (auto key : {"LING3_PREFILL_W4A4", "LING3_EXPERT_ALL_CORES", "LING3_GDN_PREFILL_DIR"}) unsetenv(key);
        for (auto key : {"LING3_PREWARM_EXPERTS", "LING3_EXPERT_BALANCED", "LING3_EXPERT_ZERO_COPY",
                         "LING3_GDN_CPU_PREFILL", "LING3_GDN_CPU_FP32_STATE", "LING3_GDN_CPU_DECODE",
                         "LING3_GDN_FULL_FP32", "LING3_MLA_SIMD", "LING3_VECTOR_MATH"}) setenv(key,"1",1);
        setenv("LING3_MLA_BACKEND", "auto", 1);
        cpu_set_t affinity; CPU_ZERO(&affinity);
        for (int core=4; core<8; ++core) CPU_SET(core,&affinity);
        if (sched_setaffinity(0,sizeof(affinity),&affinity)) throw std::runtime_error("affinity failed");
        rlimit limit{};
        if (getrlimit(RLIMIT_NOFILE,&limit)) throw std::runtime_error("getrlimit failed");
        limit.rlim_cur=std::min<rlim_t>(262144,limit.rlim_max);
        if (setrlimit(RLIMIT_NOFILE,&limit)) throw std::runtime_error("setrlimit failed");
        std::ifstream suite_file(argv[2]);
        Json suite; suite_file >> suite;
        const ling3::ModelPackage model(argv[1]);
        const auto &t=model.tensor("tokenizer");
        ling3::Tokenizer tokenizer({t.data,static_cast<std::size_t>(t.entry->data_bytes)});
        std::size_t capacity=4096;
        for(const auto &c:suite.at("cases")) capacity=std::max(capacity,c.at("prompt_ids").size()+c.at("teacher_ids").size());
        ling3::Decoder decoder(model,capacity);
        for(std::size_t rows=1;rows<=128;rows*=2)decoder.PrepareBatch(rows);
        std::vector<float> logits(model.header().vocab_size);
        decoder.Eval(16,logits); decoder.Reset();
        const std::filesystem::path output(argv[3]);
        std::filesystem::create_directories(output);
        Json results=Json::array();
        for(const auto &c:suite.at("cases")) {
            const auto name=c.at("id").get<std::string>();
            const auto prompt=c.at("prompt_ids").get<std::vector<std::uint32_t>>();
            auto teacher=c.at("teacher_ids").get<std::vector<std::uint32_t>>();
            if(prompt.empty()||teacher.empty())throw std::runtime_error("empty sequence");
            if(tokenizer.Encode(c.at("prompt_text").get<std::string>())!=prompt ||
               tokenizer.Encode(c.at("target_text").get<std::string>())!=teacher)
                throw std::runtime_error("official/C++ tokenizer mismatch: "+name);
            const auto full_teacher_tokens=teacher.size();
            if(const char* limit=std::getenv("LING3_LOSS_STEPS")){
                const auto count=std::stoul(limit);
                if(count<1)throw std::invalid_argument("LING3_LOSS_STEPS must be positive");
                teacher.resize(std::min(teacher.size(),count));
            }
            decoder.Reset();
            const auto before=decoder.AttentionStats();
            const auto start=Clock::now();
            for(std::size_t offset=0;offset<prompt.size();) {
                const auto rows=std::min<std::size_t>(128,prompt.size()-offset);
                const auto block=std::span<const std::uint32_t>(prompt).subspan(offset,rows);
                if(rows==1)decoder.Eval(block[0],logits);
                else if(offset+rows==prompt.size())decoder.EvalBatch(block,logits);
                else decoder.EvalBatchState(block);
                offset+=rows;
            }
            const auto prefill=Clock::now();
            std::ofstream values(output/(name+".rknn.f32"),std::ios::binary);
            if(!values)throw std::runtime_error("cannot open logit output");
            for(std::size_t step=0;step<teacher.size();++step) {
                if(!std::all_of(logits.begin(),logits.end(),[](float x){return std::isfinite(x);}))
                    throw std::runtime_error("nonfinite logits: "+name+" step="+std::to_string(step));
                values.write(reinterpret_cast<const char *>(logits.data()),logits.size()*sizeof(float));
                if(step+1<teacher.size())decoder.Eval(teacher[step],logits);
                if((step+1)%256==0)std::cout<<name<<" teacher_steps="<<step+1<<'/'<<teacher.size()<<std::endl;
            }
            if(!values)throw std::runtime_error("failed writing logits");
            const auto end=Clock::now(); const auto after=decoder.AttentionStats();
            Json row{{"id",name},{"thinking",c.at("thinking")},{"prompt_tokens",prompt.size()},
                {"teacher_tokens",teacher.size()},{"full_teacher_tokens",full_teacher_tokens},
                {"screening",teacher.size()!=full_teacher_tokens},{"all_logits_finite",true},{"tokenizer_match",true},
                {"context_capacity",capacity},{"mla_npu_calls",after.npu_calls-before.npu_calls},
                {"mla_fallbacks",after.fallbacks-before.fallbacks},
                {"prefill_ms",std::chrono::duration<double,std::milli>(prefill-start).count()},
                {"total_ms",std::chrono::duration<double,std::milli>(end-start).count()}};
            results.push_back(row);std::cout<<row.dump()<<std::endl;
            std::ofstream(output/"rknn-results.json")<<results.dump(2)<<'\n';
        }
        std::cout<<"PASS: all paired modes exported"<<std::endl;
    } catch(const std::exception &e){std::cerr<<e.what()<<'\n';return 1;}
}