Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,18 @@ jobs:
run: cargo clippy --workspace --locked --all-targets -- -D warnings
- name: Test
run: cargo test --workspace --locked
- name: Official Laya packing parity
env:
LAYA_TOKENIZER: ${{ runner.temp }}/laya-tokenizer.json
LAYA_PACKING_ORACLE: ${{ runner.temp }}/laya-packing.json
run: |
curl --fail --location --retry 3 \
https://huggingface.co/convaiinnovations/laya/resolve/55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851/tokenizer/tokenizer.json \
--output "$LAYA_TOKENIZER"
curl --fail --location --retry 3 \
https://raw.githubusercontent.com/linear3735/system1-omni/5e4dd4215c925ebd93bb9ce4097b27bd6375f7c0/recipe/laya/native/packing-golden.json \
--output "$LAYA_PACKING_ORACLE"
cargo test --locked -p omni-laya --test packing -- --ignored
- name: Build
run: cargo build --workspace --release --locked

Expand Down
68 changes: 67 additions & 1 deletion Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

113 changes: 113 additions & 0 deletions recipe/laya/native/export_decisions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,113 @@
"""CPU decode references from Laya 0.3.20; synthetic logits, no model loading."""

import argparse
import hashlib
import importlib.metadata
import json
import math
import platform
from pathlib import Path

import laya.agent
import laya.common
import torch


def question(qid, kind, criteria=None):
return {"id": qid, "kind": kind, "criteria": criteria}


def case(name, questions, logits, action_logits, config=None):
row = dict(name=name, questions=questions, logits=logits, action_logits=action_logits)
if config is not None:
row["config"] = config
return row


parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("checkpoint", type=Path)
parser.add_argument("output", type=Path)
args = parser.parse_args()
if importlib.metadata.version("laya") != "0.3.20":
raise SystemExit("the reference requires laya==0.3.20")
torch.set_num_threads(1)
cfg_file = args.checkpoint / "rl_agent_config.json"
cfg = json.loads(cfg_file.read_text())
config = {key: cfg[key] for key in ("temperature", "temperature_by_options")}
cases = []
for k in (1, 2, 3, 5, 6, 10, 11):
criteria = {f"option-{i}": None for i in range(k)}
cases.append(case(
f"choice-{k}", [question("q", "choice", criteria)],
[[i * 0.375 - 1.0 for i in range(k)]], [[-0.75, 0.5]],
))
cases += [
case("choice-first-tie", [question("q", "choice", {"z": None, "a": None})],
[[2.0, 2.0]], [[0.0, 0.0]]),
case("score-single", [question("q", "score", ["only"])], [[-5.0]], [[1.0, -1.0]]),
case("score-legend", [question("q", "score", ["low", {"description": "middle"}, 7])],
[[0.25, 1.5, -0.5]], [[0.3, -0.8]]),
case("noul-extremes", [question("false", "noul"), question("true", "noul")],
[[1e30, -1e30], [-1e30, 1e30]], [[1e30, -1e30], [-1e30, 1e30]]),
case("temperature-clamps", [question("choice", "choice", {"a": None, "b": None}),
question("score", "score", ["low", "mid", "high"])],
[[-0.4, 0.6], [-0.4, 0.6, 1.6]], [[-1.0, 1.0], [2.0, -2.0]],
{"temperature": [0.01, 80.0, 1.0], "temperature_by_options": {}}),
case("mixed-order", [question("z", "score", ["low", "mid", "high"]),
question("a", "choice", {"later": "", "earlier": ""}),
question("m", "noul")],
[[-0.5, 1.0, 0.75], [0.5, -0.5], [0.0, 0.0]],
[[0.0, 1.0], [2.0, -1.0], [-2.0, 0.0]]),
case("score-bucket-override", [question("q", "score", list(range(6)))],
[[-2.0, -1.0, 0.0, 0.5, 1.0, 3.0]], [[-0.25, 0.75]],
{"temperature": [1.0, 1.0, 1.0], "temperature_by_options": {"score:6-10": 4.0}}),
case("rounding-boundaries", [question(q, "choice", {"a": None, "b": None})
for q in ("prob-low", "prob-high", "entropy-low", "entropy-high")],
[[math.log(p / (1 - p)), 0.0] for p in (0.800049, 0.800051)]
+ [[1.3862165, 0.0], [1.3862353, 0.0]], [[-0.25, 0.75]] * 4,
{"temperature": [1.0, 1.0, 1.0], "temperature_by_options": {}}),
case("empty", [], [], []),
]
rounding_probe = case(
"fp32-reduction-boundary", [question("q", "choice", {str(i): None for i in range(16)})],
[[-1.7411574125289917, -0.19089631736278534, -0.6029739379882812, -0.8184939026832581,
0.16066476702690125, -0.4026077389717102, 0.343989759683609, -0.600969135761261,
0.8842262029647827, -0.26977965235710144, -0.7890094518661499, 0.2582162916660309,
0.85430908203125, -0.11924569308757782, 0.9091809988021851, -0.00020837262854911387]],
[[0.0, 0.0]], {"temperature": [1.0, 1.0, 1.0], "temperature_by_options": {}},
)
for row in cases + [rounding_probe]:
current = row.get("config", config)
agent = laya.agent.Agent.__new__(laya.agent.Agent)
agent.temperature = [laya.common.clamp_temperature(t) for t in current["temperature"]]
agent.temperature_by_options = {
key: laya.common.clamp_temperature(t) for key, t in current["temperature_by_options"].items()
}
agent.lang_temperatures = {}
logits = torch.full((len(row["logits"]), max(map(len, row["logits"]), default=0)), -1e4)
for i, values in enumerate(row["logits"]):
logits[i, :len(values)] = torch.tensor(values, dtype=torch.float32)
actions = torch.tensor(row["action_logits"], dtype=torch.float32).reshape(-1, 2)
agent._infer = lambda batch: (logits, actions)
logits_np, action_probs = agent._forward(None)
internal = {q["id"]: {"t": q["kind"], "crit": q["criteria"]} for q in row["questions"]}
items = [{"markers": list(range(len(values)))} for values in row["logits"]]
row["answers"] = agent._decode_answers(logits_np, action_probs, items, list(internal), internal, 0)

