gdelatournelle commited on
Commit
49b92fb
·
verified ·
1 Parent(s): cc94284

Mirror laya-onnx source tree + model card

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
.gitignore ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ .venv/
2
+ .upstream/
3
+ .cache/
4
+ models/
5
+ onnx/
6
+ *.safetensors
7
+ *.onnx
8
+ *.onnx.data
9
+ __pycache__/
10
+ .pytest_cache/
11
+ .ruff_cache/
12
+ *.egg-info/
13
+ build/
14
+ dist/
15
+ .DS_Store
16
+ .env
AGENTS.md ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ # laya-onnx — notes for coding agents
2
+
3
+ Python ≥ 3.11. ONNX Runtime backend for Laya typed System-1 decisions.
4
+
5
+ - Install: `pip install -e ".[demo]"` (`openvino`, `ultrafast`, `export`).
6
+ - CLI: `laya-onnx predict|convert|optimize`, `laya-onnx-snake`, `laya-onnx-ultrafast`.
7
+ - Deterministic: `load(..., deterministic=True)` or `--deterministic` (threads=1, pad off, `predict_argmax`, heuristic TYPE_TEXT).
8
+ - Snake: https://raw.githack.com/Geoking2104/laya-onnx/main/examples/snake.html
9
+ - Ultrafast: `docs/ULTRAFAST.md`, `--fixture examples/ultrafast_page.json`.
10
+ - Do not commit `.venv` or `*.onnx`.
GROK.md ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ # Grok snapshot
2
+
3
+ This tree is the Grok-maintained ONNX Runtime port of Laya (KayrosLab / Geoking2104).
4
+
5
+ - Runtime: `laya_onnx/agent.py` (`InferenceSession`, cpu / openvino / cuda)
6
+ - Export: `laya_onnx/convert.py` (torch → ONNX opset 17)
7
+ - Demo: https://raw.githack.com/Geoking2104/laya-onnx/main/examples/snake.html
8
+ - Onboard: `ONBOARD.md`
9
+
10
+ Do not commit `.venv` or Hub weights.
LICENSE ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ Licensed under the Apache License, Version 2.0 (the "License");
6
+ you may not use this file except in compliance with the License.
7
+ You may obtain a copy of the License at
8
+
9
+ http://www.apache.org/licenses/LICENSE-2.0
10
+
11
+ Unless required by applicable law or agreed to in writing, software
12
+ distributed under the License is distributed on an "AS IS" BASIS,
13
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ See the License for the specific language governing permissions and
15
+ limitations under the License.
NOTICE ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ laya-onnx
2
+ Copyright 2026 Geoffroy de La Tournelle / KayrosLab and contributors
3
+
4
+ Derived from:
5
+ - Laya (Apache-2.0) https://github.com/NandhaKishorM/laya
6
+ - laya-mlx (Apache-2.0) https://github.com/mizorewww/laya-mlx
7
+
8
+ Model weights: Convai Innovations, Apache-2.0, downloaded separately.
9
+ Default ONNX bundle: https://huggingface.co/receptron/laya-onnx
ONBOARD.md ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Onboard your AI agent to run laya-onnx
2
+
3
+ Paste this file (or the one-liner in the README) into Claude Code, Codex, Cursor, Copilot, or any coding agent. The agent must follow every step. Safe to re-run.
4
+
5
+ ## Goal
6
+
7
+ Install `laya-onnx` in the current workspace so the human can call `laya-onnx predict` and `laya-onnx-snake` on CPU (Intel i7 / 16 GB class). Do not download multi-GB weights unless the human asks.
8
+
9
+ ## Steps
10
+
11
+ 1. Confirm Python ≥ 3.11 (`python3 --version`). If missing, tell the human how to install it and stop.
12
+ 2. If this directory is not already the repo:
13
+ ```bash
14
+ git clone https://github.com/Geoking2104/laya-onnx.git
15
+ cd laya-onnx
16
+ ```
17
+ 3. Create a venv and install extras:
18
+ ```bash
19
+ python3 -m venv .venv
20
+ source .venv/bin/activate
21
+ pip install -U pip
22
+ pip install -e ".[demo,dev]"
23
+ ```
24
+ On Intel CPU also: `pip install -e ".[openvino]"` (ignore failure if wheels are missing).
25
+ Only if the human asked to export graphs: `pip install -e ".[export]"`.
26
+ 4. Verify:
27
+ ```bash
28
+ laya-onnx --help
29
+ PYTHONPATH=. python -c "from laya_onnx import load; print('import ok')"
30
+ ```
31
+ 5. Write `AGENTS.md` at the repo root if it does not exist. Point at README, `examples/`, `laya-onnx predict`, `optimize --precision int8`, and `examples/snake.html`.
32
+ 6. Optional MCP: this repo does **not** ship a hosted MCP server yet. Do not invent RunPod keys. If the human wants GPU later: `npx skills add runpod/runpod-plugins-official`.
33
+ 7. Report back Python version, install path, working commands, and that weights live on Hugging Face (`receptron/laya-onnx`).
34
+
35
+ Do not commit `.venv`, ONNX weights, or API keys.
README.md ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: onnxruntime
4
+ pipeline_tag: text-classification
5
+ base_model: convaiinnovations/laya
6
+ tags:
7
+ - onnx
8
+ - onnxruntime
9
+ - decision-model
10
+ - typed-decisions
11
+ - calibration
12
+ - cpu
13
+ - intel
14
+ - laya
15
+ ---
16
+
17
+ # laya-onnx
18
+
19
+ ONNX Runtime for [Laya](https://huggingface.co/convaiinnovations/laya) — typed
20
+ System-1 decisions, no generated tokens. PC sibling of
21
+ [laya-coreml](https://github.com/mizorewww/laya-coreml); same `choice` / `score` /
22
+ `noul` contract.
23
+
24
+ - **Source code:** <https://github.com/Geoking2104/laya-onnx>
25
+ - **Weights (not hosted in this repo, ≈1.6 GB fp32):** <https://huggingface.co/receptron/laya-onnx>
26
+ - **Benchmark · calibration · checksums:** <https://huggingface.co/datasets/gdelatournelle/laya-onnx-bench>
27
+
28
+ This Hub repository mirrors the project's source tree (runtime, CLI, Snake demo,
29
+ benchmarks, tests, docs). Model weights are **not** included here.
30
+
31
+ ## Model description
32
+
33
+ | | |
34
+ | --- | --- |
35
+ | Task | Typed decision — `choice`, `score`, `noul` |
36
+ | Output | per-question label/score + probability; no generated tokens (`output_tokens: 0`) |
37
+ | Backbone | Laya (ModernBERT-style encoder + RL decision head) |
38
+ | Runtime | ONNX Runtime — CPU, Intel OpenVINO, CUDA |
39
+ | Weights | `receptron/laya-onnx` (~1.6 GB, fp32) |
40
+ | Precision | fp32; `laya-onnx optimize --precision int8` on Intel CPU |
41
+ | License | Apache-2.0 (following Laya) |
42
+
43
+ `laya-onnx` answers **typed questions** about a state and returns probabilities —
44
+ it never writes text. `predict()` batches every question for a state into one
45
+ call and decides by argmax; `predict_argmax()` ignores the calibrated temperature
46
+ for a reproducible path.
47
+
48
+ ## Install
49
+
50
+ ```bash
51
+ git clone https://github.com/Geoking2104/laya-onnx && cd laya-onnx
52
+ python3 -m venv .venv && source .venv/bin/activate # Windows: .venv\Scripts\activate
53
+ pip install -U pip && pip install -e ".[demo,dev]"
54
+ pip install -e ".[openvino]" # Intel CPU
55
+ pip install -e ".[ultrafast]" # optional: Playwright DOM loop
56
+ ```
57
+
58
+ Model weights are not in git. The first `load("receptron/laya-onnx")` (or CLI
59
+ run) fetches them from the Hub.
60
+
61
+ ## Usage
62
+
63
+ ```python
64
+ from laya_onnx import load
65
+
66
+ agent = load("receptron/laya-onnx", providers="cpu")
67
+ print(agent.predict(
68
+ "The customer requests a refund of a duplicate payment.",
69
+ {"refund": {"type": "noul", "instructions": "Does the customer request a refund?"}},
70
+ ))
71
+ ```
72
+
73
+ Deterministic path (single thread, no padding, temperature ignored):
74
+
75
+ ```python
76
+ agent = load("receptron/laya-onnx", providers="cpu", deterministic=True)
77
+ print(agent.predict_argmax(state, questions))
78
+ ```
79
+
80
+ ## CLI
81
+
82
+ ```bash
83
+ laya-onnx predict --state "..." --questions q.json --model ./onnx
84
+ laya-onnx convert --model convaiinnovations/laya --output onnx --precision fp32
85
+ laya-onnx optimize ./onnx --precision int8 # Intel CPU
86
+ laya-onnx verify --model ./onnx # SHA-256 check vs bundled manifest
87
+ laya-onnx-ultrafast --dry-run --fixture examples/ultrafast_page.json --goal "..."
88
+ laya-onnx-snake --model ./onnx # terminal Snake
89
+ ```
90
+
91
+ ## Published benchmark
92
+
93
+ CPU run (Windows 11, Intel i7-1255U, `onnxruntime`), bundle `receptron/laya-onnx`
94
+ @ `68f27dfe`:
95
+
96
+ | question type | n | accuracy | MAE | ECE |
97
+ | --- | ---: | ---: | ---: | ---: |
98
+ | choice | 4 | 1.00 | – | 0.456 |
99
+ | noul | 4 | 1.00 | – | 0.129 |
100
+ | score | 2 | 0.50 | 0.530 | 0.112 |
101
+ | **overall** | **10** | **0.90** | – | **0.256** |
102
+
103
+ Latency per `predict()` call: mean 2808 ms / p50 1760 ms / p95 7418 ms. Small
104
+ self-check set (`benchmarks/eval/decisions.jsonl`), **not** a leaderboard — see
105
+ the [bench dataset](https://huggingface.co/datasets/gdelatournelle/laya-onnx-bench)
106
+ for the method and caveats.
107
+
108
+ ## Calibration
109
+
110
+ Confidence is calibrated with per-question-type temperatures (plus per
111
+ option-count buckets); `load()` warns and clamps any temperature outside the
112
+ accepted range. Calibration is reported as ECE in the benchmark above.
113
+
114
+ ## Verify a download
115
+
116
+ ```bash
117
+ laya-onnx verify --model ~/.cache/huggingface/hub/laya-onnx-bundles/receptron--laya-onnx
118
+ ```
119
+
120
+ Hashes every bundle file against the SHA-256 manifest shipped in
121
+ `laya_onnx/checksums/`.
122
+
123
+ ## Repository map
124
+
125
+ - `laya_onnx/` — runtime (`agent`, `hub`, `inputs`, `tokenizer`, `convert`, `optimize`, `eval`, `verify`, `cli`)
126
+ - `laya_onnx/snake/` · `laya_onnx/ultrafast/` — demos (terminal Snake; Playwright DOM loop)
127
+ - `benchmarks/` — `evaluate.py`, `pc_benchmark.py`, `eval/decisions.jsonl`, published `results.*`
128
+ - `docs/` — `DETERMINISTIC.md`, `ULTRAFAST.md`
129
+ - `examples/` — quickstart, sample state/questions, `snake.html` browser mock
130
+
131
+ ## Attribution
132
+
133
+ Apache-2.0. Derived from Laya (Apache-2.0) and laya-mlx. Ultrafast loop design
134
+ from Browser Use. See `LICENSE` and `NOTICE`.
benchmarks/eval/decisions.jsonl ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {"id": "dept-refund", "state": "Customer email: 'You charged me twice for invoice #4411. Please refund the duplicate payment today or we will cancel our plan.'", "questions": {"department": {"type": "choice", "instructions": "Which department should handle this email?", "criteria": {"billing": "invoices, payments, refunds", "technical": "bugs, outages, system errors", "sales": "pricing, new contracts", "other": "everything else"}}, "refund": {"type": "noul", "instructions": "Does the customer ask for money back?"}}, "gold": {"department": "billing", "refund": true}}
2
+ {"id": "dept-bug", "state": "Support ticket: 'The desktop app crashes every time I open a shared document since the 3.2 update.'", "questions": {"department": {"type": "choice", "instructions": "Which department should handle this ticket?", "criteria": {"billing": "invoices, payments, refunds", "technical": "bugs, outages, system errors", "sales": "pricing, new contracts", "other": "everything else"}}, "system_error": {"type": "noul", "instructions": "Is the user reporting a software bug or system error?"}}, "gold": {"department": "technical", "system_error": true}}
3
+ {"id": "dept-sales", "state": "Prospect email: 'Could you send pricing for 200 Enterprise seats and a quote for an annual contract?'", "questions": {"department": {"type": "choice", "instructions": "Which department should handle this email?", "criteria": {"billing": "invoices, payments, refunds", "technical": "bugs, outages, system errors", "sales": "pricing, new contracts", "other": "everything else"}}}, "gold": {"department": "sales"}}
4
+ {"id": "dept-other", "state": "User email: 'How do I change the display name shown on my profile?'", "questions": {"department": {"type": "choice", "instructions": "Which department should handle this email?", "criteria": {"billing": "invoices, payments, refunds", "technical": "bugs, outages, system errors", "sales": "pricing, new contracts", "other": "everything else"}}}, "gold": {"department": "other"}}
5
+ {"id": "urgency-critical", "state": "Production alert: 'Payment processing is down in production. Every checkout fails and customers are blocked right now.'", "questions": {"urgency": {"type": "score", "instructions": "How urgent is this request?", "criteria": ["not urgent", "soon", "critical deadline or blocking issue"]}}, "gold": {"urgency": 2}}
6
+ {"id": "urgency-low", "state": "Feature request: 'Sometime next quarter it would be nice to have a dark mode option.'", "questions": {"urgency": {"type": "score", "instructions": "How urgent is this request?", "criteria": ["not urgent", "soon", "critical deadline or blocking issue"]}}, "gold": {"urgency": 0}}
7
+ {"id": "refund-false", "state": "Customer email: 'Please update the billing address on our account to 12 Rue de Rivoli, Paris.'", "questions": {"refund": {"type": "noul", "instructions": "Does the customer ask for money back?"}}, "gold": {"refund": false}}
8
+ {"id": "outage-true", "state": "Ticket: 'After the latest release the CSV export button returns HTTP 500 for every user.'", "questions": {"system_error": {"type": "noul", "instructions": "Is the user reporting a software bug or system error?"}}, "gold": {"system_error": true}}
benchmarks/evaluate.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Run the labeled decision eval and print a Markdown/JSON summary.
2
+
3
+ Usage:
4
+ python benchmarks/evaluate.py <model_dir_or_hub_id> [--json benchmarks/results.json]
5
+
6
+ ``<model>`` is a local ONNX bundle or an already-cached Hub id. Runs on CPU by
7
+ default; pass ``--deterministic`` for the reproducible path.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import argparse
13
+ import json
14
+ from pathlib import Path
15
+
16
+ from laya_onnx import load
17
+ from laya_onnx.eval import format_markdown, load_items, run_eval
18
+
19
+ HERE = Path(__file__).resolve().parent
20
+ DEFAULT_EVAL = HERE / "eval" / "decisions.jsonl"
21
+
22
+
23
+ def main(argv=None) -> int:
24
+ parser = argparse.ArgumentParser(prog="benchmarks/evaluate.py", description=__doc__)
25
+ parser.add_argument("model", help="Local ONNX bundle or an already-cached Hub id")
26
+ parser.add_argument("--eval", type=Path, default=DEFAULT_EVAL, help="JSONL eval set")
27
+ parser.add_argument("--providers", default="cpu")
28
+ parser.add_argument("--threads", type=int)
29
+ parser.add_argument("--deterministic", action="store_true")
30
+ parser.add_argument("--json", type=Path, help="Write the raw results JSON here")
31
+ args = parser.parse_args(argv)
32
+
33
+ items = load_items(args.eval)
34
+ agent = load(
35
+ args.model,
36
+ providers=args.providers,
37
+ threads=args.threads,
38
+ deterministic=args.deterministic,
39
+ )
40
+ results = run_eval(agent, items)
41
+ results["model"] = {
42
+ "id": str(args.model),
43
+ "providers": args.providers,
44
+ "threads": args.threads,
45
+ "deterministic": bool(args.deterministic),
46
+ }
47
+ if args.json:
48
+ args.json.parent.mkdir(parents=True, exist_ok=True)
49
+ args.json.write_text(json.dumps(results, indent=2) + "\n", encoding="utf-8")
50
+ print(f"wrote {args.json}")
51
+ print(format_markdown(results))
52
+ return 0
53
+
54
+
55
+ if __name__ == "__main__":
56
+ raise SystemExit(main())
benchmarks/pc_benchmark.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """P50/P95 for one short multilingual decision on this PC.
2
+
3
+ Usage: python benchmarks/pc_benchmark.py <model_dir> [--calls 2000]
4
+ """
5
+
6
+ import argparse
7
+ import json
8
+ import statistics
9
+ import time
10
+
11
+ from laya_onnx import load
12
+
13
+ parser = argparse.ArgumentParser()
14
+ parser.add_argument("model")
15
+ parser.add_argument("--calls", type=int, default=2000)
16
+ parser.add_argument("--providers", default="cpu")
17
+ args = parser.parse_args()
18
+ agent = load(args.model, providers=args.providers)
19
+ questions = {"refund": {"type": "noul", "instructions": "Does the customer request a refund?"}}
20
+ agent.predict("Warmup state.", questions) # compile/warm
21
+ samples = []
22
+ for _ in range(args.calls):
23
+ t0 = time.perf_counter()
24
+ agent.predict("The customer asks for a refund of a duplicate payment.", questions)
25
+ samples.append((time.perf_counter() - t0) * 1000)
26
+ samples.sort()
27
+ print(
28
+ json.dumps(
29
+ {
30
+ "n": len(samples),
31
+ "p50_ms": samples[len(samples) // 2],
32
+ "p95_ms": samples[int(len(samples) * 0.95)],
33
+ "mean_ms": statistics.fmean(samples),
34
+ },
35
+ indent=2,
36
+ )
37
+ )
benchmarks/results.json ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "by_type": {
3
+ "choice": {
4
+ "n": 4,
5
+ "accuracy": 1.0,
6
+ "ece": 0.4559
7
+ },
8
+ "noul": {
9
+ "n": 4,
10
+ "accuracy": 1.0,
11
+ "ece": 0.1288
12
+ },
13
+ "score": {
14
+ "n": 2,
15
+ "accuracy": 0.5,
16
+ "ece": 0.1122,
17
+ "mae": 0.5302
18
+ }
19
+ },
20
+ "overall": {
21
+ "items": 8,
22
+ "questions": 10,
23
+ "answers_scored": 10,
24
+ "accuracy": 0.9,
25
+ "ece": 0.2563,
26
+ "latency_ms": {
27
+ "mean": 2807.998,
28
+ "p50": 1759.94,
29
+ "p95": 7417.614
30
+ }
31
+ },
32
+ "model": {
33
+ "id": "C:\\Users\\geoff\\.cache\\huggingface\\hub\\laya-onnx-bundles\\receptron--laya-onnx",
34
+ "providers": "cpu",
35
+ "threads": null,
36
+ "deterministic": false
37
+ }
38
+ }
benchmarks/results.md ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Benchmarks
2
+
3
+ Reproducible decision benchmarks for `laya-onnx`. Everything here runs offline
4
+ once weights are cached; nothing is a leaderboard — these are self-checks you
5
+ can rerun on your own machine.
6
+
7
+ ## Decision eval (accuracy + calibration)
8
+
9
+ ```bash
10
+ # local bundle
11
+ python benchmarks/evaluate.py ./onnx --json benchmarks/results.json
12
+ # or an already-cached Hub bundle
13
+ python benchmarks/evaluate.py receptron/laya-onnx
14
+ ```
15
+
16
+ The eval set (`benchmarks/eval/decisions.jsonl`) is a small labeled self-check:
17
+ **8 states / 10 questions** across the three typed forms (`choice`, `score`,
18
+ `noul`). Confidence values are the model's own; calibration is reported as
19
+ expected calibration error (ECE, 10 bins, lower is better).
20
+
21
+ ### Published result
22
+
23
+ Environment: Windows 11, 12th Gen Intel Core i7-1255U (12 threads, 15.8 GB RAM),
24
+ `onnxruntime` CPU, bundle `receptron/laya-onnx` @ `68f27dfe`.
25
+
26
+ | question type | n | accuracy | MAE | ECE |
27
+ | --- | ---: | ---: | ---: | ---: |
28
+ | choice | 4 | 1.00 | – | 0.456 |
29
+ | noul | 4 | 1.00 | – | 0.129 |
30
+ | score | 2 | 0.50 | 0.530 | 0.112 |
31
+ | **overall** | **10** | **0.90** | – | **0.256** |
32
+
33
+ Latency per `predict()` call: mean **2808 ms** / p50 **1760 ms** / p95 **7418 ms**.
34
+
35
+ Reading:
36
+ - All four `choice` and all four `noul` decisions were correct; one of the two
37
+ `score` items missed (MAE 0.53).
38
+ - Calibration is **not** tight at this sample size. The `choice` bucket is
39
+ *under*-confident (correct, but ~0.55 confidence), which is recorded rather
40
+ than hidden. The checkpoint also ships one calibration temperature outside the
41
+ accepted range; `load()` warns and clamps it (expected, not an error).
42
+ - With 10 questions the confidence intervals are wide — treat this as a smoke
43
+ result, not a measurement of model quality.
44
+
45
+ Raw output: [`benchmarks/results.json`](results.json).
46
+
47
+ ## Latency
48
+
49
+ ```bash
50
+ python benchmarks/pc_benchmark.py ./onnx --calls 2000
51
+ ```
52
+
53
+ P50/P95 for one short multilingual decision on this machine (no eval set
54
+ needed). See the note above for a measured sample.
55
+
56
+ ## Verify a download
57
+
58
+ Every published bundle ships a SHA-256 manifest under `laya_onnx/checksums/`.
59
+ Check a local copy:
60
+
61
+ ```bash
62
+ laya-onnx verify --model ~/.cache/huggingface/hub/laya-onnx-bundles/receptron--laya-onnx
63
+ ```
64
+
65
+ `receptron/laya-onnx` @ `68f27dfe` — sizes and SHA-256:
66
+
67
+ | file | bytes | sha256 |
68
+ | --- | ---: | --- |
69
+ | `laya.onnx` | 3,807,291 | `a874eb25…66dba1e` |
70
+ | `laya.onnx.data` | 1,685,258,240 | `48774636…3242aba` |
71
+ | `laya_config.json` | 369 | `5049005d…69bb561` |
72
+ | `tokenizer/tokenizer.json` | 3,583,228 | `6c8aaa9a…3c08d30` |
73
+ | `tokenizer/tokenizer_config.json` | 308 | `50044de6…320dd12` |
74
+
75
+ The two weight files match the SHA-256 published by Hugging Face for their LFS
76
+ objects; full digests are in
77
+ [`laya_onnx/checksums/receptron-laya-onnx.json`](../laya_onnx/checksums/receptron-laya-onnx.json).
docs/DETERMINISTIC.md ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ # Deterministic inference
2
+
3
+ - `load(..., deterministic=True)` → ORT `intra_op_num_threads=1`, no pad-to-16, CPU arena off.
4
+ - `Agent.predict_argmax(state, questions)` → ignore temperature; `choice`/`score`/`noul` from `argmax(logits)`.
5
+ - `laya-onnx-ultrafast --deterministic` → `predict_argmax` + `_heuristic_text` (never the OpenAI helper).
6
+ - Still not bit-identical across ORT versions. Pin `onnxruntime` and the ONNX graph.
docs/ULTRAFAST.md ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # laya-onnx ultrafast
2
+
3
+ Port of jev-ultrafast / laya-ultrafast onto local ONNX Laya.
4
+
5
+ `--dry-run` never opens Chrome. `--deterministic` uses `predict_argmax`, ORT threads=1, and heuristic TYPE_TEXT (no TEXT_MODEL).
6
+
7
+ ```bash
8
+ laya-onnx-ultrafast --deterministic --dry-run --fixture examples/ultrafast_page.json \
9
+ --goal "Find one-way flights from Zurich to London on 20 September 2026."
10
+ ```
11
+
12
+ Ops: CLICK TYPE_TEXT SELECT SCROLL_DOWN SCROLL_UP WAIT DONE BLOCKED
examples/questions.json ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "department": {
3
+ "type": "choice",
4
+ "instructions": "Which department should handle this email?",
5
+ "criteria": {
6
+ "billing": "invoices, payments, refunds",
7
+ "technical": "bugs, outages, system errors",
8
+ "sales": "pricing, new contracts",
9
+ "other": "everything else"
10
+ }
11
+ },
12
+ "urgency": {
13
+ "type": "score",
14
+ "instructions": "How urgent is this request?",
15
+ "criteria": ["not urgent", "soon", "critical deadline or blocking issue"]
16
+ },
17
+ "refund": {
18
+ "type": "noul",
19
+ "instructions": "Does the customer ask for money back?"
20
+ }
21
+ }
examples/quickstart.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Run with: python examples/quickstart.py"""
2
+
3
+ import json
4
+ from pathlib import Path
5
+
6
+ import laya_onnx as laya
7
+
8
+ root = Path(__file__).parent
9
+ agent = laya.load("receptron/laya-onnx", providers="cpu")
10
+ result = agent.predict(
11
+ json.loads((root / "state.json").read_text()),
12
+ json.loads((root / "questions.json").read_text()),
13
+ )
14
+ print(json.dumps(result, indent=2, ensure_ascii=False))
examples/snake.html ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <!DOCTYPE html>
2
+ <html lang="en"><head><meta charset="UTF-8"/><meta name="viewport" content="width=device-width,initial-scale=1"/>
3
+ <title>laya-onnx Snake</title>
4
+ <style>
5
+ html,body{margin:0;background:#070b12;color:#e2e8f0;font-family:system-ui,sans-serif}
6
+ .wrap{max-width:980px;margin:0 auto;padding:20px}.kicker{color:#34d399;font-size:11px;letter-spacing:.18em;text-transform:uppercase}
7
+ h1{margin:6px 0 14px;font-size:26px}.grid{display:grid;grid-template-columns:1fr 260px;gap:16px}
8
+ @media(max-width:840px){.grid{grid-template-columns:1fr}}
9
+ canvas#c{display:block;width:100%;max-width:720px;background:#020617;border:2px solid #34d399;border-radius:12px}
10
+ .panel{background:#0f172a;border:1px solid #1e293b;border-radius:12px;padding:14px}
11
+ .panel h2{margin:0 0 10px;font-size:13px;letter-spacing:.08em;text-transform:uppercase;color:#94a3b8}
12
+ .stat{display:flex;justify-content:space-between;font-size:13px;padding:5px 0;border-bottom:1px solid #1e293b}
13
+ .stat b{color:#6ee7b7;font-variant-numeric:tabular-nums}
14
+ .keys{display:none;margin-top:12px}.keys.on{display:block}
15
+ .k{display:inline-block;min-width:22px;padding:3px 7px;margin:2px;border-radius:6px;background:#1e293b;border:1px solid #334155;font-size:12px;text-align:center}
16
+ .row{display:flex;gap:8px;flex-wrap:wrap;align-items:center;margin-top:12px}
17
+ button,select{padding:8px 12px;border:0;border-radius:8px;background:#334155;color:#fff}
18
+ #play{background:#34d399;color:#052e16;font-weight:700}#spark{width:100%;height:56px;background:#020617;border-radius:8px;margin-top:8px}
19
+ .hint{color:#94a3b8;font-size:12px;margin-top:8px;line-height:1.45}
20
+ </style></head><body>
21
+ <div class="wrap"><div class="kicker">laya-onnx</div><h1>Snake bench</h1>
22
+ <div class="grid"><div>
23
+ <canvas id="c" width="720" height="480">Canvas required</canvas>
24
+ <div class="row"><button id="play">Pause</button><button id="reset">Reset</button>
25
+ <select id="mode"><option value="laya">Laya mock</option><option value="cycle">Hamilton cycle</option><option value="human">Keyboard</option></select>
26
+ <label class="hint">FPS <input id="fps" type="range" min="3" max="16" value="8"/></label></div>
27
+ <div id="keys" class="keys panel"><h2>Keyboard</h2>
28
+ <div><span class="k">↑</span><span class="k">W</span> up</div>
29
+ <div><span class="k">←</span><span class="k">A</span> left <span class="k">↓</span><span class="k">S</span> down <span class="k">→</span><span class="k">D</span> right</div>
30
+ <div><span class="k">space</span> pause <span class="k">R</span> reset</div>
31
+ <p class="hint">Choose Keyboard, then use the keys above. Reversing into the neck is illegal.</p></div>
32
+ </div>
33
+ <div class="panel"><h2>Performance</h2>
34
+ <div class="stat">Score <b id="sc">0</b></div>
35
+ <div class="stat">Best <b id="best">0</b></div>
36
+ <div class="stat">Length <b id="ln">6</b></div>
37
+ <div class="stat">Steps <b id="st">0</b></div>
38
+ <div class="stat">Eats <b id="eat">0</b></div>
39
+ <div class="stat">Deaths <b id="dead">0</b></div>
40
+ <div class="stat">Elapsed <b id="el">0.0s</b></div>
41
+ <div class="stat">Steps / s <b id="sps">0.0</b></div>
42
+ <div class="stat">Last move <b id="mv">-</b></div>
43
+ <canvas id="spark" width="232" height="56"></canvas>
44
+ <p class="hint">Sparkline = length over the last 80 steps. Mock policy, no ONNX.</p>
45
+ </div></div></div>
46
+ <script>
47
+ (function(){
48
+ var W=24,H=16,CAP=W*H,CW=720,CH=480,CW0=CW/W,CH0=CH/H;
49
+ var canvas=document.getElementById("c"),ctx=canvas.getContext("2d");
50
+ var spark=document.getElementById("spark"),sctx=spark.getContext("2d");
51
+ var DIRS={UP:[0,-1],DOWN:[0,1],LEFT:[-1,0],RIGHT:[1,0]},NAMES=["UP","DOWN","LEFT","RIGHT"];
52
+ function ham(w,h){if(h%2)return ham(h,w).map(function(p){return[p[1],p[0]]});var p=[[0,0]],y,x;for(y=0;y<h;y++){if(y%2===0)for(x=1;x<w;x++)p.push([x,y]);else for(x=w-1;x>0;x--)p.push([x,y]);}for(y=h-1;y>0;y--)p.push([0,y]);return p;}
53
+ var cycle=ham(W,H),idx={};for(var i=0;i<cycle.length;i++)idx[cycle[i][0]+","+cycle[i][1]]=i;
54
+ function rng(s){return function(){var t=(s+=0x6d2b79f5);t=Math.imul(t^t>>>15,t|1);t^=t+Math.imul(t^t>>>7,t|61);return((t^t>>>14)>>>0)/4294967296;};}
55
+ function Game(){this.rand=rng(7);var start=idx[Math.floor(W/2)+","+Math.floor(H/2)];this.body=[];for(var i=0;i<6;i++)this.body.push(cycle[(start-i+CAP)%CAP]);this.score=0;this.alive=true;this.food=this.spawn();}
56
+ Game.prototype.spawn=function(){var occ={},i,empty=[];for(i=0;i<this.body.length;i++)occ[this.body[i][0]+","+this.body[i][1]]=1;for(i=0;i<cycle.length;i++)if(!occ[cycle[i][0]+","+cycle[i][1]])empty.push(cycle[i]);return empty.length?empty[Math.floor(this.rand()*empty.length)]:null;};
57
+ Game.prototype.target=function(d){var v=DIRS[d];return[this.body[0][0]+v[0],this.body[0][1]+v[1]];};
58
+ Game.prototype.legal=function(d){var t=this.target(d),x=t[0],y=t[1],i;if(x<0||y<0||x>=W||y>=H)return"wall";if(this.body[1]&&this.body[1][0]===x&&this.body[1][1]===y)return"reverse";var occ={};for(i=0;i<this.body.length;i++)occ[this.body[i][0]+","+this.body[i][1]]=1;var eats=this.food&&this.food[0]===x&&this.food[1]===y;if(!eats){var tl=this.body[this.body.length-1];delete occ[tl[0]+","+tl[1]];}return occ[x+","+y]?"body":"legal";};
59
+ Game.prototype.step=function(d){if(!this.alive)return"dead";var why=this.legal(d);if(why!=="legal"){this.alive=false;return why;}var n=this.target(d);var eats=this.food&&n[0]===this.food[0]&&n[1]===this.food[1];this.body.unshift(n);if(eats){this.score++;this.food=this.spawn();}else this.body.pop();return eats?"eat":"ok";};
60
+ function cycleDir(g){var i=idx[g.body[0][0]+","+g.body[0][1]],n=cycle[(i+1)%CAP];var dx=n[0]-g.body[0][0],dy=n[1]-g.body[0][1];if(dx===1)return"RIGHT";if(dx===-1)return"LEFT";if(dy===1)return"DOWN";return"UP";}
61
+ function mock(g){var best=null,bd=99,i;for(i=0;i<4;i++){var d=NAMES[i];if(g.legal(d)!=="legal")continue;var t=g.target(d),dist=g.food?Math.abs(t[0]-g.food[0])+Math.abs(t[1]-g.food[1]):99;if(dist<bd){bd=dist;best=d;}}return best||cycleDir(g);}
62
+ function paint(){ctx.fillStyle="#020617";ctx.fillRect(0,0,CW,CH);ctx.strokeStyle="#1e293b";var x,y;for(x=0;x<=W;x++){ctx.beginPath();ctx.moveTo(x*CW0,0);ctx.lineTo(x*CW0,CH);ctx.stroke();}for(y=0;y<=H;y++){ctx.beginPath();ctx.moveTo(0,y*CH0);ctx.lineTo(CW,y*CH0);ctx.stroke();}if(game.food){ctx.fillStyle="#fbbf24";ctx.fillRect(game.food[0]*CW0+3,game.food[1]*CH0+3,CW0-6,CH0-6);}for(var i=game.body.length-1;i>=0;i--){var c=game.body[i];ctx.fillStyle=i===0?"#6ee7b7":"#059669";ctx.fillRect(c[0]*CW0+2,c[1]*CH0+2,CW0-4,CH0-4);}}
63
+ function sparkline(){sctx.fillStyle="#020617";sctx.fillRect(0,0,232,56);if(hist.length<2)return;var max=1,i;for(i=0;i<hist.length;i++)if(hist[i]>max)max=hist[i];sctx.strokeStyle="#34d399";sctx.beginPath();for(i=0;i<hist.length;i++){var px=i*(232/Math.max(hist.length-1,1));var py=52-(hist[i]/max)*48;if(i===0)sctx.moveTo(px,py);else sctx.lineTo(px,py);}sctx.stroke();}
64
+ var game=new Game(),running=false,humanDir="RIGHT",timer=null;
65
+ var steps=0,eats=0,deaths=0,best=0,t0=performance.now(),hist=[];
66
+ function dash(dir,r){document.getElementById("sc").textContent=game.score;if(game.score>best)best=game.score;document.getElementById("best").textContent=best;document.getElementById("ln").textContent=game.body.length;document.getElementById("st").textContent=steps;document.getElementById("eat").textContent=eats;document.getElementById("dead").textContent=deaths;var sec=(performance.now()-t0)/1000;document.getElementById("el").textContent=sec.toFixed(1)+"s";document.getElementById("sps").textContent=sec>0?(steps/sec).toFixed(1):"0.0";document.getElementById("mv").textContent=(dir||"-")+" / "+(r||"-");}
67
+ function tick(){var mode=document.getElementById("mode").value;var dir=cycleDir(game);if(mode==="laya")dir=mock(game);if(mode==="human")dir=humanDir;var r=game.step(dir);steps++;if(r==="eat")eats++;if(!game.alive)deaths++;hist.push(game.body.length);if(hist.length>80)hist.shift();paint();sparkline();dash(dir,r);if(!game.alive){stop();setTimeout(function(){game=new Game();humanDir="RIGHT";paint();dash("reset","new");if(document.getElementById("mode").value!=="human")start();},700);}}
68
+ function start(){if(running)return;running=true;document.getElementById("play").textContent="Pause";timer=setInterval(tick,1000/Number(document.getElementById("fps").value));}
69
+ function stop(){running=false;document.getElementById("play").textContent="Play";clearInterval(timer);}
70
+ function reset(){stop();game=new Game();humanDir="RIGHT";steps=0;eats=0;deaths=0;t0=performance.now();hist=[];paint();sparkline();dash("reset","ok");start();}
71
+ function syncKeys(){document.getElementById("keys").className="keys panel"+(document.getElementById("mode").value==="human"?" on":"");}
72
+ document.getElementById("play").onclick=function(){running?stop():start();};
73
+ document.getElementById("reset").onclick=reset;
74
+ document.getElementById("mode").onchange=syncKeys;
75
+ document.getElementById("fps").onchange=function(){if(running){stop();start();}};
76
+ window.addEventListener("keydown",function(e){var map={ArrowUp:"UP",ArrowDown:"DOWN",ArrowLeft:"LEFT",ArrowRight:"RIGHT",w:"UP",s:"DOWN",a:"LEFT",d:"RIGHT",W:"UP",S:"DOWN",A:"LEFT",D:"RIGHT"};if(map[e.key]){humanDir=map[e.key];if(document.getElementById("mode").value!=="human"){document.getElementById("mode").value="human";syncKeys();}e.preventDefault();}if(e.key===" "){running?stop():start();e.preventDefault();}if(e.key==="r"||e.key==="R")reset();});
77
+ syncKeys();paint();sparkline();dash("-","ready");start();
78
+ })();
79
+ </script></body></html>
examples/state.json ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ {
2
+ "from": "user@example.com",
3
+ "subject": "Duplicate charge on invoice #4411",
4
+ "body": "We were billed twice for March. Please refund the duplicate today or we will cancel our plan."
5
+ }
examples/ultrafast_goal.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Offline fixture walk — no Chrome, no Hub weights if you pass a stub agent."""
2
+
3
+ from laya_onnx.ultrafast import UltrafastAgent
4
+
5
+ GOAL = "Find one-way flights from Zurich to London on 20 September 2026."
6
+
7
+ if __name__ == "__main__":
8
+ with UltrafastAgent(
9
+ "https://www.google.com/travel/flights?hl=en",
10
+ GOAL,
11
+ fixture="examples/ultrafast_page.json",
12
+ ) as agent:
13
+ for step in agent.run(max_steps=8):
14
+ print(step["op"], step.get("target"), step["status"])
examples/ultrafast_page.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "url": "https://www.google.com/travel/flights?hl=en",
3
+ "title": "Google Flights",
4
+ "text": "Flights One way Zurich London Sat, Sep 20",
5
+ "elements": [
6
+ {"index": 0, "role": "button", "name": "Round trip", "value": "", "tag": "button", "editable": false},
7
+ {"index": 1, "role": "button", "name": "One way", "value": "", "tag": "button", "editable": false},
8
+ {"index": 2, "role": "textbox", "name": "Where from?", "value": "", "tag": "input", "editable": true},
9
+ {"index": 3, "role": "textbox", "name": "Where to?", "value": "", "tag": "input", "editable": true},
10
+ {"index": 4, "role": "textbox", "name": "Departure", "value": "", "tag": "input", "editable": true},
11
+ {"index": 5, "role": "button", "name": "Search", "value": "", "tag": "button", "editable": false}
12
+ ]
13
+ }
laya_onnx/__init__.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ """Laya typed decisions on ONNX Runtime (CPU / OpenVINO / CUDA)."""
2
+
3
+ from .agent import Agent, RLAgent, load
4
+
5
+ __version__ = "0.1.0"
6
+ __all__ = [
7
+ "Agent",
8
+ "RLAgent",
9
+ "load",
10
+ ]
laya_onnx/__main__.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ from .cli import main
2
+
3
+ main()
laya_onnx/agent.py ADDED
@@ -0,0 +1,280 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ONNX Runtime inference; prompt and result formats follow upstream Laya."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import math
7
+ import warnings
8
+ from pathlib import Path
9
+
10
+ import numpy as np
11
+
12
+ from .common import (
13
+ QTYPES,
14
+ TEMP_MAX,
15
+ TEMP_MIN,
16
+ build_sequence,
17
+ clamp_temperature,
18
+ confidence_from_probs,
19
+ render_options,
20
+ temp_bucket,
21
+ )
22
+ from .hub import find_config, find_graph, resolve_bundle
23
+ from .inputs import collate_items
24
+ from .tokenizer import Tokenizer
25
+
26
+
27
+ def _providers(name: str) -> list[str]:
28
+ if name == "openvino":
29
+ return ["OpenVINOExecutionProvider", "CPUExecutionProvider"]
30
+ if name == "cuda":
31
+ return ["CUDAExecutionProvider", "CPUExecutionProvider"]
32
+ return ["CPUExecutionProvider"]
33
+
34
+
35
+ def _session_dtypes(session) -> dict[str, np.dtype]:
36
+ mapping = {}
37
+ for inp in session.get_inputs():
38
+ raw = (inp.type or "").lower()
39
+ if "int32" in raw:
40
+ mapping[inp.name] = np.int32
41
+ elif "int64" in raw:
42
+ mapping[inp.name] = np.int64
43
+ elif "bool" in raw:
44
+ mapping[inp.name] = np.bool_
45
+ elif "float16" in raw:
46
+ mapping[inp.name] = np.float16
47
+ else:
48
+ mapping[inp.name] = None
49
+ return mapping
50
+
51
+
52
+ class Agent:
53
+ def __init__(
54
+ self,
55
+ model_id_or_path="receptron/laya-onnx",
56
+ providers="cpu",
57
+ token=None,
58
+ subfolder=None,
59
+ *,
60
+ revision=None,
61
+ batch_size=16,
62
+ threads=None,
63
+ pad_to_multiple=16,
64
+ deterministic=False,
65
+ ):
66
+ if deterministic:
67
+ threads = 1
68
+ pad_to_multiple = None
69
+ if providers not in ("cpu", "openvino", "cuda"):
70
+ raise ValueError("providers must be 'cpu', 'openvino', or 'cuda'")
71
+ if not isinstance(batch_size, int) or isinstance(batch_size, bool) or batch_size < 1:
72
+ raise ValueError("batch_size must be a positive integer")
73
+ if pad_to_multiple is not None and (
74
+ not isinstance(pad_to_multiple, int)
75
+ or isinstance(pad_to_multiple, bool)
76
+ or pad_to_multiple < 1
77
+ ):
78
+ raise ValueError("pad_to_multiple must be a positive integer or None")
79
+ self.providers = providers
80
+ self.batch_size = batch_size
81
+ self.deterministic = bool(deterministic)
82
+ self.pad_to_multiple = pad_to_multiple
83
+ self.model_id = str(model_id_or_path)
84
+ self.revision = revision
85
+ self.model_dir = resolve_bundle(
86
+ model_id_or_path, token=token, subfolder=subfolder, revision=revision
87
+ )
88
+ cfg_path = find_config(self.model_dir)
89
+ raw = json.loads(cfg_path.read_text())
90
+ self.cfg = raw
91
+ max_len = int(raw.get("max_len", 512))
92
+ head_max_len = int(raw.get("head_max_len", 192))
93
+ if not 4 < head_max_len < max_len:
94
+ raise ValueError("Expected 4 < head_max_len < max_len")
95
+ self.temperature_raw = raw.get("temperature", [1.0, 1.0, 1.0])
96
+ self.temperature_by_options_raw = raw.get("temperature_by_options", {})
97
+ if len(self.temperature_raw) != 3 or any(
98
+ not math.isfinite(float(t)) or float(t) <= 0
99
+ for t in [*self.temperature_raw, *self.temperature_by_options_raw.values()]
100
+ ):
101
+ raise ValueError("Calibration temperatures must be finite and positive")
102
+ self.temperature = [clamp_temperature(t) for t in self.temperature_raw]
103
+ self.temperature_by_options = {
104
+ k: clamp_temperature(v) for k, v in self.temperature_by_options_raw.items()
105
+ }
106
+ rejected = [
107
+ "%s=%.4g" % (k, float(v))
108
+ for k, v in self.temperature_by_options_raw.items()
109
+ if clamp_temperature(v) != float(v)
110
+ ]
111
+ rejected += [
112
+ "temperature[%d]=%.4g" % (i, float(t))
113
+ for i, t in enumerate(self.temperature_raw)
114
+ if clamp_temperature(t) != float(t)
115
+ ]
116
+ if rejected:
117
+ warnings.warn(
118
+ "laya-onnx: this checkpoint ships temperatures outside [%g, %g] which would "
119
+ "distort confidence; clamping %s." % (TEMP_MIN, TEMP_MAX, ", ".join(rejected)),
120
+ RuntimeWarning,
121
+ stacklevel=2,
122
+ )
123
+ self.tok = Tokenizer(self.model_dir / "tokenizer")
124
+
125
+ import os
126
+
127
+ import onnxruntime as ort
128
+
129
+ sess_opts = ort.SessionOptions()
130
+ sess_opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
131
+ sess_opts.enable_mem_pattern = True
132
+ sess_opts.enable_cpu_mem_arena = not self.deterministic
133
+ sess_opts.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
134
+ workers = int(threads) if threads else min(8, os.cpu_count() or 4)
135
+ sess_opts.intra_op_num_threads = workers
136
+ sess_opts.inter_op_num_threads = 1
137
+ sess_opts.add_session_config_entry("session.intra_op.allow_spinning", "0")
138
+ self.session = ort.InferenceSession(
139
+ str(find_graph(self.model_dir)),
140
+ sess_options=sess_opts,
141
+ providers=_providers(providers),
142
+ )
143
+ self._dtypes = _session_dtypes(self.session)
144
+
145
+ @staticmethod
146
+ def _to_internal(qdef):
147
+ if not isinstance(qdef, dict):
148
+ raise ValueError("Each question must be a dictionary")
149
+ kind = qdef.get("type")
150
+ if kind not in QTYPES:
151
+ raise ValueError(f"Unknown question type {kind!r}; expected choice, score, or noul")
152
+ if "instructions" not in qdef:
153
+ raise ValueError("Question is missing instructions")
154
+ criteria = qdef.get("criteria")
155
+ if kind == "choice":
156
+ if isinstance(criteria, list):
157
+ if not all(isinstance(c, str) for c in criteria):
158
+ raise ValueError("Choice labels must be strings")
159
+ if len(set(criteria)) != len(criteria):
160
+ raise ValueError("Choice labels must be unique")
161
+ criteria = dict.fromkeys(criteria)
162
+ if not isinstance(criteria, dict) or not criteria:
163
+ raise ValueError("Choice criteria must be a nonempty dictionary or list")
164
+ if not all(isinstance(k, str) for k in criteria):
165
+ raise ValueError("Choice labels must be strings")
166
+ elif kind == "score":
167
+ if not isinstance(criteria, list) or not criteria:
168
+ raise ValueError("Score criteria must be a nonempty list")
169
+ elif criteria is not None and not isinstance(criteria, dict):
170
+ raise ValueError("Noul criteria must be a dictionary with false/true descriptions")
171
+ instructions = qdef["instructions"]
172
+ if not isinstance(instructions, str):
173
+ instructions = json.dumps(instructions)
174
+ return {"t": kind, "ins": instructions, "crit": criteria}
175
+
176
+ def prepare(self, state, questions):
177
+ if not isinstance(questions, dict):
178
+ raise ValueError("questions must be a dictionary keyed by question id")
179
+ items, internal = [], []
180
+ for qid, definition in questions.items():
181
+ q = self._to_internal(definition)
182
+ ids, markers = build_sequence(
183
+ self.tok, state, q, self.cfg.get("max_len", 512), self.cfg.get("head_max_len", 192)
184
+ )
185
+ if len(markers) != len(render_options(q)):
186
+ raise ValueError(f"Question {qid!r} has too many options for the token budget")
187
+ items.append({"ids": ids, "markers": markers, "qtype": QTYPES[q["t"]]})
188
+ internal.append(q)
189
+ return items, internal
190
+
191
+ def _cast(self, batch: dict) -> dict:
192
+ feeds = {}
193
+ for name, array in batch.items():
194
+ target = self._dtypes.get(name)
195
+ if target is not None and array.dtype != target:
196
+ array = array.astype(target, copy=False)
197
+ feeds[name] = array
198
+ return feeds
199
+
200
+ def forward(self, batch):
201
+ feeds = self._cast(batch)
202
+ names = [o.name for o in self.session.get_outputs()]
203
+ outputs = self.session.run(names, feeds)
204
+ return tuple(np.asarray(o) for o in outputs)
205
+
206
+ def system_one(self, state, questions, *, argmax=False):
207
+ items, internal = self.prepare(state, questions)
208
+ answers = {}
209
+ question_ids = list(questions)
210
+ id_dt = self._dtypes.get("input_ids") or np.int64
211
+ mk_dt = self._dtypes.get("marker_pos") or np.int64
212
+ for start in range(0, len(items), self.batch_size):
213
+ chunk = items[start : start + self.batch_size]
214
+ batch = collate_items(
215
+ chunk,
216
+ self.tok.pad_token_id,
217
+ pad_to_multiple=self.pad_to_multiple,
218
+ max_length=self.cfg.get("max_len", 512),
219
+ ids_dtype=id_dt,
220
+ marker_dtype=mk_dt,
221
+ )
222
+ logits, act = self.forward(batch)
223
+ if not np.isfinite(logits).all() or not np.isfinite(act).all():
224
+ raise FloatingPointError("Non-finite model outputs")
225
+ if act.ndim == 2 and act.shape[-1] > 1 and abs(float(act[0].sum()) - 1.0) > 1e-3:
226
+ act = np.exp(act - act.max(axis=-1, keepdims=True))
227
+ act /= act.sum(axis=-1, keepdims=True)
228
+ for row, item in enumerate(chunk):
229
+ qid, q = question_ids[start + row], internal[start + row]
230
+ k, qt = len(item["markers"]), item["qtype"]
231
+ scale = (
232
+ 1.0
233
+ if argmax
234
+ else self.temperature_by_options.get(temp_bucket(qt, k), self.temperature[qt])
235
+ )
236
+ z = logits[row, :k] / scale
237
+ p = np.exp(z - z.max())
238
+ p /= p.sum()
239
+ act_p = float(act[row, 0]) if act.ndim == 2 else float(act[row])
240
+ answer = {
241
+ "type": q["t"],
242
+ "confidence": round(confidence_from_probs(p, k), 4),
243
+ "action": {"act_probability": round(act_p, 4)},
244
+ }
245
+ if q["t"] == "choice":
246
+ labels = list(q["crit"])
247
+ answer.update(
248
+ choice=labels[int(p.argmax())],
249
+ probabilities={label: round(float(v), 4) for label, v in zip(labels, p)},
250
+ )
251
+ elif q["t"] == "score":
252
+ answer.update(
253
+ score=round(float((np.arange(k) * p).sum()), 4),
254
+ legend={str(i): value for i, value in enumerate(q["crit"])},
255
+ probabilities={str(i): round(float(v), 4) for i, v in enumerate(p)},
256
+ )
257
+ else:
258
+ answer.update(
259
+ noul=round(float(p[1]), 4),
260
+ confidence=round(max(float(p[1]), 1.0 - float(p[1])), 4),
261
+ )
262
+ answers[qid] = answer
263
+ return {
264
+ "model": "laya-onnx",
265
+ "answers": answers,
266
+ "usage": {"input_tokens": sum(len(item["ids"]) for item in items), "output_tokens": 0},
267
+ }
268
+
269
+ predict = system_one
270
+
271
+ def predict_argmax(self, state, questions):
272
+ """Deterministic path: ignore calibrated temperature (see docs/DETERMINISTIC.md)."""
273
+ return self.system_one(state, questions, argmax=True)
274
+
275
+
276
+ RLAgent = Agent
277
+
278
+
279
+ def load(model_id_or_path="receptron/laya-onnx", providers="cpu", token=None, subfolder=None, deterministic=False, **kwargs):
280
+ return Agent(model_id_or_path, providers=providers, token=token, subfolder=subfolder, deterministic=deterministic, **kwargs)
laya_onnx/checksums/receptron-laya-onnx.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "repo": "receptron/laya-onnx",
3
+ "revision": "68f27dfe5a27a54fb2b1fefc432f43f972e90868",
4
+ "source": "https://huggingface.co/receptron/laya-onnx",
5
+ "algorithm": "sha256",
6
+ "files": {
7
+ "laya.onnx": {
8
+ "size": 3807291,
9
+ "sha256": "a874eb254b58b0fcb1e7ad56fbb188c29d64e08c9a46b689433e1f52c66dba1e"
10
+ },
11
+ "laya.onnx.data": {
12
+ "size": 1685258240,
13
+ "sha256": "487746363a8da57bcadb4345352997d22a0fb90d70aa22c6856668d023242aba"
14
+ },
15
+ "laya_config.json": {
16
+ "size": 369,
17
+ "sha256": "5049005dc6ae3ca5e82cc7d85c421357d5c543817300c8e8c5281ddbc69bb561"
18
+ },
19
+ "tokenizer/tokenizer.json": {
20
+ "size": 3583228,
21
+ "sha256": "6c8aaa9a542084f2457eab775d4eeb51f92a70c0fd9de28d5edb0ddec3c08d30"
22
+ },
23
+ "tokenizer/tokenizer_config.json": {
24
+ "size": 308,
25
+ "sha256": "50044de60daaa73df97d262e15a40d4faf0160e7d742df64b377877a1320dd12"
26
+ }
27
+ }
28
+ }
laya_onnx/cli.py ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Command-line prediction and checkpoint conversion."""
2
+
3
+ import argparse
4
+ import json
5
+ from pathlib import Path
6
+
7
+ from . import __version__
8
+ from .agent import Agent
9
+
10
+
11
+ def main(argv=None):
12
+ parser = argparse.ArgumentParser(prog="laya-onnx", description=__doc__)
13
+ parser.add_argument("--version", action="version", version=__version__)
14
+ commands = parser.add_subparsers(dest="command", required=True)
15
+
16
+ predict = commands.add_parser("predict")
17
+ predict.add_argument("--model", default="receptron/laya-onnx")
18
+ predict.add_argument("--subfolder")
19
+ predict.add_argument("--revision")
20
+ predict.add_argument("--providers", choices=("cpu", "openvino", "cuda"), default="cpu")
21
+ predict.add_argument("--threads", type=int)
22
+ source = predict.add_mutually_exclusive_group(required=True)
23
+ source.add_argument("--state", help="Plain text input")
24
+ source.add_argument("--state-file", type=Path, help="JSON state file")
25
+ predict.add_argument("--questions", required=True, type=Path)
26
+ predict.add_argument("--batch-size", type=int, default=16)
27
+ predict.add_argument(
28
+ "--deterministic", action="store_true", help="threads=1, pad off, predict_argmax"
29
+ )
30
+
31
+ convert = commands.add_parser("convert")
32
+ convert.add_argument("--model", default="convaiinnovations/laya")
33
+ convert.add_argument("--subfolder")
34
+ convert.add_argument("--revision")
35
+ convert.add_argument("--output", type=Path, required=True)
36
+ convert.add_argument("--precision", choices=("fp32", "int8"), default="fp32")
37
+ convert.add_argument("--opset", type=int, default=17)
38
+
39
+ verify = commands.add_parser("verify")
40
+ verify.add_argument("--model", required=True, help="Local model directory to check")
41
+ verify.add_argument(
42
+ "--checksums",
43
+ type=Path,
44
+ help="Checksum manifest JSON (default: bundled receptron/laya-onnx manifest)",
45
+ )
46
+
47
+ args = parser.parse_args(argv)
48
+ if args.command == "convert":
49
+ from .convert import convert as do_convert
50
+
51
+ result = do_convert(
52
+ args.model,
53
+ args.output,
54
+ precision=args.precision,
55
+ revision=args.revision,
56
+ subfolder=args.subfolder,
57
+ opset=args.opset,
58
+ )
59
+ print(json.dumps({"output": str(result), "precision": args.precision}))
60
+ return
61
+ if args.command == "verify":
62
+ from .verify import verify_bundle
63
+
64
+ try:
65
+ ok, rows = verify_bundle(args.model, args.checksums)
66
+ except FileNotFoundError as error:
67
+ parser.error(str(error))
68
+ for row in rows:
69
+ print("%15s %s" % (row["status"], row["file"]))
70
+ return 0 if ok else 1
71
+ state = args.state if args.state is not None else json.loads(args.state_file.read_text())
72
+ questions = json.loads(args.questions.read_text())
73
+ agent = Agent(
74
+ args.model,
75
+ providers=args.providers,
76
+ revision=args.revision,
77
+ subfolder=args.subfolder,
78
+ batch_size=args.batch_size,
79
+ threads=args.threads,
80
+ deterministic=args.deterministic,
81
+ )
82
+ result = (
83
+ agent.predict_argmax(state, questions)
84
+ if args.deterministic
85
+ else agent.predict(state, questions)
86
+ )
87
+ print(json.dumps(result, ensure_ascii=False, indent=2))
laya_onnx/common.py ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Laya prompt construction and calibration, adapted from upstream (see NOTICE)."""
2
+
3
+ import json
4
+ import math
5
+ from typing import Dict, List, Optional, Union
6
+
7
+ import numpy as np
8
+
9
+ QTYPES = {"choice": 0, "score": 1, "noul": 2}
10
+ QTYPE_NAMES = {v: k for k, v in QTYPES.items()}
11
+
12
+
13
+ def serialize_state(state: Union[str, dict, list]) -> str:
14
+ if isinstance(state, str):
15
+ return state
16
+ return json.dumps(state, ensure_ascii=False)
17
+
18
+
19
+ def render_criterion(value) -> str:
20
+ if isinstance(value, str):
21
+ return value
22
+ return json.dumps(value, ensure_ascii=False, separators=(", ", ": "), default=str)
23
+
24
+
25
+ def render_options(q: Dict) -> List[str]:
26
+ t, crit = q["t"], q.get("crit")
27
+ if t == "choice":
28
+ return [
29
+ k if v is None or v == "" else "%s: %s" % (k, render_criterion(v))
30
+ for k, v in crit.items()
31
+ ]
32
+ if t == "score":
33
+ return ["level %d: %s" % (i, render_criterion(c)) for i, c in enumerate(crit)]
34
+ crit = crit or {}
35
+ false_crit, true_crit = crit.get("false"), crit.get("true")
36
+ return [
37
+ "false: "
38
+ + (
39
+ render_criterion(false_crit)
40
+ if false_crit not in (None, "")
41
+ else "no, the statement does not hold"
42
+ ),
43
+ "true: "
44
+ + (
45
+ render_criterion(true_crit)
46
+ if true_crit not in (None, "")
47
+ else "yes, the statement holds"
48
+ ),
49
+ ]
50
+
51
+
52
+ def build_prefix(tok, q: Dict, head_max_len: int = 192, option_order=None):
53
+ mask_tok = tok.mask_token
54
+ opts = render_options(q)
55
+ order = option_order if option_order is not None else list(range(len(opts)))
56
+ ins = str(q["ins"]).replace(mask_tok, " ")
57
+ head_ids = tok("%s question: %s" % (q["t"], ins), add_special_tokens=False)["input_ids"]
58
+ opt_ids = []
59
+ for i in order:
60
+ opt_ids.append(
61
+ [tok.mask_token_id]
62
+ + tok(" " + opts[i].replace(mask_tok, " "), add_special_tokens=False)["input_ids"][:48]
63
+ )
64
+ opt_budget = head_max_len - sum(len(o) for o in opt_ids)
65
+ if opt_budget < 16:
66
+ per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
67
+ opt_ids = [o[:per] for o in opt_ids]
68
+ opt_budget = head_max_len - sum(len(o) for o in opt_ids)
69
+ head_ids = head_ids[: max(8, opt_budget)]
70
+ ids = [tok.cls_token_id] + head_ids + [tok.sep_token_id]
71
+ markers = []
72
+ for o in opt_ids:
73
+ markers.append(len(ids))
74
+ ids.extend(o)
75
+ ids.append(tok.sep_token_id)
76
+ return ids, markers
77
+
78
+
79
+ def build_sequence(
80
+ tok,
81
+ state: Union[str, dict, list],
82
+ q: Dict,
83
+ max_len: int = 512,
84
+ head_max_len: int = 192,
85
+ option_order: Optional[List[int]] = None,
86
+ truncate_left: bool = False,
87
+ ):
88
+ ids, markers = build_prefix(tok, q, head_max_len, option_order)
89
+ room = max(0, max_len - len(ids) - 1)
90
+ st = tok(serialize_state(state).replace(tok.mask_token, " "), add_special_tokens=False)[
91
+ "input_ids"
92
+ ]
93
+ st = st[-room:] if truncate_left else st[:room]
94
+ ids = ids + st + [tok.sep_token_id]
95
+ return ids[:max_len], [m for m in markers if m < max_len]
96
+
97
+
98
+ def confidence_from_probs(p: np.ndarray, k: int) -> float:
99
+ if k < 2:
100
+ return 1.0
101
+ p = p[:k]
102
+ ent = -(p * np.log(np.clip(p, 1e-12, 1.0))).sum()
103
+ return float(np.clip(1.0 - ent / math.log(k), 0.0, 1.0))
104
+
105
+
106
+ def temp_bucket(qtype: int, k: int) -> str:
107
+ size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+"
108
+ return "%s:%s" % (QTYPE_NAMES[int(qtype)], size)
109
+
110
+
111
+ TEMP_MIN = 0.5
112
+ TEMP_MAX = 5.0
113
+
114
+
115
+ def clamp_temperature(t, lo: float = TEMP_MIN, hi: float = TEMP_MAX) -> float:
116
+ try:
117
+ t = float(t)
118
+ except (TypeError, ValueError):
119
+ return 1.0
120
+ if not math.isfinite(t):
121
+ return 1.0
122
+ return min(hi, max(lo, t))
laya_onnx/convert.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Export DecisionModel torch → ONNX (dynamic batch/sequence, opset 17)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import shutil
7
+ from pathlib import Path
8
+
9
+ from .inputs import collate_items
10
+ from .torch_model import DecisionModel
11
+
12
+
13
+ def _load_torch_checkpoint(source: Path):
14
+ from safetensors.torch import load_file
15
+
16
+ encoder_cfg = json.loads((source / "encoder/config.json").read_text())
17
+ agent_cfg = json.loads((source / "rl_agent_config.json").read_text())
18
+ max_len = int(agent_cfg.get("max_len", 512))
19
+ model = DecisionModel(encoder_cfg, agent_cfg, max_len)
20
+ state = load_file(str(source / "model.safetensors"))
21
+ model.load_state_dict(state, strict=True)
22
+ model.eval()
23
+ return model, encoder_cfg, agent_cfg
24
+
25
+
26
+ def convert(
27
+ model_id_or_path,
28
+ output,
29
+ *,
30
+ precision="fp32",
31
+ revision=None,
32
+ subfolder=None,
33
+ token=None,
34
+ opset=17,
35
+ batch_size=2,
36
+ max_options=4,
37
+ max_length=None,
38
+ ):
39
+ output = Path(output).expanduser()
40
+ if output.exists():
41
+ raise FileExistsError(f"Output already exists: {output}")
42
+
43
+ import torch
44
+
45
+ src = Path(model_id_or_path).expanduser()
46
+ if not src.exists():
47
+ from huggingface_hub import snapshot_download
48
+
49
+ src = Path(
50
+ snapshot_download(
51
+ str(model_id_or_path),
52
+ token=token,
53
+ revision=revision,
54
+ allow_patterns=[
55
+ "model.safetensors",
56
+ "encoder/*",
57
+ "tokenizer/*",
58
+ "rl_agent_config.json",
59
+ ],
60
+ )
61
+ )
62
+ if subfolder:
63
+ src /= subfolder
64
+
65
+ model, encoder_cfg, agent_cfg = _load_torch_checkpoint(src)
66
+ max_len = int(agent_cfg.get("max_len", max_length or 64))
67
+ dummy_len = min(32, max_len)
68
+ items = [
69
+ {
70
+ "ids": list(range(dummy_len)),
71
+ "markers": list(range(3, 3 + min(2, max_options))),
72
+ "qtype": 0,
73
+ }
74
+ for _ in range(batch_size)
75
+ ]
76
+ shape = {
77
+ "batch_size": batch_size,
78
+ "max_length": max_len,
79
+ "min_length": 16,
80
+ "max_options": max_options,
81
+ }
82
+ arrays = collate_items(items, 0, shape=shape)
83
+ example = {k: torch.from_numpy(v) for k, v in arrays.items()}
84
+ example["marker_mask"] = example["marker_mask"].bool()
85
+
86
+ output.mkdir(parents=True)
87
+ try:
88
+ graph = output / "laya.onnx"
89
+ dynamic_axes = {
90
+ "input_ids": {0: "batch", 1: "sequence"},
91
+ "attention_mask": {0: "batch", 1: "sequence"},
92
+ "marker_pos": {0: "batch", 1: "options"},
93
+ "marker_mask": {0: "batch", 1: "options"},
94
+ "qtype": {0: "batch"},
95
+ "logits": {0: "batch", 1: "options"},
96
+ "act": {0: "batch"},
97
+ }
98
+ torch.onnx.export(
99
+ model,
100
+ (
101
+ example["input_ids"],
102
+ example["attention_mask"],
103
+ example["marker_pos"],
104
+ example["marker_mask"],
105
+ example["qtype"],
106
+ ),
107
+ str(graph),
108
+ input_names=["input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype"],
109
+ output_names=["logits", "act"],
110
+ dynamic_axes=dynamic_axes,
111
+ opset_version=opset,
112
+ dynamo=False,
113
+ )
114
+ from .optimize import optimize_bundle
115
+
116
+ optimize_bundle(graph, precision=precision, fuse=True)
117
+
118
+ if (src / "tokenizer").is_dir():
119
+ shutil.copytree(src / "tokenizer", output / "tokenizer")
120
+ keys = ("max_len", "head_max_len", "temperature", "temperature_by_options")
121
+ config = {k: agent_cfg[k] for k in keys if k in agent_cfg}
122
+ config.setdefault("max_len", max_len)
123
+ config.setdefault("head_max_len", int(agent_cfg.get("head_max_len", 32)))
124
+ config.setdefault("temperature", [1.0, 1.0, 1.0])
125
+ config.setdefault("temperature_by_options", {})
126
+ (output / "laya_config.json").write_text(json.dumps(config, indent=2) + "\n")
127
+ (output / "onnx_config.json").write_text(
128
+ json.dumps({"format": "laya-onnx", "format_version": 1, "precision": precision, "opset": opset, **config, "encoder": encoder_cfg}, indent=2)
129
+ + "\n"
130
+ )
131
+ (output / "rl_agent_config.json").write_text(json.dumps(agent_cfg, indent=2) + "\n")
132
+ if (src / "encoder").is_dir():
133
+ shutil.copytree(src / "encoder", output / "encoder")
134
+ except BaseException:
135
+ shutil.rmtree(output, ignore_errors=True)
136
+ raise
137
+ return output
laya_onnx/eval.py ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Decision evaluation and calibration metrics.
2
+
3
+ Runs a labeled eval set through a loaded :class:`laya_onnx.Agent` and reports
4
+ per-question-type accuracy, score MAE, confidence calibration (ECE) and latency
5
+ percentiles. Everything here is pure/offline so it can be unit-tested with a
6
+ stub agent.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import statistics
13
+ import time
14
+ from pathlib import Path
15
+
16
+ DEFAULT_BINS = 10
17
+
18
+
19
+ def percentile(sorted_values, fraction):
20
+ """Nearest-rank percentile. ``sorted_values`` must already be sorted."""
21
+ if not sorted_values:
22
+ return None
23
+ index = min(len(sorted_values) - 1, int(len(sorted_values) * fraction))
24
+ return sorted_values[index]
25
+
26
+
27
+ def bin_index(value, bins=DEFAULT_BINS):
28
+ return min(bins - 1, max(0, int(value * bins)))
29
+
30
+
31
+ def expected_calibration_error(confidences, correct, bins=DEFAULT_BINS):
32
+ """Expected Calibration Error over (confidence, correctness) pairs.
33
+
34
+ Lower is better; 0 means confidence matches accuracy exactly.
35
+ """
36
+ if not confidences:
37
+ return None
38
+ total = len(confidences)
39
+ error = 0.0
40
+ for b in range(bins):
41
+ low = b / bins
42
+ high = (b + 1) / bins
43
+ members = [
44
+ i
45
+ for i, confidence in enumerate(confidences)
46
+ if low <= confidence < high or (b == bins - 1 and confidence >= 1.0)
47
+ ]
48
+ if not members:
49
+ continue
50
+ accuracy = sum(1.0 for i in members if correct[i]) / len(members)
51
+ mean_confidence = sum(confidences[i] for i in members) / len(members)
52
+ error += (len(members) / total) * abs(accuracy - mean_confidence)
53
+ return error
54
+
55
+
56
+ def score_answer(qtype, answer, gold):
57
+ """Return ``(correct, error)`` for one answered question.
58
+
59
+ ``correct`` is ``None`` when the item has no scoreable answer, ``error`` is
60
+ only set for ``score`` questions.
61
+ """
62
+ if answer is None or gold is None:
63
+ return None, None
64
+ if qtype == "choice":
65
+ return (answer.get("choice") == gold), None
66
+ if qtype == "noul":
67
+ try:
68
+ predicted = float(answer.get("noul", 0.0)) >= 0.5
69
+ except (TypeError, ValueError):
70
+ return None, None
71
+ return (bool(predicted) == bool(gold)), None
72
+ if qtype == "score":
73
+ try:
74
+ error = abs(float(answer.get("score", 0.0)) - float(gold))
75
+ except (TypeError, ValueError):
76
+ return None, None
77
+ return (error <= 0.5), error
78
+ return None, None
79
+
80
+
81
+ def load_items(path):
82
+ """Read a JSONL eval set (one ``{"state", "questions", "gold"}`` per line)."""
83
+ items = []
84
+ for lineno, line in enumerate(Path(path).read_text(encoding="utf-8").splitlines(), 1):
85
+ line = line.strip()
86
+ if not line:
87
+ continue
88
+ try:
89
+ items.append(json.loads(line))
90
+ except json.JSONDecodeError as exc:
91
+ raise ValueError(f"{path}:{lineno}: {exc}")
92
+ return items
93
+
94
+
95
+ def run_eval(agent, items, *, bins=DEFAULT_BINS):
96
+ """Run ``items`` through ``agent`` and return a results dict."""
97
+ latencies = []
98
+ records = []
99
+ for item in items:
100
+ started = time.perf_counter()
101
+ out = agent.predict(item["state"], item["questions"])
102
+ latencies.append((time.perf_counter() - started) * 1000.0)
103
+ answers = out.get("answers", {})
104
+ gold = item.get("gold") or {}
105
+ for qid, definition in item["questions"].items():
106
+ answer = answers.get(qid, {})
107
+ qtype = definition["type"]
108
+ correct, error = score_answer(qtype, answer, gold.get(qid))
109
+ records.append(
110
+ {
111
+ "type": qtype,
112
+ "correct": correct,
113
+ "error": error,
114
+ "confidence": float(answer.get("confidence", 0.0) or 0.0),
115
+ }
116
+ )
117
+
118
+ buckets = {}
119
+ for record in records:
120
+ bucket = buckets.setdefault(
121
+ record["type"], {"confidences": [], "outcomes": [], "errors": [], "n": 0}
122
+ )
123
+ bucket["n"] += 1
124
+ if record["correct"] is not None:
125
+ bucket["confidences"].append(record["confidence"])
126
+ bucket["outcomes"].append(bool(record["correct"]))
127
+ if record["error"] is not None:
128
+ bucket["errors"].append(record["error"])
129
+
130
+ by_type = {}
131
+ for qtype, bucket in buckets.items():
132
+ entry = {"n": bucket["n"]}
133
+ if bucket["outcomes"]:
134
+ entry["accuracy"] = round(sum(bucket["outcomes"]) / len(bucket["outcomes"]), 4)
135
+ entry["ece"] = round(
136
+ expected_calibration_error(bucket["confidences"], bucket["outcomes"], bins), 4
137
+ )
138
+ if bucket["errors"]:
139
+ entry["mae"] = round(statistics.fmean(bucket["errors"]), 4)
140
+ by_type[qtype] = entry
141
+
142
+ confidences = [r["confidence"] for r in records if r["correct"] is not None]
143
+ outcomes = [bool(r["correct"]) for r in records if r["correct"] is not None]
144
+ latencies.sort()
145
+ overall = {
146
+ "items": len(items),
147
+ "questions": len(records),
148
+ "answers_scored": len(outcomes),
149
+ "accuracy": round(sum(outcomes) / len(outcomes), 4) if outcomes else None,
150
+ "ece": round(expected_calibration_error(confidences, outcomes, bins), 4)
151
+ if confidences
152
+ else None,
153
+ "latency_ms": {
154
+ "mean": round(statistics.fmean(latencies), 3) if latencies else None,
155
+ "p50": round(percentile(latencies, 0.5), 3) if latencies else None,
156
+ "p95": round(percentile(latencies, 0.95), 3) if latencies else None,
157
+ },
158
+ }
159
+ return {"by_type": by_type, "overall": overall}
160
+
161
+
162
+ def format_markdown(results, title="laya-onnx decision eval"):
163
+ """Render a results dict (from :func:`run_eval`) as Markdown."""
164
+ overall = results.get("overall", {})
165
+ latency = overall.get("latency_ms", {})
166
+ lines = [
167
+ f"### {title}",
168
+ "",
169
+ f"- items: **{overall.get('items')}** questions: **{overall.get('questions')}**",
170
+ f"- accuracy: **{overall.get('accuracy')}** ECE: **{overall.get('ece')}**",
171
+ f"- latency (per call, ms): mean {latency.get('mean')} / p50 {latency.get('p50')} / p95 {latency.get('p95')}",
172
+ "",
173
+ "| question type | n | accuracy | MAE | ECE |",
174
+ "| --- | ---: | ---: | ---: | ---: |",
175
+ ]
176
+ for qtype, entry in sorted(results.get("by_type", {}).items()):
177
+ lines.append(
178
+ "| {0} | {1} | {2} | {3} | {4} |".format(
179
+ qtype,
180
+ entry.get("n", ""),
181
+ entry.get("accuracy", "-"),
182
+ entry.get("mae", "-"),
183
+ entry.get("ece", "-"),
184
+ )
185
+ )
186
+ return "\n".join(lines)
laya_onnx/hub.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Resolve a local ONNX bundle or a Hugging Face snapshot."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from pathlib import Path, PurePosixPath
6
+
7
+ from huggingface_hub import snapshot_download
8
+
9
+ GRAPH_NAMES = ("laya.onnx", "model.onnx")
10
+ CONFIG_NAMES = ("onnx_config.json", "laya_config.json", "rl_agent_config.json")
11
+
12
+
13
+ def bundle_dir(model_id: str, revision=None) -> Path:
14
+ """A real, non-symlinked directory for a downloaded bundle.
15
+
16
+ ONNX Runtime validates that a model's external-data file stays inside the
17
+ model directory. The default Hugging Face cache stores snapshot entries as
18
+ symlinks into ``blobs/``, so that validation fails (seen on Windows);
19
+ downloading into a ``local_dir`` materialises regular files instead.
20
+ """
21
+ from huggingface_hub.constants import HF_HUB_CACHE
22
+
23
+ slug = str(model_id).replace("/", "--")
24
+ if revision:
25
+ slug += "@" + str(revision).replace("/", "--")
26
+ return Path(HF_HUB_CACHE) / "laya-onnx-bundles" / slug
27
+
28
+
29
+ def find_graph(path: Path) -> Path:
30
+ for name in GRAPH_NAMES:
31
+ candidate = path / name
32
+ if candidate.is_file():
33
+ return candidate
34
+ raise FileNotFoundError(f"No ONNX graph (laya.onnx / model.onnx) under {path}")
35
+
36
+
37
+ def find_config(path: Path) -> Path:
38
+ for name in CONFIG_NAMES:
39
+ candidate = path / name
40
+ if candidate.is_file():
41
+ return candidate
42
+ raise FileNotFoundError(f"No onnx_config.json / laya_config.json under {path}")
43
+
44
+
45
+ def resolve_bundle(model_id_or_path, *, token=None, subfolder=None, revision=None) -> Path:
46
+ if subfolder:
47
+ part = PurePosixPath(subfolder)
48
+ if part.is_absolute() or ".." in part.parts:
49
+ raise ValueError("subfolder must be a relative path inside the model repository")
50
+ path = Path(model_id_or_path).expanduser()
51
+ if path.exists():
52
+ if path.suffix == ".mlpackage" or str(path).endswith(".mlpackage"):
53
+ raise ValueError(
54
+ f"{path} looks like a Core ML package (.mlpackage). "
55
+ "laya-onnx expects an ONNX bundle (laya.onnx + tokenizer/). "
56
+ "Convert with `laya-onnx convert` or point to receptron/laya-onnx."
57
+ )
58
+ if subfolder:
59
+ path /= subfolder
60
+ find_graph(path)
61
+ return path
62
+ value = str(model_id_or_path)
63
+ if value.startswith(("/", "./", "../", "~")) or isinstance(model_id_or_path, Path):
64
+ raise FileNotFoundError(f"Local model directory does not exist: {value}")
65
+ prefix = subfolder.rstrip("/") + "/" if subfolder else ""
66
+ patterns = [
67
+ prefix + name
68
+ for name in (
69
+ "laya.onnx",
70
+ "laya.onnx.data",
71
+ "model.onnx",
72
+ "model.onnx.data",
73
+ "laya_config.json",
74
+ "onnx_config.json",
75
+ "rl_agent_config.json",
76
+ "tokenizer/*",
77
+ )
78
+ ]
79
+ path = Path(snapshot_download(value, token=token, revision=revision, allow_patterns=patterns, local_dir=bundle_dir(value, revision)))
80
+ if subfolder:
81
+ path /= subfolder
82
+ find_graph(path)
83
+ return path
laya_onnx/inputs.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Dynamic batch collation. Pad sequence length to a multiple of 16."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import numpy as np
6
+
7
+
8
+ def pad_length(length: int, multiple: int = 16, max_length: int | None = None) -> int:
9
+ padded = ((length + multiple - 1) // multiple) * multiple
10
+ if max_length is not None:
11
+ padded = min(max(padded, multiple), max_length)
12
+ return max(padded, multiple if length else 1)
13
+
14
+
15
+ def collate_items(
16
+ items,
17
+ pad_id,
18
+ *,
19
+ shape=None,
20
+ pad_to_multiple=16,
21
+ max_length=None,
22
+ marker_dtype=None,
23
+ ids_dtype=None,
24
+ ):
25
+ if not items:
26
+ raise ValueError("Cannot collate an empty batch")
27
+ if shape:
28
+ batch_size = int(shape.get("batch_size", len(items)))
29
+ max_length = int(shape.get("max_length", max_length or 512))
30
+ min_length = int(shape.get("min_length", pad_to_multiple))
31
+ max_options = int(shape.get("max_options", max(2, max(len(i["markers"]) for i in items))))
32
+ pad_to_multiple = min_length if min_length else pad_to_multiple
33
+ else:
34
+ batch_size = len(items)
35
+ max_options = max(2, max(len(i["markers"]) for i in items))
36
+
37
+ length = max(len(item["ids"]) for item in items)
38
+ if pad_to_multiple:
39
+ length = pad_length(length, pad_to_multiple, max_length)
40
+ if shape:
41
+ length = max(length, int(shape.get("min_length", length)))
42
+ if max_length is not None:
43
+ length = min(length, max_length)
44
+
45
+ id_dt = ids_dtype or np.int64
46
+ mk_dt = marker_dtype or np.int64
47
+ batch = {
48
+ "input_ids": np.full((batch_size, length), pad_id, dtype=id_dt),
49
+ "attention_mask": np.zeros((batch_size, length), dtype=id_dt),
50
+ "marker_pos": np.zeros((batch_size, max_options), dtype=mk_dt),
51
+ "marker_mask": np.zeros((batch_size, max_options), dtype=np.bool_),
52
+ "qtype": np.zeros((batch_size,), dtype=id_dt),
53
+ }
54
+ for i, item in enumerate(items[:batch_size]):
55
+ seq, markers = item["ids"], item["markers"]
56
+ seq = seq[:length]
57
+ batch["input_ids"][i, : len(seq)] = seq
58
+ batch["attention_mask"][i, : len(seq)] = 1
59
+ count = min(len(markers), max_options)
60
+ batch["marker_pos"][i, :count] = markers[:count]
61
+ batch["marker_mask"][i, :count] = True
62
+ batch["qtype"][i] = item["qtype"]
63
+ return batch
laya_onnx/optimize.py ADDED
@@ -0,0 +1,139 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Offline ONNX graph optimization for a PC backend (Intel CPU / 16 GB).
2
+
3
+ Order of work:
4
+ 1. Shape inference
5
+ 2. ONNX Runtime graph fusions (ORT_ENABLE_ALL) written back to disk
6
+ 3. Optional dynamic INT8 on MatMul/Gemm (best CPU win)
7
+ 4. Optional FP16 weights (useful on CUDA, usually slower on CPU)
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import json
13
+ import shutil
14
+ from pathlib import Path
15
+
16
+
17
+ def _find_graph(bundle: Path) -> Path:
18
+ for name in ("laya.onnx", "model.onnx"):
19
+ path = bundle / name if bundle.is_dir() else bundle
20
+ if bundle.is_file() and bundle.suffix == ".onnx":
21
+ return bundle
22
+ if path.is_file():
23
+ return path
24
+ raise FileNotFoundError(f"No ONNX graph under {bundle}")
25
+
26
+
27
+ def infer_shapes(graph: Path) -> Path:
28
+ import onnx
29
+
30
+ model = onnx.load(str(graph), load_external_data=True)
31
+ try:
32
+ model = onnx.shape_inference.infer_shapes(model)
33
+ onnx.save(model, str(graph))
34
+ except Exception:
35
+ pass
36
+ return graph
37
+
38
+
39
+ def fuse_graph(graph: Path, *, threads: int | None = None) -> Path:
40
+ import onnxruntime as ort
41
+
42
+ tmp = graph.with_suffix(".opt.onnx")
43
+ opts = ort.SessionOptions()
44
+ opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
45
+ opts.optimized_model_filepath = str(tmp)
46
+ if threads:
47
+ opts.intra_op_num_threads = int(threads)
48
+ opts.inter_op_num_threads = 1
49
+ sess = ort.InferenceSession(str(graph), sess_options=opts, providers=["CPUExecutionProvider"])
50
+ del sess
51
+ if tmp.is_file() and tmp.stat().st_size > 0:
52
+ tmp.replace(graph)
53
+ extra = Path(str(tmp) + ".data")
54
+ dest_extra = Path(str(graph) + ".data")
55
+ if extra.is_file():
56
+ extra.replace(dest_extra)
57
+ return graph
58
+
59
+
60
+ def quantize_dynamic_int8(graph: Path) -> Path:
61
+ from onnxruntime.quantization import QuantType, quantize_dynamic
62
+
63
+ out = graph.with_suffix(".int8.onnx")
64
+ quantize_dynamic(
65
+ str(graph),
66
+ str(out),
67
+ weight_type=QuantType.QInt8,
68
+ extra_options={"WeightSymmetric": True, "EnableSubgraph": True},
69
+ )
70
+ out.replace(graph)
71
+ return graph
72
+
73
+
74
+ def convert_weights_fp16(graph: Path) -> Path:
75
+ import numpy as np
76
+ import onnx
77
+ from onnx import TensorProto, numpy_helper
78
+
79
+ model = onnx.load(str(graph), load_external_data=True)
80
+ for init in model.graph.initializer:
81
+ if init.data_type != TensorProto.FLOAT:
82
+ continue
83
+ array = numpy_helper.to_array(init).astype(np.float16)
84
+ new = numpy_helper.from_array(array, init.name)
85
+ init.CopyFrom(new)
86
+ onnx.save(model, str(graph))
87
+ return graph
88
+
89
+
90
+ def optimize_bundle(source, output=None, *, precision="int8", threads=None, fuse=True):
91
+ source = Path(source).expanduser()
92
+ src_graph = _find_graph(source)
93
+ if output is None:
94
+ dest_graph = src_graph
95
+ dest_root = src_graph.parent
96
+ else:
97
+ output = Path(output).expanduser()
98
+ if output.suffix == ".onnx":
99
+ dest_graph = output
100
+ dest_root = output.parent
101
+ dest_root.mkdir(parents=True, exist_ok=True)
102
+ shutil.copy2(src_graph, dest_graph)
103
+ else:
104
+ if output.exists() and output.resolve() != source.resolve():
105
+ raise FileExistsError(f"Output already exists: {output}")
106
+ output.mkdir(parents=True, exist_ok=True)
107
+ if source.is_dir() and output.resolve() != source.resolve():
108
+ for item in source.iterdir():
109
+ target = output / item.name
110
+ if item.is_dir():
111
+ shutil.copytree(item, target, dirs_exist_ok=True)
112
+ else:
113
+ shutil.copy2(item, target)
114
+ dest_graph = output / src_graph.name
115
+ if not dest_graph.is_file():
116
+ shutil.copy2(src_graph, dest_graph)
117
+ dest_root = output
118
+
119
+ infer_shapes(dest_graph)
120
+ if fuse:
121
+ fuse_graph(dest_graph, threads=threads)
122
+ if precision == "int8":
123
+ quantize_dynamic_int8(dest_graph)
124
+ infer_shapes(dest_graph)
125
+ if fuse:
126
+ fuse_graph(dest_graph, threads=threads)
127
+ elif precision == "fp16":
128
+ convert_weights_fp16(dest_graph)
129
+
130
+ manifest = dest_root / "onnx_config.json"
131
+ payload = {}
132
+ if manifest.is_file():
133
+ try:
134
+ payload = json.loads(manifest.read_text())
135
+ except json.JSONDecodeError:
136
+ payload = {}
137
+ payload.update({"format": payload.get("format", "laya-onnx"), "optimized": True, "precision": precision, "graph": dest_graph.name})
138
+ manifest.write_text(json.dumps(payload, indent=2) + "\n")
139
+ return dest_graph
laya_onnx/snake/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Local Snake demonstration for Laya ONNX."""
laya_onnx/snake/__main__.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ from .cli import main
2
+
3
+ raise SystemExit(main())
laya_onnx/snake/benchmark.py ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Headless throughput benchmark for the Snake policy (no display, no TTY)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import statistics
8
+ import time
9
+
10
+ from .game import SnakeGame
11
+ from .policy import LayaPolicy
12
+
13
+
14
+ def _percentile(sorted_samples, q):
15
+ if not sorted_samples:
16
+ return None
17
+ index = min(len(sorted_samples) - 1, int(len(sorted_samples) * q))
18
+ return sorted_samples[index]
19
+
20
+
21
+ def main(argv=None) -> int:
22
+ parser = argparse.ArgumentParser(prog="laya-onnx-snake benchmark", description=__doc__)
23
+ parser.add_argument("--model", help="Local model directory or an already cached Hub ID")
24
+ parser.add_argument("--prompt", choices=("compact", "detailed"), default="compact")
25
+ parser.add_argument("--optimize", action="store_true", help="Enable 16-token buckets")
26
+ parser.add_argument("--games", type=int, default=3)
27
+ parser.add_argument("--steps", type=int, default=200, help="Max steps per game")
28
+ parser.add_argument("--width", type=int, default=16)
29
+ parser.add_argument("--height", type=int, default=12)
30
+ parser.add_argument("--seed", type=int, default=7)
31
+ parser.add_argument("--initial-length", type=int, default=6)
32
+ parser.add_argument("--warmup", type=int, default=6, help="Warm-up decisions per game")
33
+ parser.add_argument("--unassisted", action="store_true")
34
+ args = parser.parse_args(argv)
35
+ if args.games < 1:
36
+ parser.error("--games must be positive")
37
+ if args.steps < 1:
38
+ parser.error("--steps must be positive")
39
+ if args.warmup < 0:
40
+ parser.error("--warmup must be non-negative")
41
+
42
+ policy = LayaPolicy(
43
+ args.model, guarded=not args.unassisted, prompt=args.prompt, optimize=args.optimize
44
+ )
45
+ samples = []
46
+ total_steps = deaths = wins = interventions = best = 0
47
+ started = time.perf_counter()
48
+ for game_index in range(args.games):
49
+ game = SnakeGame(args.width, args.height, args.seed + game_index, args.initial_length)
50
+ for _ in range(args.warmup):
51
+ if not game.alive or game.won:
52
+ break
53
+ game.step(policy.decide(game).executed)
54
+ played = 0
55
+ while played < args.steps and game.alive and not game.won:
56
+ decision = policy.decide(game)
57
+ samples.append(decision.inference_ms)
58
+ interventions += decision.intervened
59
+ game.step(decision.executed)
60
+ played += 1
61
+ total_steps += 1
62
+ best = max(best, game.score)
63
+ if game.won:
64
+ wins += 1
65
+ elif not game.alive:
66
+ deaths += 1
67
+ seconds = time.perf_counter() - started
68
+ samples.sort()
69
+ summary = {
70
+ "model": policy.metadata,
71
+ "games": args.games,
72
+ "steps": total_steps,
73
+ "seconds": round(seconds, 3),
74
+ "steps_per_second": round(total_steps / seconds, 2) if seconds else 0,
75
+ "best_score": best,
76
+ "wins": wins,
77
+ "deaths": deaths,
78
+ "interventions": interventions,
79
+ "mean_inference_ms": round(statistics.fmean(samples), 3) if samples else None,
80
+ "p50_inference_ms": round(_percentile(samples, 0.50), 3) if samples else None,
81
+ "p95_inference_ms": round(_percentile(samples, 0.95), 3) if samples else None,
82
+ }
83
+ print(json.dumps(summary, indent=2))
84
+ return 0
85
+
86
+
87
+ if __name__ == "__main__":
88
+ raise SystemExit(main())
laya_onnx/snake/cli.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Run the live terminal demo. Also supports benchmark and export subcommands."""
2
+
3
+ import argparse
4
+ import json
5
+ import sys
6
+ import time
7
+ from collections import deque
8
+ from contextlib import nullcontext
9
+ from datetime import datetime, timezone
10
+ from pathlib import Path
11
+
12
+ from rich.console import Console
13
+ from rich.live import Live
14
+
15
+ from .game import SnakeGame
16
+ from .policy import LayaPolicy
17
+ from .ui import BG, compose, layout_size
18
+
19
+
20
+ class Keyboard:
21
+ """Non-blocking raw-key reader: cbreak on POSIX, msvcrt on Windows."""
22
+
23
+ def __enter__(self):
24
+ self._posix = sys.platform != "win32"
25
+ self.saved = None
26
+ if sys.stdin.isatty():
27
+ if self._posix:
28
+ import termios
29
+ import tty
30
+
31
+ self.fd = sys.stdin.fileno()
32
+ self.saved = termios.tcgetattr(self.fd)
33
+ tty.setcbreak(self.fd)
34
+ else:
35
+ import msvcrt # noqa: F401 (present on Windows)
36
+
37
+ self.saved = True
38
+ return self
39
+
40
+ def read(self):
41
+ if not self.saved:
42
+ return ""
43
+ if self._posix:
44
+ import os
45
+ import select
46
+
47
+ if select.select([sys.stdin], [], [], 0)[0]:
48
+ return os.read(self.fd, 128).decode(errors="ignore")
49
+ return ""
50
+ import msvcrt
51
+
52
+ chars = []
53
+ while msvcrt.kbhit():
54
+ chars.append(msvcrt.getwch())
55
+ return "".join(chars)
56
+
57
+ def __exit__(self, *_):
58
+ if self._posix and self.saved:
59
+ import termios
60
+
61
+ termios.tcsetattr(self.fd, termios.TCSADRAIN, self.saved)
62
+
63
+
64
+ def positive(value):
65
+ result = float(value)
66
+ if not 0 < result < float("inf"):
67
+ raise argparse.ArgumentTypeError("Expected a positive finite number")
68
+ return result
69
+
70
+
71
+ def play(argv=None):
72
+ parser = argparse.ArgumentParser(
73
+ prog="laya-onnx-snake",
74
+ description=__doc__,
75
+ epilog="Also: laya-onnx-snake benchmark --help | laya-onnx-snake export --help",
76
+ )
77
+ parser.add_argument("--model", help="Local model directory or an already cached Hub ID")
78
+ parser.add_argument("--prompt", choices=("compact", "detailed"), default="compact")
79
+ parser.add_argument(
80
+ "--optimize",
81
+ action="store_true",
82
+ help="Enable 16-token buckets",
83
+ )
84
+ parser.add_argument("--width", type=int, default=24)
85
+ parser.add_argument("--height", type=int, default=16)
86
+ parser.add_argument("--seed", type=int, default=7)
87
+ parser.add_argument("--initial-length", type=int, default=6)
88
+ parser.add_argument("--fps", type=positive, default=12)
89
+ parser.add_argument("--max-speed", action="store_true")
90
+ parser.add_argument("--duration", type=positive)
91
+ parser.add_argument("--steps", type=int)
92
+ parser.add_argument("--unassisted", action="store_true")
93
+ parser.add_argument("--record", type=Path)
94
+ parser.add_argument("--headless", action="store_true")
95
+ parser.add_argument("--no-alt-screen", action="store_true")
96
+ args = parser.parse_args(argv)
97
+ if args.steps is not None and args.steps < 1:
98
+ parser.error("--steps must be positive")
99
+ try:
100
+ game = SnakeGame(args.width, args.height, args.seed, args.initial_length)
101
+ except ValueError as error:
102
+ parser.error(str(error))
103
+ console = Console(style=f"on {BG}", highlight=False)
104
+ if not args.headless and not console.is_terminal:
105
+ parser.error("Interactive display needs a TTY. Use --headless.")
106
+ print("Loading local ONNX bundle; Hub only if the path is a repo id...", file=sys.stderr)
107
+ policy = LayaPolicy(
108
+ args.model, guarded=not args.unassisted, prompt=args.prompt, optimize=args.optimize
109
+ )
110
+ warm = SnakeGame(args.width, args.height, args.seed + 10000, args.initial_length)
111
+ for _ in range(6):
112
+ decision = policy.decide(warm)
113
+ warm.step(decision.executed)
114
+ if not warm.alive:
115
+ break
116
+ record = None
117
+ if args.record:
118
+ args.record.parent.mkdir(parents=True, exist_ok=True)
119
+ record = args.record.open("x")
120
+ record.write(
121
+ json.dumps(
122
+ {
123
+ "type": "metadata",
124
+ "format": "laya-onnx-snake-v1",
125
+ "created_utc": datetime.now(timezone.utc).isoformat(),
126
+ "model": policy.metadata,
127
+ "settings": {
128
+ k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()
129
+ },
130
+ }
131
+ )
132
+ + "\n"
133
+ )
134
+ started = time.perf_counter()
135
+ stats = {
136
+ "hardware": policy.metadata["hardware"],
137
+ "guarded": policy.guarded,
138
+ "interventions": 0,
139
+ "best": 0,
140
+ "round": 1,
141
+ "paused": False,
142
+ "elapsed": 0,
143
+ "steps_per_second": 0,
144
+ }
145
+ calls = total_steps = deaths = 0
146
+ timestamps = deque(maxlen=60)
147
+ inference = []
148
+ displayed_board, displayed_decision = game.snapshot(), {}
149
+ live = (
150
+ Live(console=console, screen=not args.no_alt_screen, auto_refresh=False, vertical_overflow="crop")
151
+ if not args.headless
152
+ else None
153
+ )
154
+ try:
155
+ with Keyboard() as keys, live if live else nullcontext():
156
+ while True:
157
+ now = time.perf_counter()
158
+ if (args.duration and now - started >= args.duration) or (
159
+ args.steps and total_steps >= args.steps
160
+ ):
161
+ break
162
+ pressed = keys.read().lower()
163
+ if "q" in pressed or "\x03" in pressed:
164
+ break
165
+ if " " in pressed:
166
+ stats["paused"] = not stats["paused"]
167
+ if "+" in pressed:
168
+ args.fps = min(240, args.fps + 2)
169
+ if "-" in pressed:
170
+ args.fps = max(1, args.fps - 2)
171
+ if "r" in pressed:
172
+ stats["round"] += 1
173
+ game = SnakeGame(args.width, args.height, args.seed + stats["round"] - 1, args.initial_length)
174
+ displayed_board, displayed_decision = game.snapshot(), {}
175
+ if stats["paused"]:
176
+ stats["elapsed"] = now - started
177
+ if live:
178
+ live.update(compose(displayed_board, displayed_decision, stats).rich_text(), refresh=True)
179
+ time.sleep(0.03)
180
+ continue
181
+ minimum_width, minimum_height = layout_size(game.width, game.height)
182
+ if live and (console.width < minimum_width or console.height < minimum_height):
183
+ live.update(
184
+ f"Resize terminal to at least {minimum_width}x{minimum_height}. Q quits.",
185
+ refresh=True,
186
+ )
187
+ time.sleep(0.1)
188
+ continue
189
+ decision = policy.decide(game)
190
+ calls += 1
191
+ inference.append(decision.inference_ms)
192
+ stats["interventions"] += decision.intervened
193
+ shown = time.perf_counter()
194
+ timestamps.append(shown)
195
+ stats["elapsed"] = shown - started
196
+ stats["steps_per_second"] = (
197
+ (len(timestamps) - 1) / (timestamps[-1] - timestamps[0]) if len(timestamps) > 1 else 0
198
+ )
199
+ stats["best"] = max(stats["best"], game.score)
200
+ board = game.snapshot()
201
+ displayed_board, displayed_decision = board, decision.to_dict()
202
+ if live:
203
+ live.update(compose(board, decision.to_dict(), stats).rich_text(), refresh=True)
204
+ if record:
205
+ record.write(
206
+ json.dumps(
207
+ {"type": "frame", "at": shown - started, "game": board, "decision": decision.to_dict(), "stats": dict(stats)},
208
+ separators=(",", ":"),
209
+ )
210
+ + "\n"
211
+ )
212
+ if not args.max_speed:
213
+ remaining = 1 / args.fps - (time.perf_counter() - now)
214
+ if remaining > 0:
215
+ time.sleep(remaining)
216
+ game.step(decision.executed)
217
+ total_steps += 1
218
+ stats["best"] = max(stats["best"], game.score)
219
+ if not game.alive or game.won:
220
+ deaths += not game.alive
221
+ if args.unassisted:
222
+ break
223
+ if live:
224
+ live.update(compose(game.snapshot(), {}, stats).rich_text(), refresh=True)
225
+ time.sleep(1)
226
+ stats["round"] += 1
227
+ game = SnakeGame(args.width, args.height, args.seed + stats["round"] - 1, args.initial_length)
228
+ except KeyboardInterrupt:
229
+ pass
230
+ finally:
231
+ elapsed = time.perf_counter() - started
232
+ summary = {
233
+ "steps": total_steps,
234
+ "inference_calls": calls,
235
+ "seconds": elapsed,
236
+ "steps_per_second": total_steps / elapsed if elapsed else 0,
237
+ "score": game.score,
238
+ "length": len(game.body),
239
+ "best_score": stats["best"],
240
+ "interventions": stats["interventions"],
241
+ "deaths": deaths,
242
+ "guarded": policy.guarded,
243
+ "network": "offline",
244
+ "mean_inference_ms": sum(inference) / len(inference) if inference else None,
245
+ }
246
+ if record:
247
+ record.write(json.dumps({"type": "end", "summary": summary, "game": game.snapshot()}) + "\n")
248
+ record.close()
249
+ print(json.dumps(summary, indent=2))
250
+ return 0
251
+
252
+
253
+ def main(argv=None):
254
+ argv = list(sys.argv[1:] if argv is None else argv)
255
+ if argv and argv[0] == "benchmark":
256
+ from .benchmark import main as benchmark
257
+
258
+ return benchmark(argv[1:])
259
+ if argv and argv[0] == "export":
260
+ from .replay import main as export
261
+
262
+ return export(argv[1:])
263
+ return play(argv)
laya_onnx/snake/game.py ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Deterministic, JSON-serializable Snake game used by the terminal demo.
2
+
3
+ No third-party imports: the game must be runnable (and testable) with an
4
+ empty environment. The board is a ``width`` x ``height`` grid with the head
5
+ first in ``body``. ``seed`` fully determines the food sequence.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import random
11
+
12
+ DIRECTIONS = ("UP", "DOWN", "LEFT", "RIGHT")
13
+ DELTAS = {"UP": (0, -1), "DOWN": (0, 1), "LEFT": (-1, 0), "RIGHT": (1, 0)}
14
+ MIN_SIZE = 4
15
+
16
+
17
+ class SnakeGame:
18
+ def __init__(self, width: int = 24, height: int = 16, seed: int = 7, initial_length: int = 6):
19
+ if not isinstance(width, int) or not isinstance(height, int) or isinstance(width, bool) or isinstance(height, bool):
20
+ raise ValueError("width and height must be integers")
21
+ if width < MIN_SIZE or height < MIN_SIZE:
22
+ raise ValueError(f"width and height must be at least {MIN_SIZE}")
23
+ area = width * height
24
+ if (
25
+ not isinstance(initial_length, int)
26
+ or isinstance(initial_length, bool)
27
+ or not 2 <= initial_length < area
28
+ ):
29
+ raise ValueError(f"initial_length must be an integer between 2 and {area - 1}")
30
+ self.width = width
31
+ self.height = height
32
+ self.seed = seed
33
+ self.rng = random.Random(seed)
34
+ path = self._cycle()
35
+ start = path.index((width // 2, height // 2))
36
+ span = len(path)
37
+ # Head sits at the centre; the body trails backwards along the path.
38
+ self.body = [path[(start - i) % span] for i in range(initial_length)]
39
+ self.steps = 0
40
+ self.score = 0
41
+ self.alive = True
42
+ self.won = False
43
+ self.food = self._spawn()
44
+
45
+ # -- internals ---------------------------------------------------------
46
+ def _cycle(self):
47
+ """Boustrophedon Hamiltonian path over the grid (row by row)."""
48
+ path = []
49
+ for y in range(self.height):
50
+ xs = range(self.width) if y % 2 == 0 else range(self.width - 1, -1, -1)
51
+ for x in xs:
52
+ path.append((x, y))
53
+ return path
54
+
55
+ def _spawn(self):
56
+ occupied = set(self.body)
57
+ free = [
58
+ (x, y)
59
+ for y in range(self.height)
60
+ for x in range(self.width)
61
+ if (x, y) not in occupied
62
+ ]
63
+ if not free:
64
+ return None
65
+ return self.rng.choice(free)
66
+
67
+ def _target(self, direction):
68
+ dx, dy = DELTAS[direction]
69
+ hx, hy = self.body[0]
70
+ return hx + dx, hy + dy
71
+
72
+ # -- public API --------------------------------------------------------
73
+ def legal(self, direction) -> bool:
74
+ """True if ``direction`` can be played without dying. Pure."""
75
+ if not self.alive or self.won or direction not in DELTAS:
76
+ return False
77
+ nx, ny = self._target(direction)
78
+ if not (0 <= nx < self.width and 0 <= ny < self.height):
79
+ return False
80
+ grows = self.food is not None and (nx, ny) == self.food
81
+ # Without growth the tail cell vacates, so stepping onto it is safe.
82
+ blocked = set(self.body if grows else self.body[:-1])
83
+ return (nx, ny) not in blocked
84
+
85
+ def step(self, direction):
86
+ """Advance one tick. Returns ``ok`` / ``eat`` / ``dead`` / ``over``."""
87
+ if direction not in DELTAS:
88
+ raise ValueError(f"direction must be one of {DIRECTIONS}, got {direction!r}")
89
+ if not self.alive:
90
+ return "dead"
91
+ if self.won:
92
+ return "over"
93
+ if not self.legal(direction):
94
+ self.alive = False
95
+ return "dead"
96
+ nx, ny = self._target(direction)
97
+ eats = self.food is not None and (nx, ny) == self.food
98
+ self.body.insert(0, (nx, ny))
99
+ if eats:
100
+ self.score += 1
101
+ self.food = self._spawn()
102
+ else:
103
+ self.body.pop()
104
+ self.steps += 1
105
+ if self.food is None or len(self.body) >= self.width * self.height:
106
+ self.won = True
107
+ return "eat" if eats else "ok"
108
+
109
+ def snapshot(self) -> dict:
110
+ """JSON-serializable view of the current state."""
111
+ return {
112
+ "width": self.width,
113
+ "height": self.height,
114
+ "head": list(self.body[0]),
115
+ "body": [list(cell) for cell in self.body],
116
+ "food": list(self.food) if self.food is not None else None,
117
+ "score": self.score,
118
+ "length": len(self.body),
119
+ "steps": self.steps,
120
+ "alive": self.alive,
121
+ "won": self.won,
122
+ }
laya_onnx/snake/policy.py ADDED
@@ -0,0 +1,142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Laya-backed Snake policy.
2
+
3
+ Each tick builds a typed System-1 decision: the local ONNX Agent is asked
4
+ (a ``choice`` over UP/DOWN/LEFT/RIGHT) and its answer becomes the move. In
5
+ ``guarded`` mode an illegal move is replaced by the best legal alternative
6
+ and the tick is counted as an intervention.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import os
12
+ import platform
13
+ import time
14
+
15
+ from .. import load
16
+ from .game import DIRECTIONS
17
+
18
+
19
+ def _hardware() -> str:
20
+ chip = platform.processor() or platform.machine() or "cpu"
21
+ return f"{platform.system()} {chip} x{os.cpu_count() or 1}"
22
+
23
+
24
+ class Decision:
25
+ def __init__(self, chosen, executed, intervened, probabilities, inference_ms, confidence=None):
26
+ self.chosen = chosen
27
+ self.executed = executed
28
+ self.intervened = int(intervened)
29
+ self.probabilities = probabilities
30
+ self.inference_ms = float(inference_ms)
31
+ self.confidence = confidence
32
+
33
+ def to_dict(self) -> dict:
34
+ payload = {
35
+ "chosen": self.chosen,
36
+ "executed": self.executed,
37
+ "intervened": bool(self.intervened),
38
+ "probabilities": self.probabilities,
39
+ "inference_ms": round(self.inference_ms, 3),
40
+ }
41
+ if self.confidence is not None:
42
+ payload["confidence"] = self.confidence
43
+ return payload
44
+
45
+
46
+ class LayaPolicy:
47
+ def __init__(
48
+ self,
49
+ model=None,
50
+ *,
51
+ guarded: bool = True,
52
+ prompt: str = "compact",
53
+ optimize: bool = False,
54
+ providers: str = "cpu",
55
+ threads=None,
56
+ ):
57
+ if prompt not in ("compact", "detailed"):
58
+ raise ValueError("prompt must be 'compact' or 'detailed'")
59
+ self.model_id = str(model or "receptron/laya-onnx")
60
+ self.guarded = bool(guarded)
61
+ self.prompt = prompt
62
+ self.optimize = bool(optimize)
63
+ # "optimize" keeps 16-token padding buckets; otherwise skip padding for latency.
64
+ self.agent = load(
65
+ model or "receptron/laya-onnx",
66
+ providers=providers,
67
+ threads=threads,
68
+ pad_to_multiple=16 if self.optimize else None,
69
+ )
70
+ self.metadata = {
71
+ "hardware": _hardware(),
72
+ "model": self.model_id,
73
+ "providers": self.agent.providers,
74
+ "prompt": self.prompt,
75
+ "optimized": self.optimize,
76
+ }
77
+
78
+ # -- prompt building ---------------------------------------------------
79
+ def _state(self, game) -> str:
80
+ snap = game.snapshot()
81
+ head, food = snap["head"], snap["food"]
82
+ if self.prompt == "detailed":
83
+ rows = []
84
+ body = {tuple(cell) for cell in snap["body"]}
85
+ for y in range(snap["height"]):
86
+ line = []
87
+ for x in range(snap["width"]):
88
+ point = (x, y)
89
+ if point == tuple(head):
90
+ line.append("#")
91
+ elif food is not None and point == tuple(food):
92
+ line.append("*")
93
+ elif point in body:
94
+ line.append("o")
95
+ else:
96
+ line.append(".")
97
+ rows.append("".join(line))
98
+ grid = "\n".join(rows)
99
+ return (
100
+ f"Snake {snap['width']}x{snap['height']}\n{grid}\n"
101
+ f"#=head o=body *=food head={head} food={food} "
102
+ f"length={snap['length']} score={snap['score']}"
103
+ )
104
+ return (
105
+ f"Snake {snap['width']}x{snap['height']} head={head} food={food} "
106
+ f"length={snap['length']} score={snap['score']} alive={snap['alive']}"
107
+ )
108
+
109
+ def _questions(self) -> dict:
110
+ return {
111
+ "move": {
112
+ "type": "choice",
113
+ "instructions": (
114
+ "Pick the snake's next move. Move toward the food and never "
115
+ "hit the wall or the snake's own body."
116
+ ),
117
+ "criteria": {direction: direction.lower() for direction in DIRECTIONS},
118
+ }
119
+ }
120
+
121
+ # -- decision ----------------------------------------------------------
122
+ def decide(self, game) -> Decision:
123
+ started = time.perf_counter()
124
+ out = self.agent.predict(self._state(game), self._questions())
125
+ inference_ms = (time.perf_counter() - started) * 1000.0
126
+ answer = (out.get("answers") or {}).get("move") or {}
127
+ probabilities = answer.get("probabilities") or {}
128
+ ranked = sorted(
129
+ DIRECTIONS, key=lambda d: float(probabilities.get(d, 0.0)), reverse=True
130
+ )
131
+ chosen = answer.get("choice")
132
+ if chosen not in DIRECTIONS:
133
+ chosen = ranked[0]
134
+ executed, intervened = chosen, 0
135
+ if self.guarded and not game.legal(chosen):
136
+ fallback = next((d for d in ranked if game.legal(d)), None)
137
+ if fallback is None:
138
+ fallback = next((d for d in DIRECTIONS if game.legal(d)), chosen)
139
+ executed, intervened = fallback, 1
140
+ return Decision(
141
+ chosen, executed, intervened, probabilities, inference_ms, answer.get("confidence")
142
+ )
laya_onnx/snake/replay.py ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Replay or export a recorded Snake run (the JSONL written by ``--record``)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ from pathlib import Path
8
+
9
+ from rich.console import Console
10
+ from rich.table import Table
11
+
12
+ _TEMPLATE = """<!DOCTYPE html>
13
+ <html lang="en"><head><meta charset="UTF-8"/><meta name="viewport" content="width=device-width,initial-scale=1"/>
14
+ <title>laya-onnx snake replay</title>
15
+ <style>
16
+ html,body{margin:0;background:#070b12;color:#e2e8f0;font-family:system-ui,sans-serif}
17
+ .wrap{max-width:680px;margin:0 auto;padding:20px}
18
+ canvas{display:block;width:100%;max-width:640px;background:#020617;border:2px solid #34d399;border-radius:10px}
19
+ .row{display:flex;gap:10px;align-items:center;margin-top:10px}
20
+ button{padding:8px 12px;border:0;border-radius:8px;background:#334155;color:#fff}
21
+ #play{background:#34d399;color:#052e16;font-weight:700}
22
+ #st{font-size:13px;color:#94a3b8;margin-top:8px}
23
+ </style></head><body><div class="wrap">
24
+ <h1>Snake replay</h1><p id="model"></p>
25
+ <canvas id="c" width="640" height="420"></canvas>
26
+ <div class="row"><button id="play">Pause</button><button id="reset">Reset</button>
27
+ <input id="seek" type="range" min="0" value="0" style="flex:1"/></div>
28
+ <div id="st">-</div>
29
+ <script>
30
+ var FRAMES=__FRAMES__, DECISIONS=__DECISIONS__, MODEL=__MODEL__;
31
+ var i=0, playing=true, timer=null, c=document.getElementById("c"), ctx=c.getContext("2d");
32
+ document.getElementById("model").textContent="model: "+(MODEL.model||"?")+" frames: "+FRAMES.length;
33
+ function draw(){
34
+ var g=FRAMES[i]||{}, W=g.width||1, H=g.height||1, cw=c.width/W, ch=c.height/H;
35
+ ctx.fillStyle="#020617"; ctx.fillRect(0,0,c.width,c.height);
36
+ ctx.strokeStyle="#1e293b";
37
+ for(var x=0;x<=W;x++){ctx.beginPath();ctx.moveTo(x*cw,0);ctx.lineTo(x*cw,c.height);ctx.stroke();}
38
+ for(var y=0;y<=H;y++){ctx.beginPath();ctx.moveTo(0,y*ch);ctx.lineTo(c.width,y*ch);ctx.stroke();}
39
+ var f=g.food; if(f){ctx.fillStyle="#fbbf24";ctx.fillRect(f[0]*cw+2,f[1]*ch+2,cw-4,ch-4);}
40
+ var b=g.body||[];
41
+ for(var k=b.length-1;k>=0;k--){ctx.fillStyle=k===0?"#6ee7b7":"#059669";ctx.fillRect(b[k][0]*cw+2,b[k][1]*ch+2,cw-4,ch-4);}
42
+ var d=DECISIONS[i]||{};
43
+ document.getElementById("st").textContent="frame "+(i+1)+"/"+FRAMES.length+" move "+(d.executed||"-")+(d.intervened?" (guarded)":"")+" score "+(g.score||0)+" len "+(g.length||0);
44
+ document.getElementById("seek").value=i;
45
+ }
46
+ function tick(){if(!playing)return; i=(i+1)%Math.max(FRAMES.length,1); draw();}
47
+ function start(){if(timer)return; playing=true; document.getElementById("play").textContent="Pause"; timer=setInterval(tick,120);}
48
+ function stop(){playing=false; document.getElementById("play").textContent="Play"; clearInterval(timer); timer=null;}
49
+ document.getElementById("play").onclick=function(){playing?stop():start();};
50
+ document.getElementById("reset").onclick=function(){i=0;draw();};
51
+ document.getElementById("seek").oninput=function(e){i=Number(e.target.value);draw();};
52
+ var s=document.getElementById("seek"); s.max=Math.max(FRAMES.length-1,0);
53
+ draw(); start();
54
+ </script></div></body></html>
55
+ """
56
+
57
+
58
+ def _read_record(path: Path):
59
+ meta, frames, end = None, [], None
60
+ with path.open(encoding="utf-8") as handle:
61
+ for lineno, line in enumerate(handle, 1):
62
+ line = line.strip()
63
+ if not line:
64
+ continue
65
+ try:
66
+ record = json.loads(line)
67
+ except json.JSONDecodeError as exc:
68
+ raise SystemExit(f"{path}: line {lineno} is not valid JSON: {exc}")
69
+ kind = record.get("type")
70
+ if kind == "metadata":
71
+ meta = record
72
+ elif kind == "frame":
73
+ frames.append(record)
74
+ elif kind == "end":
75
+ end = record
76
+ if meta is None:
77
+ raise SystemExit(f"{path}: missing metadata header (is this a --record file?)")
78
+ return meta, frames, end
79
+
80
+
81
+ def _render_html(meta: dict, frames: list) -> str:
82
+ games = [frame.get("game") for frame in frames]
83
+ decisions = [frame.get("decision") for frame in frames]
84
+ return (
85
+ _TEMPLATE.replace("__FRAMES__", json.dumps(games))
86
+ .replace("__DECISIONS__", json.dumps(decisions))
87
+ .replace("__MODEL__", json.dumps(meta.get("model") or {}))
88
+ )
89
+
90
+
91
+ def main(argv=None) -> int:
92
+ parser = argparse.ArgumentParser(prog="laya-onnx-snake export", description=__doc__)
93
+ parser.add_argument("record", type=Path, help="JSONL file written by --record")
94
+ parser.add_argument("--limit", type=int, help="Show only the first N frames in the table")
95
+ parser.add_argument("--html", type=Path, help="Write a self-contained animated HTML replay")
96
+ parser.add_argument("--json", action="store_true", help="Print the summary JSON only")
97
+ args = parser.parse_args(argv)
98
+ if not args.record.is_file():
99
+ parser.error(f"no such record: {args.record}")
100
+
101
+ meta, frames, end = _read_record(args.record)
102
+ summary = (end or {}).get("summary") or {}
103
+ if args.json:
104
+ print(json.dumps(summary, indent=2))
105
+ return 0
106
+
107
+ console = Console()
108
+ model = (meta.get("model") or {}).get("model", "?")
109
+ console.print(
110
+ f"[bold]{args.record}[/] [dim]{meta.get('format', '?')}[/] "
111
+ f"model={model} frames={len(frames)} created={meta.get('created_utc', '?')}"
112
+ )
113
+ if summary:
114
+ console.print(json.dumps(summary, indent=2))
115
+ shown = frames if args.limit is None else frames[: args.limit]
116
+ if shown:
117
+ table = Table(title="frames")
118
+ for column in ("at", "executed", "chosen", "guarded", "score", "length"):
119
+ table.add_column(column)
120
+ for frame in shown:
121
+ decision = frame.get("decision") or {}
122
+ game = frame.get("game") or {}
123
+ table.add_row(
124
+ f"{float(frame.get('at', 0.0)):.2f}",
125
+ str(decision.get("executed", "")),
126
+ str(decision.get("chosen", "")),
127
+ "yes" if decision.get("intervened") else "",
128
+ str(game.get("score", "")),
129
+ str(game.get("length", "")),
130
+ )
131
+ console.print(table)
132
+ if args.html:
133
+ args.html.parent.mkdir(parents=True, exist_ok=True)
134
+ args.html.write_text(_render_html(meta, frames), encoding="utf-8")
135
+ console.print(f"[green]wrote[/] {args.html}")
136
+ return 0
137
+
138
+
139
+ if __name__ == "__main__":
140
+ raise SystemExit(main())
laya_onnx/snake/ui.py ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Rich rendering for the terminal Snake demo.
2
+
3
+ ``compose`` returns a small wrapper whose ``rich_text()`` yields a
4
+ ``rich.text.Text`` renderable, matching how ``snake/cli.py`` drives
5
+ ``rich.live.Live``. Kept import-light so it can be unit-tested headlessly.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from rich.text import Text
11
+
12
+ BG = "black"
13
+
14
+ HEAD_GLYPH = "█"
15
+ BODY_GLYPH = "▪"
16
+ FOOD_GLYPH = "◆"
17
+ EMPTY_GLYPH = "·"
18
+
19
+ HEAD_STYLE = "bold bright_green"
20
+ BODY_STYLE = "green"
21
+ FOOD_STYLE = "bold yellow"
22
+ EMPTY_STYLE = "grey35"
23
+
24
+
25
+ def layout_size(width: int, height: int) -> tuple[int, int]:
26
+ """Minimum console (columns, rows) for a board of ``width`` x ``height``."""
27
+ columns = max(2 * int(width) + 2, 48)
28
+ rows = int(height) + 4
29
+ return columns, rows
30
+
31
+
32
+ def _cells(board: dict):
33
+ width = int(board.get("width") or 0)
34
+ height = int(board.get("height") or 0)
35
+ body = [tuple(cell) for cell in (board.get("body") or [])]
36
+ head = tuple(board["head"]) if board.get("head") else (body[0] if body else None)
37
+ food = tuple(board["food"]) if board.get("food") else None
38
+ occupied = set(body)
39
+ grid = []
40
+ for y in range(height):
41
+ row = []
42
+ for x in range(width):
43
+ point = (x, y)
44
+ if head is not None and point == head:
45
+ row.append((HEAD_GLYPH, HEAD_STYLE))
46
+ elif point in occupied:
47
+ row.append((BODY_GLYPH, BODY_STYLE))
48
+ elif food is not None and point == food:
49
+ row.append((FOOD_GLYPH, FOOD_STYLE))
50
+ else:
51
+ row.append((EMPTY_GLYPH, EMPTY_STYLE))
52
+ grid.append(row)
53
+ return grid
54
+
55
+
56
+ class Frame:
57
+ def __init__(self, board: dict, decision: dict, stats: dict):
58
+ self.board = board or {}
59
+ self.decision = decision or {}
60
+ self.stats = stats or {}
61
+
62
+ def rich_text(self) -> Text:
63
+ text = Text(style=f"on {BG}")
64
+ for y, row in enumerate(_cells(self.board)):
65
+ if y:
66
+ text.append("\n")
67
+ for x, (glyph, style) in enumerate(row):
68
+ text.append(glyph, style=style)
69
+ if x < len(row) - 1:
70
+ text.append(" ")
71
+ board = self.board
72
+ stats = self.stats
73
+ decision = self.decision
74
+ text.append("\n\n")
75
+ text.append(
76
+ "score {score} best {best} len {length} steps {steps} "
77
+ "{sps:.1f}/s guarded {guarded} int {interventions}".format(
78
+ score=board.get("score", 0),
79
+ best=stats.get("best", 0),
80
+ length=board.get("length", 0),
81
+ steps=board.get("steps", 0),
82
+ sps=float(stats.get("steps_per_second", 0.0) or 0.0),
83
+ guarded=stats.get("guarded", "-"),
84
+ interventions=stats.get("interventions", 0),
85
+ ),
86
+ style="dim",
87
+ )
88
+ text.append("\n")
89
+ guarded = " [guarded]" if decision.get("intervened") else ""
90
+ text.append(
91
+ "next {executed} model {chosen}{guarded} {ms:.2f} ms".format(
92
+ executed=decision.get("executed", "-"),
93
+ chosen=decision.get("chosen", "-"),
94
+ guarded=guarded,
95
+ ms=float(decision.get("inference_ms", 0.0) or 0.0),
96
+ ),
97
+ style="cyan",
98
+ )
99
+ return text
100
+
101
+
102
+ def compose(board: dict, decision: dict, stats: dict) -> Frame:
103
+ return Frame(board, decision, stats)
laya_onnx/tokenizer.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load the checkpoint tokenizer without Transformers or torch."""
2
+
3
+ import json
4
+ from pathlib import Path
5
+
6
+ from tokenizers import Tokenizer as Backend
7
+
8
+
9
+ class Tokenizer:
10
+ def __init__(self, path):
11
+ path = Path(path)
12
+ self.backend = Backend.from_file(str(path / "tokenizer.json"))
13
+ self.backend.no_padding()
14
+ self.backend.no_truncation()
15
+ config = json.loads((path / "tokenizer_config.json").read_text())
16
+ for name in ("cls_token", "sep_token", "pad_token", "mask_token"):
17
+ value = config.get(name)
18
+ if isinstance(value, dict):
19
+ value = value.get("content")
20
+ token_id = self.backend.token_to_id(value) if isinstance(value, str) else None
21
+ if token_id is None:
22
+ raise ValueError(f"Tokenizer is missing a valid {name}")
23
+ setattr(self, name, value)
24
+ setattr(self, name + "_id", token_id)
25
+
26
+ def __call__(self, text, add_special_tokens=False):
27
+ return {"input_ids": self.backend.encode(text, add_special_tokens=add_special_tokens).ids}
laya_onnx/torch_model.py ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tiny ModernBERT-style DecisionModel used for export and parity tests."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ import torch
8
+ from torch import nn
9
+
10
+
11
+ class EncoderConfig:
12
+ def __init__(self, data: dict):
13
+ self.hidden_size = int(data["hidden_size"])
14
+ self.intermediate_size = int(data.get("intermediate_size", self.hidden_size * 4))
15
+ self.vocab_size = int(data["vocab_size"])
16
+ self.num_hidden_layers = int(data["num_hidden_layers"])
17
+ self.num_attention_heads = int(data["num_attention_heads"])
18
+ self.local_attention = int(data.get("local_attention", 128))
19
+ self.global_attn_every_n_layers = int(data.get("global_attn_every_n_layers", 3))
20
+ self.max_position_embeddings = int(data.get("max_position_embeddings", 512))
21
+ self.layer_norm_eps = float(data.get("layer_norm_eps", 1e-5))
22
+
23
+ @classmethod
24
+ def from_dict(cls, data):
25
+ return cls(data)
26
+
27
+
28
+ class EncoderLayer(nn.Module):
29
+ def __init__(self, cfg: EncoderConfig, global_attn: bool):
30
+ super().__init__()
31
+ self.global_attn = global_attn
32
+ self.local_attention = cfg.local_attention
33
+ self.n_heads = cfg.num_attention_heads
34
+ self.head_dim = cfg.hidden_size // cfg.num_attention_heads
35
+ self.in_proj = nn.Linear(cfg.hidden_size, 3 * cfg.hidden_size, bias=True)
36
+ self.out_proj = nn.Linear(cfg.hidden_size, cfg.hidden_size, bias=True)
37
+ self.norm1 = nn.LayerNorm(cfg.hidden_size, eps=cfg.layer_norm_eps)
38
+ self.norm2 = nn.LayerNorm(cfg.hidden_size, eps=cfg.layer_norm_eps)
39
+ self.ff = nn.Sequential(
40
+ nn.Linear(cfg.hidden_size, cfg.intermediate_size),
41
+ nn.GELU(),
42
+ nn.Linear(cfg.intermediate_size, cfg.hidden_size),
43
+ )
44
+
45
+ def forward(self, hidden, attention_mask):
46
+ b, t, c = hidden.shape
47
+ x = self.norm1(hidden)
48
+ qkv = self.in_proj(x).view(b, t, 3, self.n_heads, self.head_dim)
49
+ q, k, v = qkv.unbind(2)
50
+ q = q.transpose(1, 2)
51
+ k = k.transpose(1, 2)
52
+ v = v.transpose(1, 2)
53
+ scale = 1.0 / math.sqrt(self.head_dim)
54
+ scores = torch.matmul(q, k.transpose(-2, -1)) * scale
55
+ pad = attention_mask.eq(0).unsqueeze(1).unsqueeze(2)
56
+ scores = scores.masked_fill(pad, torch.finfo(scores.dtype).min)
57
+ if not self.global_attn:
58
+ idx = torch.arange(t, device=hidden.device)
59
+ dist = (idx[None, :] - idx[:, None]).abs()
60
+ local = dist.gt(self.local_attention // 2).unsqueeze(0).unsqueeze(0)
61
+ scores = scores.masked_fill(local, torch.finfo(scores.dtype).min)
62
+ attn = torch.softmax(scores, dim=-1)
63
+ ctx = torch.matmul(attn, v).transpose(1, 2).contiguous().view(b, t, c)
64
+ hidden = hidden + self.out_proj(ctx)
65
+ hidden = hidden + self.ff(self.norm2(hidden))
66
+ return hidden
67
+
68
+
69
+ class Encoder(nn.Module):
70
+ def __init__(self, cfg: EncoderConfig, max_len: int):
71
+ super().__init__()
72
+ self.embed = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
73
+ self.pos = nn.Embedding(max(max_len, cfg.max_position_embeddings), cfg.hidden_size)
74
+ self.norm = nn.LayerNorm(cfg.hidden_size, eps=cfg.layer_norm_eps)
75
+ layers = []
76
+ for i in range(cfg.num_hidden_layers):
77
+ global_attn = ((i + 1) % cfg.global_attn_every_n_layers) == 0 or i == 0
78
+ layers.append(EncoderLayer(cfg, global_attn))
79
+ self.layers = nn.ModuleList(layers)
80
+
81
+ def forward(self, input_ids, attention_mask):
82
+ positions = torch.arange(input_ids.shape[1], device=input_ids.device)
83
+ hidden = self.embed(input_ids) + self.pos(positions)[None, :, :]
84
+ hidden = self.norm(hidden)
85
+ for layer in self.layers:
86
+ hidden = layer(hidden, attention_mask)
87
+ return hidden
88
+
89
+
90
+ class DecisionModel(nn.Module):
91
+ def __init__(self, encoder_cfg: dict, agent_cfg: dict, max_len: int = 512):
92
+ super().__init__()
93
+ if isinstance(encoder_cfg, EncoderConfig):
94
+ enc = encoder_cfg
95
+ else:
96
+ enc = EncoderConfig(encoder_cfg)
97
+ self.max_len = int(agent_cfg.get("max_len", max_len))
98
+ self.encoder = Encoder(enc, self.max_len)
99
+ hidden = enc.hidden_size
100
+ layers = int(agent_cfg.get("head_layers", 2))
101
+ blocks = []
102
+ for _ in range(max(1, layers)):
103
+ blocks.extend([nn.Linear(hidden, hidden), nn.GELU()])
104
+ self.option_mlp = nn.Sequential(*blocks)
105
+ self.option_out = nn.Linear(hidden, 1)
106
+ self.act_mlp = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Linear(hidden, 2))
107
+ self.qtype_embed = nn.Embedding(4, hidden)
108
+
109
+ def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype):
110
+ hidden = self.encoder(input_ids.long(), attention_mask)
111
+ gather_index = marker_pos.long().clamp(min=0, max=hidden.shape[1] - 1)
112
+ gather_index = gather_index.unsqueeze(-1).expand(-1, -1, hidden.shape[-1])
113
+ option_h = hidden.gather(1, gather_index)
114
+ option_h = option_h + self.qtype_embed(qtype.long().clamp(0, 3)).unsqueeze(1)
115
+ logits = self.option_out(self.option_mlp(option_h)).squeeze(-1)
116
+ invalid = marker_mask.eq(0) if marker_mask.dtype != torch.bool else ~marker_mask.bool()
117
+ logits = logits.masked_fill(invalid, -1e4)
118
+ pooled = (option_h * (~invalid).unsqueeze(-1).to(option_h.dtype)).sum(1)
119
+ denom = (~invalid).sum(1).clamp(min=1).unsqueeze(-1).to(option_h.dtype)
120
+ act = self.act_mlp(pooled / denom)
121
+ return logits, act
laya_onnx/ultrafast/__init__.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ """NL-goal + DOM loop using local Laya ONNX (jev-ultrafast shape)."""
2
+
3
+ from .agent import UltrafastAgent
4
+ from .ops import OPS
5
+ from .snapshot import Element, PageSnapshot
6
+
7
+ __all__ = ["UltrafastAgent", "OPS", "Element", "PageSnapshot"]
laya_onnx/ultrafast/agent.py ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Goal-driven DOM loop. Decisions from local Laya ONNX."""
2
+ from __future__ import annotations
3
+ import os
4
+ from typing import Iterator
5
+ from .browser import open_browser
6
+ from .ops import MAX_STEPS
7
+ from .policy import decide
8
+
9
+ def _heuristic_text(goal: str, el_name: str) -> str:
10
+ low = (el_name + " " + goal).lower()
11
+ if ("zurich" in goal.lower() or "zürich" in goal.lower()) and ("from" in low or "origin" in low or "depart" in low):
12
+ return "Zurich"
13
+ if "london" in goal.lower() and ("to" in low or "dest" in low):
14
+ return "London"
15
+ return goal.replace(" to ", " | ").split("|")[-1].strip(" ,.")[:40]
16
+
17
+ def _type_text(goal: str, el_name: str) -> str:
18
+ key = os.environ.get("TEXT_MODEL_API_KEY")
19
+ base = os.environ.get("TEXT_MODEL_BASE_URL")
20
+ if not key and not base:
21
+ return _heuristic_text(goal, el_name)
22
+ try:
23
+ import json, urllib.request
24
+ url = (base or "https://openrouter.ai/api/v1").rstrip("/") + "/chat/completions"
25
+ body = json.dumps({"model": os.environ.get("TEXT_MODEL", "openai/gpt-4.1-mini"), "messages": [{"role": "system", "content": "Return only the exact string to type."}, {"role": "user", "content": f"Goal: {goal}\nField: {el_name}"}], "max_tokens": 32}).encode()
26
+ req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/json", "Authorization": f"Bearer {key or ''}"})
27
+ with urllib.request.urlopen(req, timeout=20) as resp:
28
+ return json.loads(resp.read().decode())["choices"][0]["message"]["content"].strip().strip('"')
29
+ except Exception:
30
+ return _heuristic_text(goal, el_name)
31
+
32
+ class UltrafastAgent:
33
+ def __init__(self, url: str, goal: str, *, model: str = "receptron/laya-onnx", providers: str = "cpu", fixture: str | None = None, headless: bool = True, laya=None, deterministic: bool = False):
34
+ self.url, self.goal, self.fixture, self.headless = url, goal, fixture, headless
35
+ self.deterministic = bool(deterministic)
36
+ if laya is None:
37
+ from laya_onnx import load
38
+ self.laya = load(model, providers=providers, deterministic=self.deterministic)
39
+ else:
40
+ self.laya = laya
41
+ self.browser = None
42
+ def __enter__(self):
43
+ self.browser = open_browser(self.url, fixture=self.fixture, headless=self.headless)
44
+ return self
45
+ def __exit__(self, *exc):
46
+ if self.browser:
47
+ self.browser.close()
48
+ return False
49
+ def run(self, max_steps: int = MAX_STEPS) -> Iterator[dict]:
50
+ if self.browser is None:
51
+ self.browser = open_browser(self.url, fixture=self.fixture, headless=self.headless)
52
+ for step in range(max_steps):
53
+ snap = self.browser.snapshot()
54
+ decision = decide(self.laya, snap, self.goal, argmax=self.deterministic)
55
+ op, target, text, status = decision["op"], decision["target"], None, "ok"
56
+ try:
57
+ if op == "CLICK" and target is not None:
58
+ self.browser.click(target)
59
+ elif op == "TYPE_TEXT" and target is not None:
60
+ el = snap.by_index(target)
61
+ name = el.name if el else ""
62
+ text = (
63
+ _heuristic_text(self.goal, name)
64
+ if self.deterministic
65
+ else _type_text(self.goal, name)
66
+ )
67
+ self.browser.type_text(target, text)
68
+ elif op == "SELECT" and target is not None:
69
+ self.browser.select(target)
70
+ elif op in ("SCROLL_DOWN", "SCROLL_UP"):
71
+ self.browser.scroll(op)
72
+ elif op == "WAIT":
73
+ self.browser.wait(400)
74
+ elif op in ("DONE", "BLOCKED"):
75
+ yield {"step": step, "op": op, "target": target, "status": op.lower(), "url": snap.url, **decision}
76
+ return
77
+ except Exception as exc:
78
+ status = f"error:{exc}"
79
+ yield {"step": step, "op": op, "target": target, "text": text, "status": status, "url": snap.url, "ready": decision["ready"]}
80
+ yield {"step": max_steps, "op": "BLOCKED", "status": "max_steps", "url": self.url}
laya_onnx/ultrafast/browser.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Playwright driver with a fixture/mock fallback."""
2
+ from __future__ import annotations
3
+ import json
4
+ from pathlib import Path
5
+ from .snapshot import PageSnapshot, SNAPSHOT_JS
6
+
7
+ class MockBrowser:
8
+ def __init__(self, fixture):
9
+ self.snap = fixture if isinstance(fixture, PageSnapshot) else PageSnapshot.from_dict(fixture)
10
+ self.actions = []
11
+ def snapshot(self):
12
+ return self.snap
13
+ def click(self, index: int):
14
+ self.actions.append({"op": "CLICK", "target": index})
15
+ def type_text(self, index: int, text: str):
16
+ el = self.snap.by_index(index)
17
+ if el:
18
+ el.value = text
19
+ self.actions.append({"op": "TYPE_TEXT", "target": index, "text": text})
20
+ def select(self, index: int):
21
+ self.actions.append({"op": "SELECT", "target": index})
22
+ def scroll(self, direction: str):
23
+ self.actions.append({"op": "SCROLL", "direction": direction})
24
+ def wait(self, ms: int = 400):
25
+ self.actions.append({"op": "WAIT", "ms": ms})
26
+ def close(self):
27
+ pass
28
+
29
+ class PlaywrightBrowser:
30
+ def __init__(self, url: str, headless: bool = True):
31
+ from playwright.sync_api import sync_playwright
32
+ self._pw = sync_playwright().start()
33
+ self._browser = self._pw.chromium.launch(headless=headless)
34
+ self._page = self._browser.new_page()
35
+ self._page.goto(url, wait_until="domcontentloaded", timeout=30000)
36
+ def snapshot(self):
37
+ return PageSnapshot.from_dict(self._page.evaluate(SNAPSHOT_JS))
38
+ def click(self, index: int):
39
+ self._page.evaluate("(i) => { const n = document.querySelectorAll('a,button,input,select,textarea,[role=button],[role=link]')[i]; n && n.click(); }", int(index))
40
+ def type_text(self, index: int, text: str):
41
+ self._page.evaluate("([i, t]) => { const n = document.querySelectorAll('input,textarea,select,[role=textbox]')[i]; if (n) { n.focus(); n.value = t; n.dispatchEvent(new Event('input', {bubbles:true})); } }", [int(index), text])
42
+ def select(self, index: int):
43
+ self.click(index)
44
+ def scroll(self, direction: str):
45
+ self._page.mouse.wheel(0, 600 if direction == "SCROLL_DOWN" else -600)
46
+ def wait(self, ms: int = 400):
47
+ self._page.wait_for_timeout(ms)
48
+ def close(self):
49
+ self._browser.close(); self._pw.stop()
50
+
51
+ def open_browser(url: str, *, fixture=None, headless: bool = True):
52
+ if fixture:
53
+ return MockBrowser(json.loads(Path(fixture).read_text()))
54
+ try:
55
+ return PlaywrightBrowser(url, headless=headless)
56
+ except Exception as exc:
57
+ raise RuntimeError("Playwright missing. pip install -e '.[ultrafast]' && playwright install chromium, or pass --fixture") from exc
laya_onnx/ultrafast/cli.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+ import argparse, json
3
+ from .agent import UltrafastAgent
4
+
5
+ def main(argv=None) -> int:
6
+ p = argparse.ArgumentParser(prog="laya-onnx-ultrafast")
7
+ p.add_argument("--url", default="https://example.com")
8
+ p.add_argument("--goal", required=True)
9
+ p.add_argument("--model", default="receptron/laya-onnx")
10
+ p.add_argument("--providers", default="cpu")
11
+ p.add_argument("--fixture")
12
+ p.add_argument("--dry-run", action="store_true")
13
+ p.add_argument("--max-steps", type=int, default=20)
14
+ p.add_argument("--headed", action="store_true")
15
+ p.add_argument("--deterministic", action="store_true")
16
+ args = p.parse_args(argv)
17
+ if args.dry_run and not args.fixture:
18
+ raise SystemExit("--dry-run needs --fixture")
19
+ with UltrafastAgent(args.url, args.goal, model=args.model, providers=args.providers, fixture=args.fixture, headless=not args.headed, deterministic=args.deterministic) as agent:
20
+ for step in agent.run(max_steps=args.max_steps):
21
+ print(json.dumps(step, ensure_ascii=False))
22
+ if step.get("op") in ("DONE", "BLOCKED") or step.get("status") == "max_steps":
23
+ break
24
+ return 0
25
+
26
+ if __name__ == "__main__":
27
+ raise SystemExit(main())
laya_onnx/ultrafast/ops.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ OPS = ("CLICK", "TYPE_TEXT", "SELECT", "SCROLL_DOWN", "SCROLL_UP", "WAIT", "DONE", "BLOCKED")
2
+ MAX_ELEMENTS = 250
3
+ MAX_STEPS = 60
laya_onnx/ultrafast/policy.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+ from .ops import OPS
3
+ from .snapshot import PageSnapshot
4
+
5
+ def questions(snap: PageSnapshot, goal: str) -> dict:
6
+ labels = [el.label() for el in snap.elements[:32]] or ["(none)"]
7
+ return {
8
+ "op": {"type": "choice", "instructions": f"Goal: {goal}\nURL: {snap.url}\nPick the next browser operation.", "criteria": list(OPS)},
9
+ "ready": {"type": "noul", "instructions": f"Is the goal already satisfied? Goal: {goal}. Text: {snap.text[:400]}"},
10
+ "target": {"type": "choice", "instructions": f"Goal: {goal}. Which control?", "criteria": labels[:16]},
11
+ }
12
+
13
+ def decide(agent, snap: PageSnapshot, goal: str, *, argmax: bool = False) -> dict:
14
+ state = f"GOAL\n{goal}\n\nPAGE {snap.url}\n{snap.title}\n{snap.table(40)}"
15
+ ask = agent.predict_argmax if argmax and hasattr(agent, "predict_argmax") else agent.predict
16
+ out = ask(state, questions(snap, goal))
17
+ answers = out["answers"]
18
+ ready = float(answers.get("ready", {}).get("noul", 0))
19
+ op = answers.get("op", {}).get("choice", "WAIT")
20
+ if ready >= 0.72:
21
+ op = "DONE"
22
+ if op not in OPS:
23
+ op = "WAIT"
24
+ target_label = answers.get("target", {}).get("choice")
25
+ target = None
26
+ if target_label:
27
+ for el in snap.elements:
28
+ if el.label() == target_label:
29
+ target = el.index
30
+ break
31
+ if target is None and snap.elements:
32
+ target = snap.elements[0].index
33
+ return {"op": op, "target": target, "ready": ready, "raw": answers, "usage": out.get("usage", {})}
laya_onnx/ultrafast/snapshot.py ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+ from dataclasses import asdict, dataclass, field
3
+
4
+ @dataclass
5
+ class Element:
6
+ index: int
7
+ role: str
8
+ name: str
9
+ value: str = ""
10
+ tag: str = ""
11
+ editable: bool = False
12
+ def label(self) -> str:
13
+ bits = [f"#{self.index}", self.role or self.tag or "el", self.name[:80]]
14
+ if self.value:
15
+ bits.append(f"= {self.value[:40]}")
16
+ return " ".join(x for x in bits if x)
17
+
18
+ @dataclass
19
+ class PageSnapshot:
20
+ url: str
21
+ title: str = ""
22
+ elements: list = field(default_factory=list)
23
+ text: str = ""
24
+ def table(self, cap: int = 250) -> str:
25
+ return "\n".join(el.label() for el in self.elements[:cap]) or "(no interactive elements)"
26
+ def by_index(self, idx: int):
27
+ for el in self.elements:
28
+ if el.index == idx:
29
+ return el
30
+ return None
31
+ def to_dict(self) -> dict:
32
+ return {"url": self.url, "title": self.title, "text": self.text, "elements": [asdict(e) for e in self.elements]}
33
+ @classmethod
34
+ def from_dict(cls, raw: dict):
35
+ els = [e if isinstance(e, Element) else Element(**e) for e in raw.get("elements", [])]
36
+ return cls(url=raw.get("url", ""), title=raw.get("title", ""), text=raw.get("text", ""), elements=els)
37
+
38
+ SNAPSHOT_JS = """() => {
39
+ const seen = new WeakSet(); const out = [];
40
+ const nodes = document.querySelectorAll('a,button,input,select,textarea,[role=button],[role=link],[role=textbox],[contenteditable=true]');
41
+ for (const n of nodes) {
42
+ if (seen.has(n) || out.length >= 250) break; seen.add(n);
43
+ const r = n.getBoundingClientRect(); if (r.width < 2 || r.height < 2) continue;
44
+ const role = n.getAttribute('role') || n.tagName.toLowerCase();
45
+ const name = (n.getAttribute('aria-label') || n.getAttribute('placeholder') || n.innerText || n.value || '').trim();
46
+ out.push({index: out.length, role, name: name.slice(0,120), value: String(n.value||'').slice(0,80), tag: n.tagName.toLowerCase(), editable: !!(n.isContentEditable || /input|textarea|select/.test(n.tagName.toLowerCase()))});
47
+ }
48
+ return {url: location.href, title: document.title, text: (document.body.innerText||'').slice(0,1500), elements: out};
49
+ }"""
laya_onnx/verify.py ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Verify a downloaded ONNX bundle against a SHA-256 manifest."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import json
7
+ from pathlib import Path
8
+
9
+ MANIFEST_DIR = Path(__file__).with_name("checksums")
10
+ DEFAULT_MANIFEST = MANIFEST_DIR / "receptron-laya-onnx.json"
11
+
12
+ CHUNK = 1 << 20
13
+
14
+
15
+ def sha256_file(path, chunk=CHUNK) -> str:
16
+ digest = hashlib.sha256()
17
+ with Path(path).open("rb") as handle:
18
+ for block in iter(lambda: handle.read(chunk), b""):
19
+ digest.update(block)
20
+ return digest.hexdigest()
21
+
22
+
23
+ def load_manifest(path=None) -> dict:
24
+ manifest = Path(path) if path else DEFAULT_MANIFEST
25
+ if not manifest.is_file():
26
+ raise FileNotFoundError(f"No checksum manifest at {manifest}")
27
+ return json.loads(manifest.read_text(encoding="utf-8"))
28
+
29
+
30
+ def verify_bundle(path, manifest=None):
31
+ """Check every file in the manifest against ``path``.
32
+
33
+ Returns ``(ok, rows)`` where each row is ``{"file", "status", ...}`` and
34
+ status is one of ``ok`` / ``missing`` / ``size-mismatch`` / ``sha256-mismatch``.
35
+ """
36
+ data = manifest if isinstance(manifest, dict) else load_manifest(manifest)
37
+ root = Path(path).expanduser()
38
+ rows = []
39
+ ok = True
40
+ for name, expected in data.get("files", {}).items():
41
+ target = root / name
42
+ if not target.is_file():
43
+ rows.append({"file": name, "status": "missing"})
44
+ ok = False
45
+ continue
46
+ size = target.stat().st_size
47
+ if expected.get("size") is not None and size != expected["size"]:
48
+ rows.append(
49
+ {
50
+ "file": name,
51
+ "status": "size-mismatch",
52
+ "expected_size": expected["size"],
53
+ "size": size,
54
+ }
55
+ )
56
+ ok = False
57
+ continue
58
+ digest = sha256_file(target)
59
+ if expected.get("sha256") and digest != expected["sha256"]:
60
+ rows.append(
61
+ {
62
+ "file": name,
63
+ "status": "sha256-mismatch",
64
+ "expected": expected["sha256"],
65
+ "actual": digest,
66
+ }
67
+ )
68
+ ok = False
69
+ else:
70
+ rows.append({"file": name, "status": "ok", "sha256": digest})
71
+ return ok, rows
pyproject.toml ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [build-system]
2
+ requires = ["hatchling"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "laya-onnx"
7
+ version = "0.1.0"
8
+ description = "ONNX Runtime inference for Laya typed decision models (CPU, OpenVINO, CUDA)"
9
+ readme = "README.md"
10
+ requires-python = ">=3.11"
11
+ license = "Apache-2.0"
12
+ dependencies = [
13
+ "onnxruntime>=1.18",
14
+ "onnx>=1.16",
15
+ "numpy>=1.26",
16
+ "huggingface-hub>=0.34,<2",
17
+ "tokenizers>=0.21,<1",
18
+ ]
19
+
20
+ [project.optional-dependencies]
21
+ openvino = ["onnxruntime-openvino>=1.18"]
22
+ gpu = ["onnxruntime-gpu>=1.18"]
23
+ export = ["torch>=2.4", "transformers>=4.44", "safetensors>=0.4"]
24
+ dev = ["pytest>=8", "ruff>=0.12", "build>=1"]
25
+ demo = ["rich>=13,<16", "Pillow>=10"]
26
+ ultrafast = ["playwright>=1.47"]
27
+
28
+ [project.scripts]
29
+ laya-onnx = "laya_onnx.cli:main"
30
+ laya-onnx-snake = "laya_onnx.snake.cli:main"
31
+ laya-onnx-ultrafast = "laya_onnx.ultrafast.cli:main"
32
+
33
+ [project.urls]
34
+ Repository = "https://github.com/Geoking2104/laya-onnx"
35
+
36
+ [tool.hatch.build.targets.wheel]
37
+ packages = ["laya_onnx"]