Ling-3.0-tiny-RKNN / tests /checkpoint_model_test.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
7.17 kB
#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;
}
}