File size: 7,170 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
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
#include "ling3/decoder.h"
#include "ling3/model_package.h"
#include <algorithm>
#include <cmath>
#include <cstdlib>
#include <iostream>
#include <sched.h>
#include <stdexcept>
#include <sys/resource.h>
#include <vector>
#include <unistd.h>

int main(int argc, char ** argv) {
    try {
        if (argc != 2) throw std::invalid_argument("usage: ling3-checkpoint-model-test MODEL.l3r");
        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);
        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");
        const ling3::ModelPackage package(argv[1]);
        auto instance = std::make_unique<ling3::Decoder>(package, 512);
        auto & decoder = *instance;
        decoder.PrepareBatch(128);
        decoder.PrepareBatch(16);
        std::vector<float> logits(package.header().vocab_size);
        std::vector<std::uint32_t> prefix(128, 16), suffix(16, 17);
        decoder.EvalBatch(prefix, logits);
        decoder.EvalBatch(prefix, logits);
        auto checkpoint = decoder.SaveCheckpoint();
        auto full = decoder.SaveState();
        decoder.EvalBatch(suffix, logits);
        const auto expected_batch = logits;
        auto future = decoder.SaveCheckpoint();
        decoder.Eval(18, logits);
        const auto expected_decode = logits;
        if (decoder.RestoreCheckpoint(*checkpoint) != 256) throw std::runtime_error("wrong restored position");
        decoder.EvalBatch(suffix, logits);
        if (logits != expected_batch) throw std::runtime_error("batch logits differ after checkpoint restore");
        decoder.Eval(18, logits);
        if (logits != expected_decode) throw std::runtime_error("decode logits differ after checkpoint restore");
        for (auto value : logits) if (!std::isfinite(value)) throw std::runtime_error("non-finite logits");
        bool rejected = false;
        try { decoder.RestoreCheckpoint(*future); } catch (const std::invalid_argument &) { rejected = true; }
        if (!rejected) throw std::runtime_error("overwritten checkpoint was accepted");
        decoder.RestoreCheckpoint(*checkpoint);
        decoder.Eval(19, logits);
        decoder.Reset();
        rejected = false;
        try { decoder.RestoreCheckpoint(*checkpoint); } catch (const std::invalid_argument &) { rejected = true; }
        if (!rejected) throw std::runtime_error("checkpoint survived Reset");
        decoder.EvalBatch(std::vector<std::uint32_t>(128, 21), logits);
        auto other = decoder.SaveState();
        decoder.Eval(22, logits); const auto expected_other = logits;
        decoder.RestoreState(*full);
        decoder.EvalBatch(suffix, logits);
        if (logits != expected_batch) throw std::runtime_error("A-B-A batch state mismatch");
        decoder.Eval(18, logits);
        if (logits != expected_decode) throw std::runtime_error("A-B-A decode state mismatch");
        decoder.RestoreState(*other); decoder.Eval(22, logits);
        if (logits != expected_other) throw std::runtime_error("B restore mismatch");
        char temporary[] = "/tmp/ling3-model-state-XXXXXX";
        const int fd = mkstemp(temporary); if (fd < 0) throw std::runtime_error("mkstemp failed"); close(fd);
        try {
            ling3::WriteDecoderState(*full, temporary);
            auto loaded = ling3::ReadDecoderState(temporary, decoder.StateSignature(), 512);
            decoder.Reset(); decoder.RestoreState(loaded); decoder.EvalBatch(suffix, logits);
            if (logits != expected_batch) throw std::runtime_error("disk state mismatch");
            unlink(temporary);
        } catch (...) { unlink(temporary); throw; }
        const auto mla=decoder.AttentionStats();
        if(!mla.npu_calls || mla.fallbacks)throw std::runtime_error("NPU attention was not validated");
        std::cout << "MLA npu_calls=" << mla.npu_calls << " fallbacks=" << mla.fallbacks << '\n';
        std::cout << "PASS: all " << logits.size() << " batch/decode logits identical after restore; "
                  << "overwritten and reset checkpoints rejected; snapshot_bytes="
                  << ling3::Decoder::CheckpointBytes(*checkpoint) << '\n';
        std::cout << "PASS: complete A-B-A and disk state restore; full_state_bytes=" << full->bytes() << '\n';
        // A generated prefix is valid continuation state, but it is not the
        // same numerical path as re-prefilling those tokens in a batch.
        decoder.Reset(); decoder.EvalBatch(prefix, logits);
        auto seed = decoder.SaveState();
        std::vector<std::uint32_t> generated_tokens(32);
        for (std::size_t i=0; i<generated_tokens.size(); ++i) {
            generated_tokens[i] = 16 + i%7;
            decoder.Eval(generated_tokens[i], logits);
        }
        auto continuation = decoder.SaveState();
        decoder.EvalBatch(suffix, logits); const auto continued_logits = logits;
        decoder.Reset(); decoder.RestoreState(*continuation); decoder.EvalBatch(suffix, logits);
        if (logits != continued_logits) throw std::runtime_error("generated state restore mismatch");
        decoder.RestoreState(*seed); decoder.EvalBatch(generated_tokens, logits); decoder.EvalBatch(suffix, logits);
        double absolute = 0, maximum = 0;
        std::size_t differing = 0;
        for (std::size_t i=0; i<logits.size(); ++i) {
            if (!std::isfinite(logits[i])) throw std::runtime_error("non-finite batch reference logits");
            const double d = std::abs(double(logits[i])-continued_logits[i]);
            absolute += d; maximum = std::max(maximum, d); differing += d != 0;
        }
        std::cout << "PASS: generated state restore all logits identical\n";
        std::cout << "generated_vs_batch={\"mae\":" << absolute/logits.size()
                  << ",\"max_absolute\":" << maximum << ",\"differing_logits\":" << differing << "}\n";
        // Sequential instances: never load two sets of NPU weights at once.
        instance.reset();
        ling3::Decoder replacement(package, 512);
        replacement.PrepareBatch(16);
        replacement.RestoreState(*full);
        replacement.EvalBatch(suffix, logits);
        if (logits != expected_batch) throw std::runtime_error("cross-instance state mismatch");
        std::cout << "PASS: complete state restores into a new decoder instance\n";
    } catch (const std::exception & e) {
        std::cerr << e.what() << '\n';
        return 1;
    }
}