From 7be881f8e174101a514e5ca27513806db9d3cc92 Mon Sep 17 00:00:00 2001 From: linear3735 Date: Mon, 28 Sep 2026 15:45:04 +0800 Subject: [PATCH 1/3] Add CPU checkpoint loading for Laya --- Cargo.lock | 219 +++++++++++++++++++++++++-- Cargo.toml | 2 +- README.md | 4 +- recipe/laya/native/export_weights.py | 27 ++++ src/models/laya/Cargo.toml | 17 +++ src/models/laya/README.md | 19 ++- src/models/laya/src/config.rs | 91 +++++++++++ src/models/laya/src/lib.rs | 2 + src/models/laya/src/weights.rs | 157 +++++++++++++++++++ src/models/laya/tests/checkpoint.rs | 180 ++++++++++++++++++++++ src/models/laya/tests/weights.rs | 93 ++++++++++++ 11 files changed, 796 insertions(+), 15 deletions(-) create mode 100644 recipe/laya/native/export_weights.py create mode 100644 src/models/laya/Cargo.toml create mode 100644 src/models/laya/src/config.rs create mode 100644 src/models/laya/src/lib.rs create mode 100644 src/models/laya/src/weights.rs create mode 100644 src/models/laya/tests/checkpoint.rs create mode 100644 src/models/laya/tests/weights.rs diff --git a/Cargo.lock b/Cargo.lock index be37964..9212bac 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,12 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "anyhow" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + [[package]] name = "atomic-waker" version = "1.1.2" @@ -78,6 +84,15 @@ version = "2.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3ded4057c258ba199e2d26386d3af3780957ecaee6c4ef4041c6b4b8b97c0b06" +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + [[package]] name = "bumpalo" version = "3.20.3" @@ -119,10 +134,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.3.1", "rand_core", ] +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + [[package]] name = "cpufeatures" version = "0.3.1" @@ -132,6 +156,32 @@ dependencies = [ "libc", ] +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + [[package]] name = "displaydoc" version = "0.2.7" @@ -140,7 +190,7 @@ checksum = "c6232dd377dcc64799954cbd3a9bb882e9cdc1308ccd87b1c098f1fb2eaf82a8" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -153,6 +203,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "fastrand" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" + [[package]] name = "find-msvc-tools" version = "0.1.13" @@ -197,7 +253,7 @@ checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -228,6 +284,16 @@ dependencies = [ "slab", ] +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + [[package]] name = "getrandom" version = "0.2.17" @@ -255,6 +321,17 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + [[package]] name = "http" version = "1.5.0" @@ -494,6 +571,12 @@ version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litemap" version = "0.8.3" @@ -524,6 +607,15 @@ version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" +[[package]] +name = "memmap2" +version = "0.9.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1219ed1b7f229ee7104d281dd01d6802fe28bb6e95d292942c4daacdeb798c0" +dependencies = [ + "libc", +] + [[package]] name = "mime" version = "0.3.17" @@ -551,6 +643,20 @@ dependencies = [ "tokio", ] +[[package]] +name = "omni-laya" +version = "0.1.0" +dependencies = [ + "anyhow", + "half", + "memmap2", + "safetensors", + "serde", + "serde_json", + "sha2", + "tempfile", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -745,6 +851,19 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" +[[package]] +name = "rustix" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.23.45" @@ -792,6 +911,16 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "safetensors" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "172dd94c5a87b5c79f945c863da53b2ebc7ccef4eca24ac63cca66a41aab2178" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "serde" version = "1.0.229" @@ -799,6 +928,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", + "serde_derive", ] [[package]] @@ -818,7 +948,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -857,6 +987,17 @@ dependencies = [ "serde", ] +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + [[package]] name = "shlex" version = "2.0.1" @@ -907,6 +1048,17 @@ version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "syn" version = "3.0.6" @@ -935,7 +1087,20 @@ checksum = "901704edd0dfe137f1987838ee4f259e4e063c31371bdb423f7ae38ec6f77f02" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", +] + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", ] [[package]] @@ -955,7 +1120,7 @@ checksum = "fe5197923287db20a58125f0bc85c062f7f2c892de97b18c356f9efb14b28524" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -998,7 +1163,7 @@ checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -1096,6 +1261,12 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + [[package]] name = "unicode-ident" version = "1.0.26" @@ -1126,6 +1297,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + [[package]] name = "want" version = "0.3.1" @@ -1184,7 +1361,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn", + "syn 3.0.6", "wasm-bindgen-shared", ] @@ -1352,10 +1529,30 @@ checksum = "33811428bee40dbceb6d545e95754741d17a6aef9a4849f0fd62e2ba4f412a78" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] +[[package]] +name = "zerocopy" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6df92bf3d9227be3d53173901ddbffac2babc27ae50f397776ffd6dc33f800cb" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac4f328cf2f05d084e496c3e9c3f33ed0a183656a16e1fcec4d464d8373aec82" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "zerofrom" version = "0.1.8" @@ -1373,7 +1570,7 @@ checksum = "f75b4683f6c7f45248d4d64056a24298c6281e0993356d7d1b4a1a962ef10d4a" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] @@ -1413,7 +1610,7 @@ checksum = "34df6fc39dbd26ddc9c10e6a2984476e13acce22e64e4487636ef494369225da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 8036e4b..4f33792 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend"] +members = ["src/frontend", "src/models/laya"] resolver = "3" diff --git a/README.md b/README.md index d12aa29..7b2a777 100644 --- a/README.md +++ b/README.md @@ -47,7 +47,7 @@ Implementation code lives under `src/`; recipes and documentation stay at the re | [`recipe/`](recipe/) | Model setup instructions, launch commands, configuration examples, and example requests. | | [`docs/`](docs/) | Project documentation and architecture assets. | -The frontend is a Cargo workspace member. Model and backend directories currently document planned work; they do not prescribe process boundaries. +The frontend and Laya checkpoint reader are Cargo workspace members. Model execution and GPU backends remain planned; these directories do not prescribe process boundaries. ## Supported models @@ -55,7 +55,7 @@ LAYA can run as an external Python worker for text requests. Its in-repository m | Model | Status | | --- | --- | -| LAYA | [External worker](recipe/laya/README.md); model engine planned | +| LAYA | [External worker](recipe/laya/README.md); [CPU checkpoint reader](src/models/laya/README.md); model execution planned | CUDA and Metal coverage will be documented per model as implementations are added and validated. diff --git a/recipe/laya/native/export_weights.py b/recipe/laya/native/export_weights.py new file mode 100644 index 0000000..34aed02 --- /dev/null +++ b/recipe/laya/native/export_weights.py @@ -0,0 +1,27 @@ +"""CPU reference hashes for each tensor's FP32, FP16 and BF16 conversions.""" + +import argparse, hashlib, json +from pathlib import Path +import torch +from safetensors import safe_open + +p = argparse.ArgumentParser() +p.add_argument("checkpoint", type=Path) +p.add_argument("output", type=Path) +a = p.parse_args() +torch.set_num_threads(4) +rows = [] +with safe_open(a.checkpoint / "model.safetensors", framework="pt", device="cpu") as f: + for name in f.keys(): + x = f.get_tensor(name) + row = {"name": name, "shape": list(x.shape), "source_dtype": str(x.dtype)} + for key, dtype in [ + ("f32", torch.float32), + ("f16", torch.float16), + ("bf16", torch.bfloat16), + ]: + y = x.to(torch.float32).to(dtype).contiguous() + row[key] = hashlib.sha256(y.view(torch.uint8).numpy().tobytes()).hexdigest() + rows.append(row) +a.output.write_text(json.dumps(rows, indent=2)) +print("WEIGHT_ORACLE", len(rows), flush=True) diff --git a/src/models/laya/Cargo.toml b/src/models/laya/Cargo.toml new file mode 100644 index 0000000..d9e73e8 --- /dev/null +++ b/src/models/laya/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "omni-laya" +version = "0.1.0" +edition = "2024" +publish = false + +[dependencies] +anyhow = "1" +half = "2" +memmap2 = "0.9" +safetensors = "0.6" +serde = { version = "1", features = ["derive"] } +serde_json = "1" + +[dev-dependencies] +sha2 = "0.10" +tempfile = "3" diff --git a/src/models/laya/README.md b/src/models/laya/README.md index ec10b58..858da19 100644 --- a/src/models/laya/README.md +++ b/src/models/laya/README.md @@ -4,4 +4,21 @@ LAYA is the first planned System1-Omni model. This directory owns its complete r GPU operations and kernel implementations belong in [`backends/cuda/`](../../backends/cuda/) and [`backends/metal/`](../../backends/metal/). Setup and usage examples belong in the top-level [`recipe/`](../../../recipe/) directory. -Status: planned; no model implementation or validated GPU backend support yet. +The `omni-laya` crate currently reads and checks the English Laya 0.3.20 checkpoint. `Config::load` validates the architecture and temperatures; `Weights` checks tensor names and shapes and converts FP32, FP16 and BF16 values. `checkpoint_tensors()` lists the 206 expected tensors. Each backend chooses its own storage precision. + +Keep checkpoint files unchanged while `Weights` holds a read-only memory mapping. This crate does not yet execute inference. + +## CPU checks + +The normal workspace tests cover configuration errors, malformed tensors, inventory mismatches and conversion boundaries without downloading weights. + +To check the complete checkpoint, use `convaiinnovations/laya` revision `55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851` and a Python environment with PyTorch, safetensors and NumPy: + +```sh +export LAYA_CHECKPOINT=/path/to/laya/snapshot +export LAYA_WEIGHT_ORACLE=/tmp/laya-weight-oracle.json +python recipe/laya/native/export_weights.py "$LAYA_CHECKPOINT" "$LAYA_WEIGHT_ORACLE" +cargo test --release --locked -p omni-laya --test weights -- --ignored +``` + +These two CPU tests check all 206 tensor names and shapes, 618 conversion hashes, and the legacy temperature buffer. The normal CI job skips them because it does not download the full checkpoint. diff --git a/src/models/laya/src/config.rs b/src/models/laya/src/config.rs new file mode 100644 index 0000000..8b85c39 --- /dev/null +++ b/src/models/laya/src/config.rs @@ -0,0 +1,91 @@ +use anyhow::{Context, Result, ensure}; +use serde::Deserialize; +use std::{collections::HashMap, fs, path::Path}; + +#[derive(Debug, Deserialize)] +pub struct AgentConfig { + pub max_len: usize, + pub head_max_len: usize, + pub head_layers: usize, + pub temperature: Vec, + pub temperature_by_options: HashMap, +} + +#[derive(Debug, Deserialize)] +pub struct EncoderConfig { + pub hidden_size: usize, + pub intermediate_size: usize, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub vocab_size: usize, + pub norm_eps: f32, + pub local_attention: usize, + pub layer_types: Vec, + pub rope_parameters: serde_json::Value, +} + +pub struct Config { + pub agent: AgentConfig, + pub encoder: EncoderConfig, +} +impl Config { + pub fn load(dir: &Path) -> Result { + let read = |name| fs::read(dir.join(name)).with_context(|| format!("read {name}")); + let agent: AgentConfig = serde_json::from_slice(&read("rl_agent_config.json")?)?; + let encoder_bytes = read("encoder/config.json")?; + let raw: serde_json::Value = serde_json::from_slice(&encoder_bytes)?; + ensure!( + raw["model_type"] == "modernbert" + && raw["hidden_activation"] == "gelu" + && raw["attention_bias"] == false + && raw["mlp_bias"] == false + && raw["norm_bias"] == false, + "unsupported encoder activation, bias or model type" + ); + let encoder: EncoderConfig = serde_json::from_slice(&encoder_bytes)?; + ensure!( + agent.max_len == 512 && agent.head_max_len == 192 && agent.head_layers == 2, + "native Laya supports max_len=512, head_max_len=192, head_layers=2" + ); + ensure!( + encoder.hidden_size == 1024 + && encoder.intermediate_size == 2624 + && encoder.num_attention_heads == 16 + && encoder.num_hidden_layers == 28 + && encoder.vocab_size == 50368 + && encoder.local_attention == 128 + && encoder.norm_eps == 1e-5, + "unsupported encoder configuration" + ); + let expected: Vec<_> = (0..28) + .map(|i| { + if i % 3 == 0 { + "full_attention" + } else { + "sliding_attention" + } + }) + .collect(); + ensure!( + encoder.layer_types == expected, + "unsupported attention schedule" + ); + for (kind, theta) in [("full_attention", 160000.0), ("sliding_attention", 10000.0)] { + let r = &encoder.rope_parameters[kind]; + ensure!( + r["rope_type"] == "default" && r["rope_theta"].as_f64() == Some(theta), + "unsupported RoPE configuration" + ); + } + ensure!( + agent.temperature.len() == 3 + && agent + .temperature + .iter() + .chain(agent.temperature_by_options.values()) + .all(|t| t.is_finite() && *t > 0.0), + "invalid temperatures" + ); + Ok(Self { agent, encoder }) + } +} diff --git a/src/models/laya/src/lib.rs b/src/models/laya/src/lib.rs new file mode 100644 index 0000000..fbc136f --- /dev/null +++ b/src/models/laya/src/lib.rs @@ -0,0 +1,2 @@ +pub mod config; +pub mod weights; diff --git a/src/models/laya/src/weights.rs b/src/models/laya/src/weights.rs new file mode 100644 index 0000000..69d3be6 --- /dev/null +++ b/src/models/laya/src/weights.rs @@ -0,0 +1,157 @@ +use anyhow::{Result, ensure}; +use half::{bf16, f16}; +use memmap2::Mmap; +use safetensors::{Dtype, SafeTensors}; +use std::{fs::File, path::Path}; + +pub struct Weights { + data: Mmap, +} +impl Weights { + /// Reject omitted, extra or duplicate names in the expected checkpoint inventory. + pub fn validate_names<'a>(&self, names: impl IntoIterator) -> Result<()> { + let tensors = SafeTensors::deserialize(&self.data)?; + let mut expected = std::collections::BTreeSet::new(); + for name in names { + ensure!(expected.insert(name), "duplicate expected tensor: {name}"); + } + let names = tensors.names(); + let actual: std::collections::BTreeSet<_> = names.into_iter().collect(); + ensure!( + expected == actual, + "tensor inventory does not match checkpoint" + ); + Ok(()) + } + /// The checkpoint must remain immutable while the mapping exists. + pub fn open(path: &Path) -> Result { + let file = File::open(path)?; + // SAFETY: model files are read-only inputs; no mutable mapping is created. + let data = unsafe { Mmap::map(&file)? }; + SafeTensors::deserialize(&data)?; + Ok(Self { data }) + } + pub fn f32(&self, name: &str, shape: &[usize]) -> Result> { + let tensors = SafeTensors::deserialize(&self.data)?; + let t = tensors.tensor(name)?; + ensure!( + t.shape() == shape, + "{name}: expected {shape:?}, got {:?}", + t.shape() + ); + let out = match t.dtype() { + Dtype::F16 => t + .data() + .as_chunks::<2>() + .0 + .iter() + .map(|b| f16::from_bits(u16::from_le_bytes([b[0], b[1]])).to_f32()) + .collect(), + Dtype::BF16 => t + .data() + .as_chunks::<2>() + .0 + .iter() + .map(|b| bf16::from_bits(u16::from_le_bytes([b[0], b[1]])).to_f32()) + .collect(), + Dtype::F32 => t + .data() + .as_chunks::<4>() + .0 + .iter() + .map(|b| f32::from_le_bytes(*b)) + .collect(), + dt => anyhow::bail!("{name}: unsupported dtype {dt:?}"), + }; + Ok(out) + } + pub fn bf16(&self, name: &str, shape: &[usize]) -> Result> { + Ok(self + .f32(name, shape)? + .into_iter() + .map(|f| bf16::from_f32(f).to_bits()) + .collect()) + } + pub fn f16(&self, name: &str, shape: &[usize]) -> Result> { + Ok(self + .f32(name, shape)? + .into_iter() + .map(|f| f16::from_f32(f).to_bits()) + .collect()) + } +} + +/// Names and shapes in the supported checkpoint; storage precision is chosen by the caller. +pub struct TensorSpec { + pub name: String, + pub shape: Vec, +} +pub fn checkpoint_tensors() -> Vec { + const D: usize = 1024; + let mut tensors = Vec::new(); + let mut add = |name: &str, shape: &[usize]| { + tensors.push(TensorSpec { + name: name.into(), + shape: shape.into(), + }); + }; + add("encoder.embeddings.tok_embeddings.weight", &[50368, D]); + add("encoder.embeddings.norm.weight", &[D]); + add("encoder.final_norm.weight", &[D]); + for i in 0..28 { + let p = format!("encoder.layers.{i}"); + if i > 0 { + add(&format!("{p}.attn_norm.weight"), &[D]); + } + add(&format!("{p}.mlp_norm.weight"), &[D]); + for (name, shape) in [ + ("attn.Wqkv.weight", vec![3 * D, D]), + ("attn.Wo.weight", vec![D, D]), + ("mlp.Wi.weight", vec![5248, D]), + ("mlp.Wo.weight", vec![D, 2624]), + ] { + add(&format!("{p}.{name}"), &shape); + } + } + add("type_emb.weight", &[3, D]); + for i in 0..2 { + let p = format!("head.layers.{i}"); + for n in ["norm1.weight", "norm1.bias", "norm2.weight", "norm2.bias"] { + add(&format!("{p}.{n}"), &[D]); + } + for (n, rows, cols) in [ + ("self_attn.in_proj_weight", 3 * D, D), + ("self_attn.out_proj.weight", D, D), + ("linear1.weight", 4 * D, D), + ("linear2.weight", D, 4 * D), + ] { + add(&format!("{p}.{n}"), &[rows, cols]); + } + for (n, len) in [ + ("self_attn.in_proj_bias", 3 * D), + ("self_attn.out_proj.bias", D), + ("linear1.bias", 4 * D), + ("linear2.bias", D), + ] { + add(&format!("{p}.{n}"), &[len]); + } + } + for n in ["scorer.0.weight", "scorer.0.bias"] { + add(n, &[D]); + } + for (p, n, k) in [ + ("scorer.1", D, D), + ("scorer.3", 1, D), + ("act_head.0", 256, 1028), + ("act_head.2", 2, 256), + ] { + add(&format!("{p}.weight"), &[n, k]); + add(&format!("{p}.bias"), &[n]); + } + + // Laya 0.3.20 common.py registers this legacy buffer but forward does not + // consume it. Agent decoding uses fitted config temperatures instead. + // Keep it in the inventory for complete checkpoint accounting. + add("temperature", &[3]); + tensors +} diff --git a/src/models/laya/tests/checkpoint.rs b/src/models/laya/tests/checkpoint.rs new file mode 100644 index 0000000..2b3493b --- /dev/null +++ b/src/models/laya/tests/checkpoint.rs @@ -0,0 +1,180 @@ +use omni_laya::{config::Config, weights::Weights}; +use serde_json::{Value, json}; +use std::fs; +use tempfile::{TempDir, tempdir}; + +fn configs() -> (Value, Value) { + let encoder = json!({ + "model_type": "modernbert", "hidden_activation": "gelu", + "attention_bias": false, "mlp_bias": false, "norm_bias": false, + "hidden_size": 1024, "intermediate_size": 2624, + "num_attention_heads": 16, "num_hidden_layers": 28, + "vocab_size": 50368, "norm_eps": 0.00001, "local_attention": 128, + "layer_types": (["full_attention", "sliding_attention", "sliding_attention"] + .repeat(10)[..28]), + "rope_parameters": { + "full_attention": {"rope_type": "default", "rope_theta": 160000.0}, + "sliding_attention": {"rope_type": "default", "rope_theta": 10000.0} + } + }); + let agent = json!({ + "max_len": 512, "head_max_len": 192, "head_layers": 2, + "temperature": [0.5, 1.0, 2.0], "temperature_by_options": {"4": 1.5} + }); + (encoder, agent) +} + +fn load_config(encoder: &Value, agent: &Value) -> anyhow::Result { + let dir = tempdir()?; + fs::create_dir(dir.path().join("encoder"))?; + fs::write(dir.path().join("encoder/config.json"), encoder.to_string())?; + fs::write(dir.path().join("rl_agent_config.json"), agent.to_string())?; + Config::load(dir.path()) +} + +#[test] +fn config_accepts_supported_model_and_rejects_incompatible_inputs() { + let (encoder, agent) = configs(); + let loaded = load_config(&encoder, &agent).unwrap(); + assert_eq!(loaded.agent.temperature, [0.5, 1.0, 2.0]); + assert_eq!(loaded.encoder.hidden_size, 1024); + for (pointer, value) in [ + ("/model_type", json!("bert")), + ("/hidden_activation", json!("relu")), + ("/attention_bias", json!(true)), + ("/hidden_size", json!(768)), + ("/num_hidden_layers", json!(27)), + ("/layer_types/1", json!("full_attention")), + ("/rope_parameters/full_attention/rope_theta", json!(10000)), + ( + "/rope_parameters/sliding_attention/rope_type", + json!("linear"), + ), + ] { + let mut bad = encoder.clone(); + *bad.pointer_mut(pointer).unwrap() = value; + assert!(load_config(&bad, &agent).is_err(), "accepted {pointer}"); + } + let mut missing = encoder.clone(); + missing.as_object_mut().unwrap().remove("vocab_size"); + assert!(load_config(&missing, &agent).is_err()); + for (pointer, value) in [ + ("/max_len", json!(513)), + ("/head_layers", json!(3)), + ("/temperature", json!([1.0, 1.0])), + ("/temperature/0", json!(0.0)), + ("/temperature/1", json!(-1.0)), + ("/temperature/2", json!(1e100)), + ("/temperature_by_options/4", json!(0.0)), + ] { + let mut bad = agent.clone(); + *bad.pointer_mut(pointer).unwrap() = value; + assert!(load_config(&encoder, &bad).is_err(), "accepted {pointer}"); + } +} + +fn tensor(dtype: &str, shape: &[usize], data: &[u8]) -> (TempDir, Weights) { + let dir = tempdir().unwrap(); + let path = dir.path().join("model.safetensors"); + let mut header = json!({"w": { + "dtype": dtype, "shape": shape, "data_offsets": [0, data.len()] + }}) + .to_string(); + header.extend(std::iter::repeat_n(' ', (8 - header.len() % 8) % 8)); + let mut bytes = (header.len() as u64).to_le_bytes().to_vec(); + bytes.extend_from_slice(header.as_bytes()); + bytes.extend_from_slice(data); + fs::write(&path, bytes).unwrap(); + let weights = Weights::open(&path).unwrap(); + (dir, weights) +} + +#[test] +fn f32_conversion_preserves_zero_and_rounds_ties_to_even() { + // Independent IEEE encodings: signed zero, subnormals and half-way values. + let bits: [u32; 12] = [ + 0x00000000, 0x80000000, 0x33800000, 0x33000000, 0x33c00000, 0x3f801000, 0x3f803000, + 0x3f808000, 0x3f818000, 0x00010000, 0x00008000, 0x00018000, + ]; + let bytes: Vec<_> = bits.iter().flat_map(|x| x.to_le_bytes()).collect(); + let (_dir, weights) = tensor("F32", &[12], &bytes); + let actual: Vec<_> = weights + .f32("w", &[12]) + .unwrap() + .into_iter() + .map(f32::to_bits) + .collect(); + assert_eq!(actual, bits); + assert_eq!( + weights.f16("w", &[12]).unwrap(), + [0, 0x8000, 1, 0, 2, 0x3c00, 0x3c02, 0x3c04, 0x3c0c, 0, 0, 0] + ); + assert_eq!( + weights.bf16("w", &[12]).unwrap(), + [ + 0, 0x8000, 0x3380, 0x3300, 0x33c0, 0x3f80, 0x3f80, 0x3f80, 0x3f82, 1, 0, 2 + ] + ); +} + +#[test] +fn half_precision_inputs_decode_exactly() { + for (dtype, bits, expected) in [ + ( + "F16", + [0_u16, 0x8000, 1, 0x03ff, 0x0400, 0x3c00, 0xbc00], + [ + 0, 0x80000000, 0x33800000, 0x387fc000, 0x38800000, 0x3f800000, 0xbf800000, + ], + ), + ( + "BF16", + [0_u16, 0x8000, 1, 0x007f, 0x0080, 0x3f80, 0xbf80], + [ + 0, 0x80000000, 0x00010000, 0x007f0000, 0x00800000, 0x3f800000, 0xbf800000, + ], + ), + ] { + let bytes: Vec<_> = bits.iter().flat_map(|x| x.to_le_bytes()).collect(); + let (_dir, weights) = tensor(dtype, &[7], &bytes); + let actual: Vec<_> = weights + .f32("w", &[7]) + .unwrap() + .into_iter() + .map(f32::to_bits) + .collect(); + assert_eq!(actual, expected, "{dtype}"); + } +} + +#[test] +fn tensor_shape_names_and_dtype_are_checked() { + let (_dir, weights) = tensor("F32", &[2], &[0; 8]); + weights.validate_names(["w"]).unwrap(); + assert!(weights.validate_names([]).is_err()); + assert!(weights.validate_names(["w", "extra"]).is_err()); + assert!(weights.validate_names(["w", "w"]).is_err()); + assert!(weights.f32("w", &[1, 2]).is_err()); + assert!(weights.f32("missing", &[2]).is_err()); + let (_dir, integers) = tensor("I64", &[1], &[0; 8]); + assert!(integers.f32("w", &[1]).is_err()); + let dir = tempdir().unwrap(); + let path = dir.path().join("broken.safetensors"); + fs::write(&path, b"not a safetensors file").unwrap(); + assert!(Weights::open(&path).is_err()); +} + +#[test] +fn inventory_includes_gated_projection_and_legacy_buffer() { + let tensors = omni_laya::weights::checkpoint_tensors(); + let names: std::collections::HashSet<_> = tensors.iter().map(|t| &t.name).collect(); + assert_eq!(tensors.len(), 206); + assert_eq!(names.len(), tensors.len()); + let gated = tensors + .iter() + .find(|t| t.name == "encoder.layers.0.mlp.Wi.weight") + .unwrap(); + assert_eq!(gated.shape, [5248, 1024]); + let legacy = tensors.iter().find(|t| t.name == "temperature").unwrap(); + assert_eq!(legacy.shape, [3]); +} diff --git a/src/models/laya/tests/weights.rs b/src/models/laya/tests/weights.rs new file mode 100644 index 0000000..cc22ac4 --- /dev/null +++ b/src/models/laya/tests/weights.rs @@ -0,0 +1,93 @@ +use omni_laya::weights::Weights; +use sha2::{Digest, Sha256}; +#[test] +#[ignore = "requires LAYA_CHECKPOINT and LAYA_WEIGHT_ORACLE; CPU only"] +fn every_weight_conversion_matches_torch() { + let checkpoint = std::path::PathBuf::from(std::env::var_os("LAYA_CHECKPOINT").unwrap()); + let oracle = std::fs::read(std::env::var_os("LAYA_WEIGHT_ORACLE").unwrap()).unwrap(); + let rows: Vec = serde_json::from_slice(&oracle).unwrap(); + assert_eq!(rows.len(), 206, "oracle must cover the frozen checkpoint"); + let mut names = std::collections::HashSet::new(); + let weights = Weights::open(&checkpoint.join("model.safetensors")).unwrap(); + let inventory = omni_laya::weights::checkpoint_tensors(); + weights + .validate_names(inventory.iter().map(|t| t.name.as_str())) + .unwrap(); + let checkpoint_names: std::collections::HashSet<_> = + inventory.iter().map(|t| t.name.as_str()).collect(); + assert_eq!(checkpoint_names.len(), rows.len()); + for row in rows { + let name = row["name"].as_str().unwrap(); + assert!( + checkpoint_names.contains(name), + "unaccounted checkpoint tensor: {name}" + ); + assert!( + names.insert(name.to_owned()), + "duplicate oracle tensor: {name}" + ); + let shape: Vec = serde_json::from_value(row["shape"].clone()).unwrap(); + let spec = inventory.iter().find(|spec| spec.name == name).unwrap(); + assert_eq!(spec.shape, shape, "{name}: checkpoint shape"); + for dtype in ["f32", "f16", "bf16"] { + let bytes: Vec = match dtype { + "f32" => weights + .f32(name, &shape) + .unwrap() + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect(), + "f16" => weights + .f16(name, &shape) + .unwrap() + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect(), + _ => weights + .bf16(name, &shape) + .unwrap() + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect(), + }; + assert_eq!( + format!("{:x}", Sha256::digest(&bytes)), + row[dtype].as_str().unwrap(), + "{name} {dtype}" + ); + } + } +} + +#[test] +#[ignore = "requires LAYA_CHECKPOINT; CPU only"] +fn checkpoint_inventory_covers_legacy_temperature() { + let checkpoint = std::path::PathBuf::from(std::env::var_os("LAYA_CHECKPOINT").unwrap()); + let weights = Weights::open(&checkpoint.join("model.safetensors")).unwrap(); + let inventory = omni_laya::weights::checkpoint_tensors(); + weights + .validate_names(inventory.iter().map(|t| t.name.as_str())) + .unwrap(); + assert!( + weights + .validate_names( + inventory + .iter() + .filter(|t| t.name != "temperature") + .map(|t| t.name.as_str()) + ) + .is_err() + ); + assert!( + weights + .validate_names(inventory.iter().map(|t| t.name.as_str()).chain(["unknown"])) + .is_err() + ); + let legacy = weights.f32("temperature", &[3]).unwrap(); + assert_eq!(legacy, [1.0, 1.0, 1.0]); + let config = omni_laya::config::Config::load(&checkpoint).unwrap(); + assert_ne!( + legacy, config.agent.temperature, + "legacy buffer is not the fitted calibration source" + ); +} From c8cebc4db70bea6855fac1c2913f335f4df23f95 Mon Sep 17 00:00:00 2001 From: xiaoyu-xyz Date: Mon, 28 Sep 2026 16:19:51 +0800 Subject: [PATCH 2/3] [laya] Name the tensors that differ in an inventory error The inventory mismatch reported only that the sets were unequal, so the most likely failure -- pointing the engine at a checkpoint that is not the frozen one -- said nothing about which tensor was wrong or in which direction. Both sets were already in scope. Report the expected-only names as "missing" and the checkpoint-only names as "unexpected", sorted, so the message is deterministic. The neighbouring errors already name their tensor (duplicate expected tensor, shape mismatch, unsupported dtype); this was the one that did not. Adds a CPU test over the synthetic safetensors fixture that asserts both directions and the ordering. fmt, clippy -D warnings, and the workspace tests pass; the checkpoint test that exercises this path still passes against the frozen checkpoint. --- src/models/laya/src/weights.rs | 14 ++++++++++---- src/models/laya/tests/checkpoint.rs | 25 +++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 4 deletions(-) diff --git a/src/models/laya/src/weights.rs b/src/models/laya/src/weights.rs index 69d3be6..d3e5a1f 100644 --- a/src/models/laya/src/weights.rs +++ b/src/models/laya/src/weights.rs @@ -17,10 +17,16 @@ impl Weights { } let names = tensors.names(); let actual: std::collections::BTreeSet<_> = names.into_iter().collect(); - ensure!( - expected == actual, - "tensor inventory does not match checkpoint" - ); + if expected != actual { + let joined = |set: std::collections::BTreeSet<&str>| { + set.into_iter().collect::>().join(", ") + }; + let missing = joined(expected.difference(&actual).copied().collect()); + let extra = joined(actual.difference(&expected).copied().collect()); + anyhow::bail!( + "tensor inventory does not match checkpoint: missing [{missing}], unexpected [{extra}]" + ); + } Ok(()) } /// The checkpoint must remain immutable while the mapping exists. diff --git a/src/models/laya/tests/checkpoint.rs b/src/models/laya/tests/checkpoint.rs index 2b3493b..5eb8059 100644 --- a/src/models/laya/tests/checkpoint.rs +++ b/src/models/laya/tests/checkpoint.rs @@ -164,6 +164,31 @@ fn tensor_shape_names_and_dtype_are_checked() { assert!(Weights::open(&path).is_err()); } +#[test] +fn inventory_mismatch_names_the_tensors_that_differ() { + let (_dir, weights) = tensor("F32", &[2], &[0; 8]); + // The checkpoint holds exactly "w". expected = what the caller asked for, + // actual = what the file has, so each direction has its own side. + let message = weights + .validate_names(["w", "absent"]) + .unwrap_err() + .to_string(); + assert!(message.contains("missing [absent]"), "{message}"); + assert!(message.contains("unexpected []"), "{message}"); + // Asking for nothing leaves the file's tensor unaccounted for. + let message = weights.validate_names([]).unwrap_err().to_string(); + assert!(message.contains("missing []"), "{message}"); + assert!(message.contains("unexpected [w]"), "{message}"); + let message = weights + .validate_names(["w", "absent", "also_absent"]) + .unwrap_err() + .to_string(); + assert!( + message.contains("missing [absent, also_absent]"), + "{message}" + ); +} + #[test] fn inventory_includes_gated_projection_and_legacy_buffer() { let tensors = omni_laya::weights::checkpoint_tensors(); From 219808b315c2c36b6bcbf7f640de4b660bf1d0a3 Mon Sep 17 00:00:00 2001 From: linear3735 Date: Thu, 1 Oct 2026 17:40:37 +0800 Subject: [PATCH 3/3] Move Laya checkpoint tests to the root test directory --- src/models/laya/Cargo.toml | 8 ++++++++ {src/models/laya/tests => tests/laya}/checkpoint.rs | 0 {src/models/laya/tests => tests/laya}/weights.rs | 0 3 files changed, 8 insertions(+) rename {src/models/laya/tests => tests/laya}/checkpoint.rs (100%) rename {src/models/laya/tests => tests/laya}/weights.rs (100%) diff --git a/src/models/laya/Cargo.toml b/src/models/laya/Cargo.toml index d9e73e8..0eba920 100644 --- a/src/models/laya/Cargo.toml +++ b/src/models/laya/Cargo.toml @@ -15,3 +15,11 @@ serde_json = "1" [dev-dependencies] sha2 = "0.10" tempfile = "3" + +[[test]] +name = "checkpoint" +path = "../../../tests/laya/checkpoint.rs" + +[[test]] +name = "weights" +path = "../../../tests/laya/weights.rs" diff --git a/src/models/laya/tests/checkpoint.rs b/tests/laya/checkpoint.rs similarity index 100% rename from src/models/laya/tests/checkpoint.rs rename to tests/laya/checkpoint.rs diff --git a/src/models/laya/tests/weights.rs b/tests/laya/weights.rs similarity index 100% rename from src/models/laya/tests/weights.rs rename to tests/laya/weights.rs