Skip to content
Closed
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.

69 changes: 69 additions & 0 deletions recipe/laya/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -56,3 +56,72 @@ if the worker requires a bearer token.

See the [frontend documentation](../../src/frontend/README.md) for configuration
and transport behavior.

## Native CPU packing check

The `omni-laya` preprocessor packs English Laya 0.3.20 requests without weights
or a GPU.

```sh
cargo test --locked -p omni-laya --test preprocess
```

Native callers use `Request::from_json(&str)` for a single top-level JSON request,
or `Request::from_value(Value)` for an existing structured value. `Request`
retains its public fields and `Serialize`; it does not implement generic
`Deserialize`. The JSON entry checks the complete request against serde_json's
default nesting limit. The value entry preserves existing nested values without
reparsing. Both preserve literal private Number/RawValue object keys and reject
unknown request fields.

The normal tests cover validation, JSON rendering, question and option order,
and truncation with a small tokenizer:
Pass raw JSON directly to `from_json`.

For the official 17-case comparison, use the same pinned inputs as CPU CI.
The test checks both files by SHA-256 before comparing:

```sh
LAYA_PACKING_DIR=$(mktemp -d)
export LAYA_TOKENIZER="$LAYA_PACKING_DIR/tokenizer.json"
export LAYA_PACKING_ORACLE="$LAYA_PACKING_DIR/packing-golden.json"
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
```

Existing copies of these pinned files can be supplied through `LAYA_TOKENIZER`
and `LAYA_PACKING_ORACLE` instead. The comparison covers every token, marker,
question type, row length, question order and usage count; it excludes backend
padding and bucket dimensions. The [reference generator and inputs](https://github.com/linear3735/system1-omni/tree/5e4dd4215c925ebd93bb9ce4097b27bd6375f7c0/recipe/laya/native)
use `laya==0.3.20`. Packing parity does not measure model quality or execute
native model inference.

## Native decoder validation

Run the CPU decoding checks without Python, weights or a GPU:

```sh
cargo test --locked -p omni-laya --test decision
```

The checked-in reference covers all three question types, temperatures, ordering,
ties and rounding boundaries. Its 16 cases require exact rounded answers; a
separate FP32 reduction probe allows one displayed decimal unit for probabilities.
These fixed logits check decoding, not model quality or GPU execution.

To regenerate the reference, use `laya==0.3.20`, `torch==2.14.0` and
`numpy==2.5.3`, matching the recorded fixture. Supply the English checkpoint
`convaiinnovations/laya` at revision `55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851`:

```sh
export LAYA_CHECKPOINT=/path/to/laya/snapshot
python recipe/laya/native/export_decisions.py "$LAYA_CHECKPOINT" /path/to/new-decisions.json
```

The generator reads only `rl_agent_config.json` and calls Laya's official CPU
decoding functions. The output records source hashes and package versions.
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"] }
Loading
Loading