diff --git a/Cargo.lock b/Cargo.lock index 16864a1..4b9adda 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -119,6 +119,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" @@ -169,7 +178,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.3.1", "rand_core 0.10.1", ] @@ -188,6 +197,15 @@ dependencies = [ "static_assertions", ] +[[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" @@ -228,6 +246,16 @@ 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 = "darling" version = "0.20.11" @@ -303,6 +331,16 @@ dependencies = [ "syn 2.0.119", ] +[[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" @@ -435,6 +473,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" @@ -903,7 +951,7 @@ dependencies = [ "half", "libloading", "memmap2", - "safetensors", + "safetensors 0.8.0", "serde", "serde_json", "tokenizers", @@ -920,6 +968,20 @@ dependencies = [ "tokio", ] +[[package]] +name = "omni-laya" +version = "0.1.0" +dependencies = [ + "anyhow", + "half", + "memmap2", + "safetensors 0.6.2", + "serde", + "serde_json", + "sha2", + "tempfile", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -1318,6 +1380,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 = "safetensors" version = "0.8.0" @@ -1398,6 +1470,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" @@ -1718,6 +1801,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" diff --git a/Cargo.toml b/Cargo.toml index 8c100b4..661c83e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend", "src/models/cua_s1/native"] +members = ["src/frontend", "src/models/cua_s1/native", "src/models/laya"] resolver = "3" diff --git a/README.md b/README.md index da46691..09d9d90 100644 --- a/README.md +++ b/README.md @@ -49,7 +49,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 and the Cua-S1 native worker are Cargo workspace members. The other model and backend directories currently document planned work; they do not prescribe process boundaries. +The frontend, Cua-S1 native worker and Laya checkpoint reader are Cargo workspace members. The other model and backend directories currently document planned work; they do not prescribe process boundaries. ## Supported models @@ -57,7 +57,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 | | Cua-S1 4B 0.2 (`text` adapter) | [Python worker](recipe/cua_s1/text.md); [native worker](recipe/cua_s1/native.md), CUDA, run on sm_89 | 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..0eba920 --- /dev/null +++ b/src/models/laya/Cargo.toml @@ -0,0 +1,25 @@ +[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" + +[[test]] +name = "checkpoint" +path = "../../../tests/laya/checkpoint.rs" + +[[test]] +name = "weights" +path = "../../../tests/laya/weights.rs" 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..d3e5a1f --- /dev/null +++ b/src/models/laya/src/weights.rs @@ -0,0 +1,163 @@ +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(); + 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. + 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/tests/laya/checkpoint.rs b/tests/laya/checkpoint.rs new file mode 100644 index 0000000..5eb8059 --- /dev/null +++ b/tests/laya/checkpoint.rs @@ -0,0 +1,205 @@ +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_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(); + 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/tests/laya/weights.rs b/tests/laya/weights.rs new file mode 100644 index 0000000..cc22ac4 --- /dev/null +++ b/tests/laya/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" + ); +}