Mirror laya-onnx source tree + model card
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitignore +16 -0
- AGENTS.md +10 -0
- GROK.md +10 -0
- LICENSE +15 -0
- NOTICE +9 -0
- ONBOARD.md +35 -0
- README.md +134 -0
- benchmarks/eval/decisions.jsonl +8 -0
- benchmarks/evaluate.py +56 -0
- benchmarks/pc_benchmark.py +37 -0
- benchmarks/results.json +38 -0
- benchmarks/results.md +77 -0
- docs/DETERMINISTIC.md +6 -0
- docs/ULTRAFAST.md +12 -0
- examples/questions.json +21 -0
- examples/quickstart.py +14 -0
- examples/snake.html +79 -0
- examples/state.json +5 -0
- examples/ultrafast_goal.py +14 -0
- examples/ultrafast_page.json +13 -0
- laya_onnx/__init__.py +10 -0
- laya_onnx/__main__.py +3 -0
- laya_onnx/agent.py +280 -0
- laya_onnx/checksums/receptron-laya-onnx.json +28 -0
- laya_onnx/cli.py +87 -0
- laya_onnx/common.py +122 -0
- laya_onnx/convert.py +137 -0
- laya_onnx/eval.py +186 -0
- laya_onnx/hub.py +83 -0
- laya_onnx/inputs.py +63 -0
- laya_onnx/optimize.py +139 -0
- laya_onnx/snake/__init__.py +1 -0
- laya_onnx/snake/__main__.py +3 -0
- laya_onnx/snake/benchmark.py +88 -0
- laya_onnx/snake/cli.py +263 -0
- laya_onnx/snake/game.py +122 -0
- laya_onnx/snake/policy.py +142 -0
- laya_onnx/snake/replay.py +140 -0
- laya_onnx/snake/ui.py +103 -0
- laya_onnx/tokenizer.py +27 -0
- laya_onnx/torch_model.py +121 -0
- laya_onnx/ultrafast/__init__.py +7 -0
- laya_onnx/ultrafast/agent.py +80 -0
- laya_onnx/ultrafast/browser.py +57 -0
- laya_onnx/ultrafast/cli.py +27 -0
- laya_onnx/ultrafast/ops.py +3 -0
- laya_onnx/ultrafast/policy.py +33 -0
- laya_onnx/ultrafast/snapshot.py +49 -0
- laya_onnx/verify.py +71 -0
- pyproject.toml +37 -0
.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"]
|