reference = {
"python": platform.python_version(),
**{name: importlib.metadata.version(name) for name in ("laya", "torch", "numpy")},
"source_sha256": {
Path(module.__file__).name: hashlib.sha256(Path(module.__file__).read_bytes()).hexdigest()
for module in (laya.agent, laya.common)
},
"config_sha256": hashlib.sha256(cfg_file.read_bytes()).hexdigest(),
"scope": "Synthetic boundary logits; official CPU decode parity, not model quality or performance.",
}
args.output.parent.mkdir(parents=True, exist_ok=True)
dump = lambda value: json.dumps(value, ensure_ascii=False, allow_nan=False, separators=(",", ":"))
args.output.write_text(
'{"reference":' + dump(reference) + ',"config":' + dump(config) + ',"cases":[\n'
+ ",\n".join(dump(row) for row in cases) + '\n],"rounding_probe":' + dump(rounding_probe) + "}\n"
)
print(f"{len(cases)} exact cases and one rounding probe written to {args.output}")
2 changes: 1 addition & 1 deletion src/models/cua_s1/native/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ libloading = "0.8"
memmap2 = "0.9.9"
safetensors = "0.8.0"
serde = "1"
serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order"] }
serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order", "raw_value"] }
# the onig regex backend, as in the Python tokenizers wheel
tokenizers = { version = "=0.22.2", default-features = false, features = ["onig"] }
tokio = { version = "1.49.0", features = ["macros", "net", "rt-multi-thread", "sync"] }
64 changes: 53 additions & 11 deletions src/models/cua_s1/native/src/json.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,14 +10,14 @@ use std::fmt::Write as _;
use std::io;

use serde::Serialize;
use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor};
use serde::de::{self, DeserializeSeed, Deserializer, MapAccess, SeqAccess, Visitor};
use serde_json::{Map, Number, Value};

