Download tests/checkpoint_model_test.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 7.17 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tests/checkpoint_model_test.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tests/checkpoint_model_test.cpp
-
curl -L -o checkpoint_model_test.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tests/checkpoint_model_test.cpp
7.17 kB
| 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; | |
| } | |
| } | |