#include "ling3/decoder.h" #include "ling3/model_package.h" #include #include #include #include #include #include #include #include #include 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(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(package, 512); auto & decoder = *instance; decoder.PrepareBatch(128); decoder.PrepareBatch(16); std::vector logits(package.header().vocab_size); std::vector 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(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 generated_tokens(32); for (std::size_t i=0; i