/// Decode a request body into its top-level object; the error is the 400 message.
pub fn parse(raw: &[u8]) -> Result<Map<String, Value>, String> {
let mut de = serde_json::Deserializer::from_slice(raw);
let value = de
.deserialize_any(NoDuplicates)
let value = NoDuplicates(0)
.deserialize(&mut de)
.and_then(|v| de.end().map(|()| v))
.map_err(|e| format!("request body is not valid JSON: {e}"))?;
match value {
Expand All @@ -27,16 +27,37 @@ pub fn parse(raw: &[u8]) -> Result<Map<String, Value>, String> {
}

/// Builds a `Value` like serde_json does, but fails on a repeated key.
struct NoDuplicates;
struct NoDuplicates(usize);

impl<'de> de::Deserialize<'de> for Wrapped {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
d.deserialize_any(NoDuplicates).map(Wrapped)
impl<'de> DeserializeSeed<'de> for NoDuplicates {
type Value = Value;

fn deserialize<D: Deserializer<'de>>(self, d: D) -> Result<Value, D::Error> {
let raw = <&serde_json::value::RawValue as de::Deserialize>::deserialize(d)?;
let text = raw.get();
// Another workspace member may enable arbitrary_precision. Read number
// tokens directly so its private map encoding cannot become user data.
if matches!(text.as_bytes()[0], b'-' | b'0'..=b'9') {
if text != "-0" {
if let Ok(n) = text.parse::<i64>() {
return self.visit_i64(n);
}
if let Ok(n) = text.parse::<u64>() {
return self.visit_u64(n);
}
}
let n = serde_json::from_str::<f64>(text).map_err(de::Error::custom)?;
return self.visit_f64(n);
}
if self.0 >= 127 && matches!(text.as_bytes()[0], b'{' | b'[') {
return Err(de::Error::custom("recursion limit exceeded"));
}
serde_json::Deserializer::from_str(text)
.deserialize_any(self)
.map_err(de::Error::custom)
}
}

struct Wrapped(Value);

impl<'de> Visitor<'de> for NoDuplicates {
type Value = Value;

Expand Down Expand Up @@ -68,15 +89,15 @@ impl<'de> Visitor<'de> for NoDuplicates {
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Value, A::Error> {
let mut items = Vec::new();
while let Some(Wrapped(v)) = seq.next_element()? {
while let Some(v) = seq.next_element_seed(NoDuplicates(self.0 + 1))? {
items.push(v);
}
Ok(Value::Array(items))
}
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Value, A::Error> {
let mut obj = Map::new();
while let Some(key) = map.next_key::<String>()? {
let Wrapped(v) = map.next_value()?;
let v = map.next_value_seed(NoDuplicates(self.0 + 1))?;
if obj.contains_key(&key) {
return Err(de::Error::custom(format_args!(
"duplicate key {}",
Expand Down Expand Up @@ -107,6 +128,14 @@ pub fn dumps(value: &Value) -> String {
struct PyFormatter;

impl serde_json::ser::Formatter for PyFormatter {
fn write_number_str<W: ?Sized + io::Write>(&mut self, w: &mut W, n: &str) -> io::Result<()> {
if n.contains(['.', 'e', 'E']) {
let x = serde_json::from_str::<f64>(n).map_err(io::Error::other)?;
self.write_f64(w, x)
} else {
w.write_all(n.as_bytes())
}
}
fn begin_array_value<W: ?Sized + io::Write>(
&mut self,
w: &mut W,
Expand Down Expand Up @@ -214,6 +243,19 @@ mod tests {
);
}

#[test]
fn number_encoding_does_not_change_user_objects() {
let input =
r#"{"a": {"$serde_json::private::Number": "1.5"}, "b": -0, "c": 18446744073709551616}"#;
let value = Value::Object(parse(input.as_bytes()).unwrap());
assert_eq!(value["a"]["$serde_json::private::Number"], "1.5");
assert_eq!(
dumps(&value),
r#"{"a": {"$serde_json::private::Number": "1.5"}, "b": -0.0, "c": 1.8446744073709552e+19}"#
);
assert!(err(r#"{"a": {"x": 1, "\u0078": 2}}"#).contains("duplicate key"));
}

#[test]
fn rejects_what_the_contract_rejects() {
assert_eq!(err("[]"), "request body must be a JSON object");
Expand Down
15 changes: 14 additions & 1 deletion src/models/laya/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@ half = "2"
memmap2 = "0.9"
safetensors = "0.6"
serde = { version = "1", features = ["derive"] }
serde_json = "1"
serde_json = { version = "1", features = ["preserve_order", "arbitrary_precision"] }
tokenizers = { version = "0.23.2", default-features = false, features = ["fancy-regex"] }

[dev-dependencies]
sha2 = "0.10"
Expand All @@ -23,3 +24,15 @@ path = "../../../tests/laya/checkpoint.rs"
[[test]]
name = "weights"
path = "../../../tests/laya/weights.rs"

[[test]]
name = "preprocess"
path = "../../../tests/laya/preprocess.rs"

[[test]]
name = "packing"
path = "../../../tests/laya/packing.rs"

[[test]]
name = "decision"
path = "../../../tests/laya/decision.rs"
Loading
Loading