diff --git a/docs/components.md b/docs/components.md index 82b629f..822d221 100644 --- a/docs/components.md +++ b/docs/components.md @@ -16,7 +16,7 @@ Registered via `register_controller`; created with `ControllerFactory`. All implement `compute(current, target)` and `reset()`. | Registered name | Class | File | Config | -|-----------------|-------|------|--------| +| ----------------- | ------- | ------ | -------- | | `LQR` | `LQR` | `controllers/lqr.py` | `configs/controllers/lqr_base.toml` | | `PID` | `PIDController` | `controllers/pid.py` | `configs/controllers/pid_arm.toml` | | `MPC_LTI` | `MPC_LTI_Base` | `controllers/mpc_lti.py` | `configs/controllers/mpc_lti_base.toml` | @@ -29,6 +29,11 @@ implement `compute(current, target)` and `reset()`. > `MPC_LTI` and `MPC_DeltaU` are both built on `MPC_LTI` in `mpc_lti.py`: > `MPC_DeltaU` adds Δu (control-rate) regularization. The two names are > distinct registrations, not aliases. +> +> `onnx_rl` is the only controller that may omit `[estimator]` in a scenario, +> and the only one with a compiled deployment path: an ONNX policy imports to a +> standalone graph, which the adapter runs eagerly from `model_path` or through +> a `make compile` kernel when `artifact_dir` is set (see the config header). ## Plants @@ -38,7 +43,7 @@ linearization set `input_dim` and expose `dynamics(x, u)` (see `docs/how-it-works.md` and `utils/linearization.py`). | Registered name | Class | File | Config | -|-----------------|-------|------|--------| +| ----------------- | ------- | ------ | -------- | | `ArmRobot` | `ArmRobot` | `plants/armrobot.py` | `configs/plants/armrobot.toml` | | `HolonomicMobileRobot` | `HolonomicMobileRobot` | `plants/holonomicmobilerobot.py` | `configs/plants/holonomic_base.toml` | | `InvertedPendulum` | `InvertedPendulum` | `plants/inverted_pendulum.py` | `configs/plants/inverted_pendulum.toml` | @@ -66,7 +71,7 @@ Registered via `register_trajectory`; created with `TrajectoryFactory`. All implement `generate(...)` and `position_at(t)`. | Registered name | Class | File | Config | -|-----------------|-------|------|--------| +| ----------------- | ------- | ------ | -------- | | `cubic_segments` | `CubicPolynomial` | `trajectories/cubic_polynomial.py` | `configs/trajectories/arm_extension.toml` | | `quintic_segments` | `QuinticPolynomial` / `QuinticPolynomialConfigAdapter` | `trajectories/quintic_polynomial.py` | `configs/trajectories/arm_quintic.toml` | | `waypoints` | `WaypointSchedule` | `trajectories/quintic_polynomial.py` | `configs/trajectories/arm_lift.toml`, `base_straight.toml`, `base_triangle.toml` | @@ -87,7 +92,7 @@ via `plant.physics_engine(engine)`. MuJoCo requires the optional Not registry-based; import directly from `shinro.utils`. | Symbol | Module | Purpose | -|--------|--------|---------| +| -------- | -------- | --------- | | `ArrayBackend`, `NumpyBackend`, `TorchBackend` | `utils/array_backend.py` | Backend-agnostic array abstraction; `parse_matrix` converts TOML lists to matrices | | `BatchedDynamicsAdapter` | `utils/batched_adapter.py` | Batches N parallel trajectory rollouts for sampling-based controllers (MPPI) | | `linearize`, `linearize_plant` | `utils/linearization.py` | Numeric linearization of plant dynamics around an operating point | @@ -100,7 +105,7 @@ Tracing/composition/lowering pipeline. See `docs/codegen.md` for the full walkthrough. | Symbol | Module | Purpose | -|--------|--------|---------| +| -------- | -------- | --------- | | `Tracer`, `Graph`, `Node` | `codegen/tracing.py` | Abstract values + graph records; operator overloads emit nodes | | `TraceBackend` | `codegen/trace_backend.py` | Recording `ArrayBackend` used during tracing | | `trace_node`, `trace_node_with_state` | `codegen/trace_node.py` | Trace one component call | diff --git a/lab-notes/daily/2026-09-17.md b/lab-notes/daily/2026-09-17.md index d11c28c..a1424a1 100644 --- a/lab-notes/daily/2026-09-17.md +++ b/lab-notes/daily/2026-09-17.md @@ -42,3 +42,265 @@ Make recipe (`$$(...)`, Make syntax) as shell parse errors (SC2276/SC1036/SC1088), and its default `yamllint` caps CI lines at 80 chars. Both reproduce on `HEAD` and are pre-existing; the repo has no yamllint or shellcheck config, so they are not addressed here. + +### 2026-09-17 22:51 UTC — ONNX RL policies compiled to Zig, and where the comptime VM stops scaling + +**Why.** `onnx_rl` needed `onnxruntime` (~10–20 MB C++ runtime) at deploy time, +and its observation encoding / action post-processing lived in Python. Replace +that with the framework's own path: an ONNX graph is *already* a dataflow graph, +so translate it into a shinro `Graph`, bake the encoder + post-processing as +arithmetic on constants, and compile to a dependency-free `.so` the host +dlopens. The `onnx` schema package (compile-time only) replaces `onnxruntime`. + +**What (steps 1–5; steps 1–4 committed `81a86a5`…`ad81d05`, step 5 uncommitted).** + +1. `onnx-rl` extra: `onnxruntime` → `onnx` (and `onnx` into `dev`) — also fixes +the silent `importorskip` skip the adapter tests had. +1. `codegen/onnx_import.py`: `import_onnx_policy` — a **tracer-free** ONNX→`Graph` +translator (`Gemm`/`MatMul`/`Add`/`Relu`/`Tanh`; `Sigmoid` composed from +`exp/neg/add/div`). Every emitted node is evaluated eagerly with the ops +registry's numpy handler, so its declared shape comes from real semantics +rather than a parallel shape table. The obs encoder (index selection → baked +0/1 matmul, mean/std, clip) is folded in; unreachable nodes are ignored; +unsupported ops/attrs, multi-input models, and batched outputs fail loudly. +Implemented as one `_OnnxImporter` object rather than free functions threading +a builder through every call. +1. Action space + `epsilon` port: continuous scale/bias, discrete argmax +one-hot, stochastic `[mean; log_std]`. Non-deterministic spaces take host noise +through an `epsilon` C-ABI port (Gumbel for discrete, so `argmax(logits+g)` is +exact categorical sampling; standard normal for stochastic) — RNG stays on the +host. Single-sided action clip is rejected: a `±inf` bound cannot be emitted +(Zig has no `inf` literal). +1. `controllers/onnx_rl_adapter.py` rewritten: frozen `OnnxRLConfig` (kills the +registry's missing-Config warning) and **two interchangeable backends** — eager +(`model_path` → `interpret`) and compiled (`artifact_dir` → ctypes `shinro_step` + +- the graph manifest for port layout). Host-side seeded noise; +backend-agnostic. Actions are now f64, not f32. + +1. `scenario_gen` policy-only branch (no `[estimator]` for `onnx_rl`), `.onnx` +sha256 pinned in the manifest provenance, and `[compile].artifact_name` threaded +through `build.zig`/`oracle`/`stamp`/`scenario_build` — the kernel installs as +`lib/.so`. (Zig's `addLibrary` prefixes `lib`, so an explicit install +sub-path is needed or `lib_neural_network` becomes `liblib_neural_network`.) + +**Toy fixture.** `tests/fixtures/models/toy_mlp.onnx` (316 B; 26 params; +`action = [tanh(x0)+0.5, tanh(x1)-0.5]`), generated deterministically by +`scripts/gen_toy_onnx.py`, plus a controller config and a policy-only scenario +under `tests/fixtures/configs/`. A drift guard regenerates it (through the +script's CLI) and compares the graph, so the fixture cannot rot. +`make compile SCENARIO=tests/fixtures/configs/scenarios/toy_mlp_policy.toml` now +works out of the box. + +**Scale sweep** (`scripts/measure_onnx_policy_scale.py`, `--optimize release`, +**cold** build — the shared local zig cache is cleared per point, because zig +caches comptime work and a warm cache turns minutes into seconds; 256×4 measured +2.3 s warm vs 649 s cold): + +| params | arch (h×depth) | onnx | graph_data.zig | .so | gen s | build s | nodes | VM stack | consts | +| --- | --- | --- | --- | --- | --- | --- | --- | --- | --- | +| 9,360 | 64×2 | 36.9 KB | 0.2 MB | 234 KB | 1.6 | 10.6 | 35 | 152 KB | 73 KB | +| 86,544 | 256×2 | 338 KB | 1.9 MB | 2.0 MB | 1.2 | 175 | 35 | 1.34 MB | 676 KB | +| 218,128 | 256×4 | 853 KB | 4.7 MB | 5.1 MB | 1.4 | 649 | 55 | 3.4 MB | 1.7 MB | +| 829,456 | 512×4 | 3.2 MB | 17.8 MB | — | 2.2 | 27 | — | — | — (comptime quota) | +| 1,354,768 | 512×6 | 5.2 MB | 29.1 MB | — | 2.9 | 30 | — | — | — (comptime quota) | +| 1,880,080 | 512×8 | 7.3 MB | 40.4 MB | — | 3.6 | 34 | — | — | — (comptime quota) | + +Every compiled point passes oracle B at ≤5e-16. Node count is independent of +width (the batch lives in node *shapes*), so all the growth is unrolled loops + +the baked const blob. + +**Where it stops scaling.** + +- **Hard wall at ~830k params**: `lower.zig:123` → `error: evaluation exceeded + 1000000 backwards branches`. The VM's per-element `inline for` loops (const + copies, elementwise ops) unroll at comptime and blow + `@setEvalBranchQuota(1_000_000)`; anything 512-wide × 4+ hidden layers fails + within seconds. +- **Compile time is superlinear**: 10.6 s → 175 s → 649 s for 9k → 86k → 218k. + A realistic 1M-param actor is hours away before the quota wall even matters, + so the comptime VM is a ~10^5-parameter design, not a 10^6 one. The follow-up + is runtime loops over comptime-known sizes for the const-copy/elementwise + arms (plus a compact const encoding). +- **Source size ≈ 22 bytes/param**: constants are emitted as hex-float literals, + so 1M params is a ~22 MB `graph_data.zig` (1.88M → 40 MB). `build.zig`'s + `readFile` cap of **1 MiB** made every policy ≥ ~50k params fail with the + misleading "graph_data.zig has no has_solve_qp flag; regenerate it" panic; + raised to 256 MiB — a real bug fix, not a tweak. +- **Binary ≈ const blob + unrolled code**: f64 weights are 8 B/param (2× the f32 + `.onnx`), so 86k params is 676 KB of constants + ~1.4 MB of code. + +**Verification.** `make test` 1233 passed / 5 skipped; `make test-zig` 77 passed +/ 2 skipped; `make lint` ruff + pyrefly clean; the committed toy scenario +compiles, oracle-verifies, stamps, and verifies end-to-end; eager vs compiled +agree to 1.1e-16. + +**Files.** `src/shinro/codegen/onnx_import.py` (new); `controllers/onnx_rl_adapter.py` +(rewritten); `codegen/{scenario_gen,scenario_build,oracle,stamp,cli}.py`; +`runtime/build.zig`; `configs/scenarios/_template.toml`; tests +`tests/unit/test_onnx_import.py` (new), `tests/unit/test_onnx_rl_adapter.py`, +`tests/test_compile_scenario.py`; fixtures `tests/fixtures/models/toy_mlp.onnx`, +`tests/fixtures/configs/{controllers/onnx_toy.toml,scenarios/toy_mlp_policy.toml}`; +tooling `scripts/gen_toy_onnx.py`, `scripts/measure_onnx_policy_scale.py`. + +**Step 6 — committed Zig oracle.** `TestOnnxPolicyOracle` in +`tests/test_zig_lowering.py` imports the committed toy policy and compiles its +four baked variants (continuous, discrete-deterministic, discrete-sampling, and +stochastic-sampling) into tmp-path kernels, then asserts the ``.so`` matches +``interpret()`` on seeded random states — exact for the continuous/one-hot paths +and ≤1e-14 for the noise paths — plus the closed form +(``[tanh(x0)+0.5, tanh(x1)-0.5]``), numpy ``argmax``, and the Gumbel/stochastic +formulas. This is the first oracle whose subject is a *learned* policy rather +than a hand-written control law. `make test-zig` 81 passed / 2 skipped. + +**Step 7 — docs.** `configs/controllers/onnx_rl.toml` now documents both +backends (eager `model_path` / compiled `artifact_dir`), the policy-only compile +command, the epsilon / host-noise contract, and that `add_batch_dim` is legacy; +the stale `--controller onnx_rl` demo line is gone (no such base config exists). +`docs/components.md` records that onnx_rl is the only controller that may omit +`[estimator]` and the only one with a compiled deployment path. + +### 2026-09-17 — MPPI production-kernel sizes, and the comptime-VM element-loop fix + +**Question.** How big is a production MPPI kernel on the real plants, and why +were the builds taking minutes? + +**Measured (before the fix).** Real plant rollouts wired through +`BatchedDynamicsAdapter` via `attach_plant`, production config N=200 / K=15, +built `ReleaseFast`. Standalone MPPI and the composed KF+MPPI closed-loop +kernel, both reported: + +| plant | D_x×D_u | nodes (SA / +KF) | VM buf (SA / +KF) | `.so` (SA / +KF) | compile (SA / +KF) | +| ----- | ------- | ---------------- | ----------------- | ---------------- | ------------------ | +| InvertedPendulum (nonlinear) | 2×1 | 654 / 794 | 993 / 1119 KiB | 1.9 / 2.4 MiB | 411 / 437 s | +| CartPole (nonlinear) | 4×1 | 1134 / 1274 | 2034 / 2213 KiB | 4.5 / 5.2 MiB | 1154 / 1066 s | +| DoublePendulum (nonlinear) | 4×2 | 1224 / 1364 | 2458 / 2663 KiB | 5.3 / 6.1 MiB | 1306 / 1128 s | +| HolonomicMobileRobot (LTI) | 3×3 | 354 / 494 | 934 / 1137 KiB | — / — | — / — | + +Quadrotor (12×4) and ArmRobot are not standalone-buildable (not implemented / +needs sim-injected `engine`+`joint_groups`); the LTI `measure` tool brackets +12×4×200×15 at 2722 KiB VM-buffer, ~5.5 MiB `.so` by the ~2× `so`/`buf` ratio. +`readelf` showed the artifact is ~99.9% `.text` (IP: 1,956,286 of 1,959,696 +bytes), stripped, zero debug info — the elements *are* machine code. + +**Root cause.** `runtime/lower.zig` used `inline for` for **both** the outer +node-table loop *and* every per-element loop (38 `inline for` sites, 0 runtime +`for`). The outer unroll is cheap (~1300 nodes); the inner ones emit one +statement per element, so `shinro_step` was a single function with ≈`buf_len` +statements (100k–340k here). Zig's comptime evaluator then LLVM had to process +that giant straight-line function: compile tracks `buf_len` (~0.4–0.5 s/KiB) +and `.text` ≈ 2× `buf_bytes`. `@setEvalBranchQuota(1_000_000)` was the same +wall showing up as a symptom. + +**Fix.** Kept the outer `inline for (g.nodes)` (op tag + shapes stay comptime, +buffer offsets stay comptime constants), changed all 34 inner element loops to +runtime `for` over those comptime-known sizes. `bcast_flat`/`ew2` stay `inline` +so shapes and the op tag still constant-fold. Three design comments rewritten +(`lower.zig` header, `shinro_step` doc, `ew2` doc). One file changed, 34 +lines swapped. + +**Result.** Same four kernels, rebuilt cold after clearing +`src/shinro/runtime/.zig-cache`: + +| plant / kernel | before | after | shrink | compile before → after | speedup | +| -------------- | ------ | ----- | ------ | ---------------------- | ------- | +| IP standalone | 1914 KiB | 246 KiB | 7.8× | 411 → 8.5 s | 49× | +| IP +KF | 2402 KiB | 459 KiB | 5.2× | 437 → 12.3 s | 35× | +| CartPole standalone | 4507 KiB | 393 KiB | 11.5× | 1154 → 10.9 s | 106× | +| CartPole +KF | 5200 KiB | 618 KiB | 8.4× | 1066 → 18.1 s | 59× | +| DoublePendulum standalone | 5333 KiB | 496 KiB | 10.8× | 1306 → 15.2 s | 86× | +| DoublePendulum +KF | 6076 KiB | 715 KiB | 8.5× | 1128 → 24.9 s | 45× | +| Holonomic standalone | — | 270 KiB | — | — → 2.3 s | — | +| Holonomic +KF | — | 481 KiB | — | — → 9.2 s | — | + +The whole 8-kernel matrix now builds in ~90 s total; before, one CartPole +kernel took ~19 min. **Tick latency did not regress — it improved 1.4–2.1×** +(`scripts/bench_tick.py`, min ns/tick): IP 171→99 µs, IP+KF 176→85 µs, CartPole +629→388 µs, CartPole+KF 708→400 µs, DoublePendulum 500→275 µs. The unrolled +function was blowing I-cache/register pressure; the runtime loops fit and +vectorize. So the change is a size, compile-time, *and* runtime win — the +comptime element unrolling was paying for nothing. + +**Verification.** `make test-zig` 81 passed / 2 skipped; `make test` 1241 +passed / 5 skipped (identical to the pre-change baseline); `make lint` ruff + +pyrefly clean. Shipped `src/shinro/runtime/graph_data.zig` untouched +(`git status` shows only `lower.zig`). The three-way MPPI/SMC/ONNX oracles and +the closed-loop oracles (KF+PID, Luenberger+LQR, KF+MPC-DeltaU) all still pass, +so every op path in the VM is bit-consistent. + +**Open follow-up (not done).** The size/compile metric in `shinro.codegen.measure` +still only measures a hand-written LTI MPPI graph; the real-plant numbers above +came from a throwaway `build/mppi-investigation/` probe. Folding +`attach_plant`-driven plants into `make measure-kernels` would make the +production-size number reproducible (Quadrotor pending its plant port). + +### 2026-09-17 — ONNX/MLP scale sweep re-run after the element-loop fix + +The `scripts/measure_onnx_policy_scale.py` sweep documented the wall that +motivated the fix: ≥830k params FAILED with `evaluation exceeded 1000000 +backwards branches`, and 9.4k / 86.5k / 218k took 10.6 / 175 / 649 s to build +234 KB / 2.0 MB / 5.1 MB. Re-run post-fix (cold cache per point, oracle-verified): + +| params | arch | prod `.so` before → after | build before → after | +| ------ | ---- | ------------------------- | -------------------- | +| 9,360 | 64×2 | 234 KB → 85 KB | 10.6 → 4.3 s | +| 86,544 | 256×2 | 2.0 MB → 687 KB | 175 → 4.7 s | +| 218,128 | 256×4 | 5.1 MB → 1.7 MB | 649 → 5.7 s | +| 829,456 | 512×4 | **FAILED → 6.3 MB** | 27 s (fail) → 11.1 s | +| 1,354,768 | 512×6 | **FAILED → 10.4 MB** | 30 s (fail) → 16.2 s | +| 1,880,080 | 512×8 | **FAILED → 14.4 MB** | 34 s (fail) → 22.0 s | + +All six now compile and pass `oracle B` (max abs err ≤1.3e-15). The comptime +branch-quota wall at ~830k params is gone; the binary is ≈3× smaller and the +build 2.5–114× faster (superlinear → roughly linear in buffer size). + +**New limiting factor — the runtime stack buffer.** `shinro_step` declares +`buf: [buf_len]f64` on the stack: 20.8 MiB at 512×6 and 28.8 MiB at 512×8, +larger than the default 16 MB process stack. The sweep's oracle dlopens the +`.so` and calls it in-process, so those two points segfault (exit 139, no +error message — the harness's error summarizer then printed a stray warning +line, which is why the JSON looked like a build failure). Re-running the whole +sweep under `ulimit -s unlimited` compiles, oracle-verifies, stamps and verifies +all six. **Deployment implication:** the host must provision a stack ≥ +`buf_bytes`, or `buf` must move off the stack (static / thread-local / heap) — +the follow-up for 10^6-param policy deployment. The compiler is no longer the +limit. + +**Tests.** `pytest tests/unit/test_onnx_import.py tests/unit/test_onnx_rl_adapter.py +tests/test_codegen.py::TestTraceMLPPolicy tests/test_zig_lowering.py::TestOnnxPolicyOracle` +→ 90 passed. + +### 2026-09-17 — move the VM workspace off the stack (file-scope buffer) + +The element-loop fix removed the compile-time wall, which exposed the next one: +`shinro_step` declared `buf: [g.buf_len]f64` as a **stack local**, so calling it +reserved `buf_bytes` on the caller's stack. At 512×6 that is 20.8 MiB and at +512×8 28.8 MiB, past the default 16 MB stack, so the in-process oracle +segfaulted — the compiler was fine. Those two sweep points only passed under +`ulimit -s unlimited`. + +**Fix (user decision: plain global).** The kernel runs as a sequential control +loop — one caller, one tick at a time — so a process-global workspace is +acceptable. The declaration moved to file scope: + +```zig +var workspace: [g.buf_len]f64 align(16) = undefined; +``` + +`shinro_step` no longer declares a local; every `&buf` call site became +`&workspace`, and the global pointer is what the helpers receive (their +parameter stays `buf`). The buffer is `.bss` — no initializer, no heap — and is +documented as **not reentrant / thread-safe** (one call at a time), matching the +QP path's existing static `solver` global. It was named `workspace` rather than +`buf` because Zig rejects a function parameter shadowing a container-scope +declaration. + +**Verified.** + +- 512×6 (1.35 M params, 20.8 MiB workspace) builds and passes `oracle B` + (1.17e-15) under the **default 16 MB stack** — the overflow is gone. Its `.so` + is 15,354 B of `.text` + 21,776,000 B of `.bss` (exactly `buf_bytes`). +- `make test-zig` 81 passed / 2 skipped; `make test` 1241 passed / 5 skipped. +- Tick latency unchanged within noise (interleaved, CPU-pinned A/B: best-of + 41.8 µs stack vs 46.2 µs global, but the per-run spread is larger — medians + swing 106–190 µs on the same build). diff --git a/pyproject.toml b/pyproject.toml index 7e59d43..904162b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,12 +26,13 @@ dependencies = [ mujoco = ["mujoco>=3.0"] torch = ["torch>=2.0"] lerobot = ["lerobot>=0.5"] -onnx-rl = ["onnxruntime>=1.17"] +onnx-rl = ["onnx>=1.16"] media = ["imageio>=2.30", "Pillow>=10.0", "matplotlib>=3.8"] dev = [ "pytest>=8.0", "ruff>=0.5", "pyrefly>=1.3", + "onnx>=1.16", "build>=1.2", "setuptools-scm[toml]>=8.0", "git-cliff>=2.0", diff --git a/scripts/bench_tick.py b/scripts/bench_tick.py new file mode 100644 index 0000000..4739b04 --- /dev/null +++ b/scripts/bench_tick.py @@ -0,0 +1,226 @@ +"""Benchmark `shinro_step` wall time per control tick for a compiled kernel. + +The lowering oracle proves a kernel is *correct*; this measures whether it is +*fast*, so lowering changes can be A/B'd on runtime rather than argued about. +It dlopens a compiled artifact, drives the C ABI with seeded inputs, feeds the +recurrent ``state_*`` outputs back into their matching inputs (a realistic +rollout), and reports ns/tick. + +Timing samples are integer nanoseconds and the summary is plain arithmetic, so +nothing here depends on numpy's float types. + +Pass ``--onnx-model`` to also time the eager numpy path for the same policy +(``OnnxRLAdapter`` with ``model_path``), which is the honest yardstick for an +NN policy: the compiled kernel must beat it to be worth deploying. + +Usage:: + + python3 scripts/bench_tick.py --artifact-dir build/fence_before + python3 scripts/bench_tick.py --artifact-dir build/onnx_scale/H256x4/out \\ + --kernel lib_neural_network --onnx-model build/onnx_scale/H256x4/policy.onnx +""" + +from __future__ import annotations + +import argparse +import ctypes +import json +import sys +import time +from collections.abc import Callable +from pathlib import Path + +import numpy as np + +REPO = Path(__file__).resolve().parents[1] + + +def _flat(shape) -> int: + dims = list(shape or []) + n = 1 + for d in dims: + n *= d + return n + + +def _port_offsets(ports: list[dict]) -> list[tuple[int, int]]: + """Flat (start, stop) offsets of each port, in port order.""" + bounds: list[tuple[int, int]] = [] + offset = 0 + for port in ports: + size = _flat(port["shape"]) + bounds.append((offset, offset + size)) + offset += size + return bounds + + +def _load_manifest(path: Path) -> dict: + """Read a graph manifest, failing with a build hint rather than a traceback.""" + if not path.exists(): + raise SystemExit(f"no graph manifest at {path} — build the artifact first (make compile / zig build)") + try: + return json.loads(path.read_text()) + except json.JSONDecodeError as exc: + raise SystemExit(f"graph manifest at {path} is corrupt: {exc}") from exc + + +#: Aim for ~0.2 ms of work per timed sample, so clock overhead is <1% of it. +#: (A single bare call on this class of host can be shorter than the timer's +#: effective resolution, which inflates per-call medians by orders of magnitude.) +_TARGET_SAMPLE_NS = 200_000 + + +def _time_ticks(step: Callable[[], object], ticks: int, warmup: int) -> list[float]: + """Per-call nanosecond samples, timed in batches sized to dwarf timer overhead.""" + for _ in range(warmup): + step() + + start = time.perf_counter_ns() + step() + single = max(1, time.perf_counter_ns() - start) + batch = max(1, min(ticks, _TARGET_SAMPLE_NS // single)) + + samples: list[float] = [] + remaining = ticks + while remaining > 0: + n = min(batch, remaining) + start = time.perf_counter_ns() + for _ in range(n): + step() + samples.append((time.perf_counter_ns() - start) / n) + remaining -= n + return samples + + +def _stats(samples: list[float]) -> dict: + """median / mean / p99 / rate from per-tick nanosecond samples (ints).""" + ordered = sorted(samples) + n = len(ordered) + median = ordered[n // 2] if n % 2 else (ordered[n // 2 - 1] + ordered[n // 2]) / 2 + p99_index = min(n - 1, (n * 99) // 100) + return { + "ticks": n, + # min is the robust estimator on a noisy host: the fastest pass is the + # one least perturbed by scheduling, so it tracks true cost best. + "ns_per_tick_min": ordered[0], + "ns_per_tick_median": median, + "ns_per_tick_mean": sum(ordered) / n, + "ns_per_tick_p99": ordered[p99_index], + "hz_median": 1e9 / median, + } + + +def bench_so(artifact_dir: Path, kernel: str, ticks: int, warmup: int, seed: int, manifest_path: Path | None = None) -> dict: + """Time one tick of ``lib.so`` from ``artifact_dir``.""" + manifest = _load_manifest(manifest_path or (artifact_dir / "graph_data_manifest.json")) + so_path = artifact_dir / "lib" / f"{kernel}.so" + lib = ctypes.CDLL(str(so_path)) + ptr = ctypes.POINTER(ctypes.c_double) + lib.shinro_step.argtypes = [ptr, ptr, ptr] + lib.shinro_step.restype = None + + in_slices = _port_offsets(manifest["inputs"]) + state_slices = _port_offsets(manifest["state_outputs"]) + n_in = sum(stop - start for start, stop in in_slices) + n_out = sum(_flat(p["shape"]) for p in manifest["outputs"]) + n_state = sum(_flat(p["shape"]) for p in manifest["state_outputs"]) + + rng = np.random.default_rng(seed) + inputs = rng.normal(0.0, 0.1, max(n_in, 1)).copy() + outputs = np.zeros(max(n_out, 1)) + state = np.zeros(max(n_state, 1)) + + # Recurrent feedback: each state_* output is also a state_* input next tick. + in_names = [p["name"] for p in manifest["inputs"]] + feedback: list[tuple[int, int, int, int]] = [] + for name, (sstart, sstop) in zip([p["name"] for p in manifest["state_outputs"]], state_slices): + if name in in_names: + istart = in_slices[in_names.index(name)][0] + feedback.append((istart, istart + (sstop - sstart), sstart, sstop)) + + def step() -> None: + lib.shinro_step( + inputs.ctypes.data_as(ptr), + outputs.ctypes.data_as(ptr), + state.ctypes.data_as(ptr), + ) + for i0, i1, s0, s1 in feedback: + inputs[i0:i1] = state[s0:s1] + + samples = _time_ticks(step, ticks, warmup) + + return { + "artifact": f"{artifact_dir}/lib/{kernel}.so", + "node_count": manifest.get("nodes_total"), + "buf_bytes": manifest.get("buf_bytes"), + "const_blob_bytes": manifest.get("const_blob_bytes"), + **_stats(samples), + } + + +def bench_eager(model_path: str, ticks: int, warmup: int, seed: int, n_x: int | None) -> dict: + """Time one tick of the eager numpy path (the yardstick for an NN policy).""" + from shinro.controllers.onnx_rl_adapter import OnnxRLAdapter + + cfg: dict = {"model_path": model_path} + if n_x is not None: + cfg["n_x"] = n_x + ctrl = OnnxRLAdapter.from_config(cfg) + rng = np.random.default_rng(seed) + state = rng.normal(0.0, 0.1, ctrl.policy.state_size) + + samples = _time_ticks(lambda: ctrl.compute(state), ticks, warmup) + + return {"artifact": f"eager(interpret) {model_path}", **_stats(samples)} + + +def _print_row(row: dict) -> None: + label = Path(row["artifact"]).name if "/" in row["artifact"] else row["artifact"] + print( + f" {label:<30} min {row['ns_per_tick_min']:>11,.0f} ns" + f" median {row['ns_per_tick_median']:>11,.0f} ns" + f" p99 {row['ns_per_tick_p99']:>11,.0f} ns" + ) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--artifact-dir", required=True, help="dir with graph_data_manifest.json + lib/.so") + parser.add_argument("--manifest", help="graph manifest path (default: /graph_data_manifest.json)") + parser.add_argument("--kernel", default="libbase", help="artifact stem (default libbase)") + parser.add_argument("--onnx-model", help="also time the eager numpy path for this .onnx policy") + parser.add_argument("--n-x", type=int, help="plant state dim for the eager path (default: derived from the model)") + parser.add_argument("--ticks", type=int, default=20000, help="timed ticks (default 20000)") + parser.add_argument("--warmup", type=int, default=2000, help="warmup ticks (default 2000)") + parser.add_argument("--seed", type=int, default=0, help="input RNG seed") + parser.add_argument("--json", help="also write the rows to this JSON path") + args = parser.parse_args() + + rows = [ + bench_so( + Path(args.artifact_dir), + args.kernel, + args.ticks, + args.warmup, + args.seed, + Path(args.manifest) if args.manifest else None, + ) + ] + if args.onnx_model: + rows.append(bench_eager(args.onnx_model, args.ticks, args.warmup, args.seed, args.n_x)) + + print(f"=== tick benchmark: {args.ticks} ticks, {args.warmup} warmup ===") + for row in rows: + _print_row(row) + if len(rows) == 2: + compiled, eager = rows + ratio = eager["ns_per_tick_min"] / compiled["ns_per_tick_min"] + print(f" -> compiled is {ratio:.2f}x {'faster' if ratio > 1 else 'SLOWER'} than eager numpy (min)") + if args.json: + Path(args.json).write_text(json.dumps(rows, indent=2) + "\n") + print(f"wrote {args.json}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/gen_toy_onnx.py b/scripts/gen_toy_onnx.py new file mode 100644 index 0000000..fb75363 --- /dev/null +++ b/scripts/gen_toy_onnx.py @@ -0,0 +1,82 @@ +"""Generate the committed toy ONNX policy fixture. + +Deterministic — fixed hand-written weights, no RNG — so re-running reproduces +the same graph and the fixture is reviewable by reading this file. The model is +a 3 -> 4 -> 2 tanh MLP:: + + h = tanh(W1 x + b1) + action = W2 h + b2 + +with ``W1 = [[1,0,0],[0,1,0],[0,0,1],[1,1,0]]``, ``b1 = 0`` and +``W2 = [[1,0,0,0],[0,1,0,0]]``, ``b2 = [0.5, -0.5]``. Only the first two hidden +units feed the output, but all four are kept so the Gemm dimensions are +non-trivial (a real transpose, not a degenerate 1xN). The closed form is:: + + action = [tanh(x0) + 0.5, tanh(x1) - 0.5] + +It is used as the policy fixture for the ONNX importer/adapter unit tests and +for the committed policy-only compile scenario +(``tests/fixtures/configs/scenarios/toy_mlp_policy.toml``). + +Usage:: + + python3 scripts/gen_toy_onnx.py # the committed fixture + python3 scripts/gen_toy_onnx.py --out path/to/policy.onnx +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +DEFAULT_OUT = "tests/fixtures/models/toy_mlp.onnx" + + +def build_toy_mlp(): + """Build the toy 3 -> 4 -> 2 tanh MLP as an ``onnx.ModelProto``.""" + import numpy as np + from onnx import TensorProto, helper + + w1 = np.array([[1, 0, 0], [0, 1, 0], [0, 0, 1], [1, 1, 0]], dtype=np.float32) + b1 = np.zeros(4, dtype=np.float32) + w2 = np.array([[1, 0, 0, 0], [0, 1, 0, 0]], dtype=np.float32) + b2 = np.array([0.5, -0.5], dtype=np.float32) + + def init(name, array): + a = np.asarray(array, dtype=np.float32) + return helper.make_tensor(name, TensorProto.FLOAT, a.shape, a.flatten().tolist()) + + obs = helper.make_tensor_value_info("obs", TensorProto.FLOAT, [None, 3]) + action = helper.make_tensor_value_info("action", TensorProto.FLOAT, [None, 2]) + nodes = [ + helper.make_node("Gemm", ["obs", "w1", "b1"], ["h"], transB=1), + helper.make_node("Tanh", ["h"], ["a"]), + helper.make_node("Gemm", ["a", "w2", "b2"], ["action"], transB=1), + ] + graph = helper.make_graph( + nodes, + "toy_mlp", + [obs], + [action], + [init("w1", w1), init("b1", b1), init("w2", w2), init("b2", b2)], + ) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + model.ir_version = 8 + return model + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", default=DEFAULT_OUT, help=f"output path (default: {DEFAULT_OUT})") + args = parser.parse_args() + + import onnx + + out = Path(args.out) + out.parent.mkdir(parents=True, exist_ok=True) + onnx.save(build_toy_mlp(), str(out)) + print(f"wrote {out} ({out.stat().st_size} bytes)") + + +if __name__ == "__main__": + main() diff --git a/scripts/measure_onnx_policy_scale.py b/scripts/measure_onnx_policy_scale.py new file mode 100644 index 0000000..c367fd3 --- /dev/null +++ b/scripts/measure_onnx_policy_scale.py @@ -0,0 +1,229 @@ +"""ONNX policy size sweep: parameter count -> artifact size, compile time, footprint. + +Generates realistic RL-actor MLPs (``obs -> [hidden]*depth -> act``) of +increasing parameter count and compiles each through the real pipeline +(``gen_scenario.py`` → ``build_scenario.py --optimize release``), reporting the +compiled ``.so`` size, the stage wall-clock times, and the lowered VM's buffer +sizes (stack buffer, baked constants, node count). A stage that fails or times +out is recorded as data, not raised — the scaling wall is the point. + +The ONNX analogue of ``scripts/measure_kernels.py``: a one-off measurement +harness, not part of ``make test``. Everything lands under the gitignored +``build/onnx_scale/``. + + python3 scripts/measure_onnx_policy_scale.py + python3 scripts/measure_onnx_policy_scale.py --archs 64:2,256:4,512:4 + python3 scripts/measure_onnx_policy_scale.py --json build/onnx_scale.json +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import shutil +import subprocess +import sys +import time +from pathlib import Path + +REPO = Path(__file__).resolve().parents[1] +WORK = REPO / "build" / "onnx_scale" +#: The zig build's LOCAL cache (relative to the runtime build root). Clearing it +#: forces a cold compile — zig caches comptime work, and a warm cache turns a +#: 3-minute compile into ~2 seconds. +_LOCAL_ZIG_CACHE = REPO / "src" / "shinro" / "runtime" / ".zig-cache" + +#: Observation / action dimensions of the fixed policy interface. +OBS, ACT = 64, 16 +#: Default ``hidden:depth`` architectures — a realistic shallow-to-deep RL actor. +DEFAULT_ARCHS = "64:2,256:2,256:4,512:4,512:6,512:8" + + +def params_for(obs: int, act: int, hidden: int, depth: int) -> int: + """Parameter count of the ``obs -> [hidden]*depth -> act`` MLP.""" + return (obs + 1) * hidden + (depth - 1) * (hidden + 1) * hidden + (hidden + 1) * act + + +def make_mlp(obs: int, act: int, hidden: int, depth: int, seed: int): + """Build an ``obs -> [hidden]*depth -> act`` tanh MLP as an onnx ModelProto.""" + import numpy as np + from onnx import TensorProto, helper + + rng = np.random.default_rng(seed) + nodes: list = [] + inits: list = [] + + def init(name: str, array) -> object: + a = np.asarray(array, dtype=np.float32) + return helper.make_tensor(name, TensorProto.FLOAT, a.shape, a.flatten().tolist()) + + prev = "obs" + for layer in range(depth + 1): + in_dim = obs if layer == 0 else hidden + out_dim = act if layer == depth else hidden + out = "action" if layer == depth else f"h{layer}" + nodes.append(helper.make_node("Gemm", [prev, f"w{layer}", f"b{layer}"], [out], transB=1)) + inits += [ + init(f"w{layer}", rng.normal(0.0, 1.0 / math.sqrt(in_dim), (out_dim, in_dim))), + init(f"b{layer}", rng.normal(0.0, 0.1, out_dim)), + ] + if out != "action": + nodes.append(helper.make_node("Tanh", [out], [f"a{layer}"])) + prev = f"a{layer}" + + x = helper.make_tensor_value_info("obs", TensorProto.FLOAT, [None, obs]) + y = helper.make_tensor_value_info("action", TensorProto.FLOAT, [None, act]) + graph = helper.make_graph(nodes, "sweep", [x], [y], inits) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + model.ir_version = 8 + return model + + +def _env(extra: dict | None = None) -> dict: + env = {**os.environ, "PYTHONPATH": f"{REPO / 'src'}{os.pathsep}{REPO}"} + if extra: + env.update(extra) + return env + + +def _run(cmd: list[str], timeout: float, env: dict | None = None) -> tuple[float, str, str]: + """Run a stage; return ``(seconds, stdout, error)`` — ``stdout`` is "" on failure.""" + start = time.perf_counter() + try: + result = subprocess.run(cmd, capture_output=True, text=True, env=env or _env(), timeout=timeout) + except subprocess.TimeoutExpired: + return time.perf_counter() - start, "", f"TIMEOUT >{timeout:.0f}s" + elapsed = time.perf_counter() - start + if result.returncode != 0: + return elapsed, result.stdout, _summarize_error(result.stderr) + return elapsed, result.stdout, "" + + +def _summarize_error(stderr: str) -> str: + """Pull the most informative line out of a zig/python failure.""" + for line in stderr.splitlines(): + low = line.lower() + if "error:" in low or "exceeded" in low or "out of memory" in low or "panic" in low: + return line.strip()[:180] + lines = stderr.strip().splitlines() + return lines[-1][:180] if lines else "unknown failure" + + +def measure_point(hidden: int, depth: int, timeout: float, seed: int) -> dict: + """Compile one ``hidden x depth`` policy and return its measurements (or error).""" + import onnx + + n_params = params_for(OBS, ACT, hidden, depth) + row: dict = {"params": n_params, "hidden": hidden, "depth": depth} + root = WORK / f"H{hidden}x{depth}" + root.mkdir(parents=True, exist_ok=True) + policy = root / "policy.onnx" + onnx.save(make_mlp(OBS, ACT, hidden, depth, seed), str(policy)) + + ctrl = root / "ctrl.toml" + ctrl.write_text( + f'type = "onnx_rl"\nmodel_path = "{policy}"\naction_space = "continuous"\n' + f"\n[observation]\nstate_keys = {list(range(OBS))}\n" + ) + scenario = root / "scenario.toml" + scenario.write_text( + f'[controller]\nconfig = "{ctrl}"\n' + f'[compile]\nn_x = {OBS}\nn_u = {ACT}\nartifact_name = "lib_neural_network"\n' + ) + + out = root / "out" + gen_sec, _gen_out, gen_err = _run([sys.executable, "scripts/gen_scenario.py", str(scenario), "--out", str(out)], timeout) + row["onnx_bytes"] = policy.stat().st_size + row["gen_s"] = round(gen_sec, 2) + if gen_err: + row["error"] = f"gen: {gen_err}" + return row + + graph_src = out / "graph_data.zig" + row["graph_src_bytes"] = graph_src.stat().st_size if graph_src.exists() else 0 + + # Cold compile: drop the shared local cache so this point is not measured + # against another point's cached comptime evaluation. The only expected + # failure is a missing dir (first build); anything else is reported, because + # an uncleared cache would silently make the timing warm. + cold = True + if _LOCAL_ZIG_CACHE.exists(): + try: + shutil.rmtree(_LOCAL_ZIG_CACHE) + except OSError as exc: + cold = False + print(f"warning: could not clear {_LOCAL_ZIG_CACHE} ({exc}); this point may be warm", file=sys.stderr) + row["cold_cache"] = cold + build_sec, build_out, build_err = _run( + [sys.executable, "scripts/build_scenario.py", str(out), "--scenario", str(scenario), "--optimize", "release"], + timeout, + ) + row["build_s"] = round(build_sec, 2) + if build_err: + row["error"] = f"build: {build_err}" + return row + + try: + manifest = json.loads((out / "graph_data_manifest.json").read_text()) + except (OSError, json.JSONDecodeError) as exc: + row["error"] = f"manifest unreadable: {exc}" + return row + row.update( + so_bytes=(out / "lib" / "lib_neural_network.so").stat().st_size, + nodes=manifest["nodes_total"], + buf_bytes=manifest["buf_bytes"], + const_bytes=manifest["const_blob_bytes"], + oracle=next((ln.strip() for ln in build_out.splitlines() if "oracle B" in ln), ""), + ) + return row + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--archs", default=DEFAULT_ARCHS, help=f"comma-separated hidden:depth (default {DEFAULT_ARCHS})") + parser.add_argument("--timeout", type=float, default=900.0, help="per-stage timeout in seconds (default 900)") + parser.add_argument("--seed", type=int, default=0, help="weight RNG seed") + parser.add_argument("--json", help="also write the rows to this JSON path") + args = parser.parse_args() + + archs: list = [] + try: + archs = [tuple(int(p) for p in spec.split(":")) for spec in args.archs.split(",")] + except ValueError: + archs = [] + if not archs or any(len(a) != 2 for a in archs): + parser.error(f"--archs must be comma-separated hidden:depth pairs (got {args.archs!r})") + hdr = ( + f"{'params':>9} {'arch':>10} {'onnx':>9} {'graph.zig':>10} {'so':>9} " + f"{'gen s':>6} {'build s':>8} {'nodes':>6} {'vm buf':>8} {'consts':>8}" + ) + print(hdr) + print("-" * len(hdr)) + rows = [] + for hidden, depth in archs: + row = measure_point(hidden, depth, args.timeout, args.seed) + rows.append(row) + arch = f"{hidden}x{depth}" + if "error" in row: + src_mb = row.get("graph_src_bytes", 0) / 1024 / 1024 + print( + f"{row['params']:>9} {arch:>10} {row['onnx_bytes'] / 1024:>8.1f}K {src_mb:>9.1f}M" + f" FAILED: {row['error']}" + ) + continue + print( + f"{row['params']:>9} {arch:>10} {row['onnx_bytes'] / 1024:>8.1f}K" + f" {row['graph_src_bytes'] / 1024 / 1024:>9.1f}M {row['so_bytes'] / 1024:>8.1f}K" + f" {row['gen_s']:>6.2f} {row['build_s']:>8.2f} {row['nodes']:>6}" + f" {row['buf_bytes'] / 1024:>7.1f}K {row['const_bytes'] / 1024:>7.1f}K" + ) + if args.json: + Path(args.json).write_text(json.dumps(rows, indent=2) + "\n") + print(f"wrote {args.json}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/src/shinro/codegen/cli.py b/src/shinro/codegen/cli.py index f3d2d6a..bf6c72a 100644 --- a/src/shinro/codegen/cli.py +++ b/src/shinro/codegen/cli.py @@ -29,6 +29,7 @@ def main() -> int: parser.add_argument("--optimize", choices=["debug", "release"], help="override [compile].optimize") parser.add_argument("--target", help="override [compile].target (zig triple, e.g. aarch64-linux-gnu)") parser.add_argument("--solver-dir", help="override [compile].solver_dir (baked OSQP solver dir)") + parser.add_argument("--artifact-name", help="override [compile].artifact_name (kernel installs as lib/.so)") parser.add_argument("--samples", type=int, default=20, help="random inputs for the oracle (default 20)") parser.add_argument("--seed", type=int, default=0, help="RNG seed for the oracle (default 0)") args = parser.parse_args() @@ -57,6 +58,7 @@ def main() -> int: optimize=args.optimize, target=args.target, solver_dir=args.solver_dir, + artifact_name=args.artifact_name, samples=args.samples, seed=args.seed, ) diff --git a/src/shinro/codegen/onnx_import.py b/src/shinro/codegen/onnx_import.py new file mode 100644 index 0000000..810f4ce --- /dev/null +++ b/src/shinro/codegen/onnx_import.py @@ -0,0 +1,701 @@ +"""Import an ONNX-exported policy as a shinro graph. + +Unlike every other component in the framework, this does **not** go through the +tracer. ONNX is already a dataflow graph — ``onnx.load(path).graph`` is a +topologically-sorted list of ``(op_type, inputs, outputs, attributes)`` records +with the weights inlined as initializers — so the importer translates it +directly into :class:`~shinro.codegen.tracing.Graph` nodes. The result is a +memoryless :class:`~shinro.codegen.compose.ComposedGraph` that runs through the +same two execution paths as every other compiled graph: + +- :func:`shinro.codegen.interpreter.interpret` — the pure-numpy f64 reference, + and also the eager controller path when no ``.so`` has been built, and +- :func:`shinro.codegen.lower_zig.lower_zig` + ``zig build`` — the deployable + C-ABI ``.so``. + +The policy's observation encoder (integer index selection, mean/std +normalization, clipping) is folded into the graph as arithmetic on baked +constants, so the compiled artifact's only input port is the raw plant state +(``state``). ``onnxruntime`` is not involved anywhere: ``onnx`` is a +compile-time-only dependency (the graph reader), and the deployed artifact has +no runtime dependency at all. + +Supported ONNX surface — everything else raises ``NotImplementedError``: + +- ``Gemm`` (``alpha`` / ``beta`` / ``transB``; ``transA=1`` is rejected), + ``MatMul``, ``Add`` +- ``Relu``, ``Tanh``, ``Sigmoid`` (composed from ``exp`` / ``neg`` / ``add`` / + ``div`` so no new VM op is needed) + +Nodes that do not contribute to the declared output are ignored, so an +exporter's stray logging/cast node does not fail the import. + +The batch axis is interpreted as a single sample: the policy's declared input +shape ``(None, n_obs)`` becomes the 1-D ``state`` port, and the emitted graph +stays 1-D (or 2-D where an initializer forces it) exactly like the classical +controllers. ``observation.add_batch_dim`` is consequently a no-op in the +compiled path — it only mattered for feeding ``onnxruntime``. + +The action space is baked too (``action_cfg``): ``continuous`` applies +scale/bias, ``discrete`` emits an argmax one-hot, and ``stochastic`` splits +``[mean; log_std]``. A non-deterministic ``discrete`` / ``stochastic`` policy +gains an ``epsilon`` input port — the host supplies the noise (Gumbel for +discrete, standard normal for stochastic) and the kernel does only the +arithmetic, exactly like MPPI's port. Deterministic policies have no +``epsilon`` port at all. + +Usage:: + + from shinro.codegen.interpreter import interpret + from shinro.codegen.onnx_import import import_onnx_policy + + cg = import_onnx_policy("policy.onnx", n_x=6, obs_cfg={"state_keys": [0, 1, 2]}) + u = interpret(cg.graph, {"state": state})["u"] +""" + +from __future__ import annotations + +from typing import Any + +import numpy as np + +from shinro.codegen.compose import ComposedGraph +from shinro.codegen.ops import OP_HANDLERS, has_op, missing_op_error +from shinro.codegen.tracing import Graph, Node + +#: C-ABI input port carrying the raw plant state (mirrors ``Controller.compute(state, ...)``). +STATE_PORT = "state" +#: C-ABI output port carrying the policy's action. +OUTPUT_PORT = "u" +#: C-ABI input port carrying host noise for a non-deterministic action space. +EPSILON_PORT = "epsilon" + +#: Action-space names the importer understands. +_ACTION_SPACES = frozenset({"continuous", "discrete", "stochastic"}) +#: Observation-config keys the importer reads. Unknown keys are rejected so a +#: typo (``obs_means``) cannot silently drop normalization or clipping. +_OBS_KEYS = frozenset({"input_name", "state_keys", "normalize", "obs_mean", "obs_std", "clip", "add_batch_dim"}) +#: ONNX ops the importer can translate. Everything else is rejected loudly. +_SUPPORTED_OPS = frozenset({"Gemm", "MatMul", "Add", "Relu", "Tanh", "Sigmoid"}) +#: ONNX ops translated straight to a same-named shinro op. +_UNARY_OPS = {"Relu": "relu", "Tanh": "tanh"} +#: Attributes Gemm may carry; any other attribute is rejected. +_GEMM_ATTRS = frozenset({"alpha", "beta", "transA", "transB"}) + + +class _OnnxImporter: + """Translate one ONNX policy into a memoryless shinro graph. + + One object owns the whole translation state — the graph under construction, + the eagerly-computed value of every node (the shape oracle), the ONNX + tensor-name → node-id map, the baked initializers, and the resolved + observation config — so the op translators are methods referring to a + single graph rather than free functions threading a builder through every + call. + + Each emitted node is immediately evaluated with its registered numpy + handler from :mod:`shinro.codegen.ops`, and the resulting array's shape + becomes the node's declared shape. The ops registry is therefore the single + source of truth for interpreter *and* lowering semantics — there is no + parallel hand-written shape table to drift — and an unregistered op fails + here, at import time, with the registry's actionable message. + + Initializers are emitted lazily (first :meth:`tensor` reference), so weights + the graph never reads do not land in the compiled constant blob. + """ + + def __init__( + self, + model_path: str, + *, + n_x: int | None = None, + obs_cfg: dict | None = None, + action_cfg: dict | None = None, + output_name: str | None = None, + ) -> None: + """Store the import request; parsing and emission happen in :meth:`build`. + + Args: + model_path: Path to the ``.onnx`` policy file. + n_x: Plant state dimension — the length of the ``state`` port. + Defaults to ``max(state_keys) + 1`` (or the model's observation + dimension when no ``state_keys`` are given). + obs_cfg: Observation-encoder config, mirroring the old adapter's + ``[observation]`` TOML table. + action_cfg: Action-space config, mirroring the old adapter's + top-level TOML fields. + output_name: ONNX tensor to import as the action. Defaults to the + model's first declared output. + """ + self.model_path = model_path + self.requested_n_x = n_x + self.obs_cfg = dict(obs_cfg or {}) + self.action_cfg = dict(action_cfg or {}) + self.output_name = output_name + + self.g = Graph() + self.values: dict[int, np.ndarray] = {} + self.tensors: dict[str, int] = {} + self.inputs: dict[str, np.ndarray] = {} + self.initializers: dict[str, np.ndarray] = {} + + # Overwritten by _resolve_action_cfg(); the defaults keep the object + # introspectable before build() runs. + self.action_space = "continuous" + self.deterministic = True + self.action_scale = np.asarray(1.0, dtype=np.float64) + self.action_bias = np.asarray(0.0, dtype=np.float64) + self.action_clip: tuple[float, float] | None = None + + # ── orchestration ───────────────────────────────────────────────────── + + def build(self) -> ComposedGraph: + """Parse the model, emit the graph, and return the composed result. + + Returns: + A :class:`ComposedGraph` with a ``state`` input port (plus + ``epsilon`` when the action space samples), a ``u`` output port, + and no recurrent state. + + Raises: + ValueError: On a malformed model/config, a multi-input policy, an + unresolvable tensor, a batched (leading dim ≠ 1) output, or an + inconsistent action config. + NotImplementedError: On an unsupported ONNX op or attribute. + """ + self._resolve_action_cfg() + onnx, numpy_helper = _onnx_modules() + graph_proto = onnx.load(self.model_path).graph + self.initializers = {t.name: np.asarray(numpy_helper.to_array(t), dtype=np.float64) for t in graph_proto.initializer} + + input_name, state_keys, n_x = self._resolve_ports(graph_proto) + resolved_output = self.output_name or graph_proto.output[0].name + nodes = list(graph_proto.node) + needed = _needed_node_indices(nodes, resolved_output, self.initializers, input_name) + + self.inputs = {STATE_PORT: np.ones(n_x, dtype=np.float64)} + state_id = self.emit("input", [], name=STATE_PORT) + # The encoder output takes the place of the ONNX input, so the network + # reads the encoded observation while the graph's only port stays the + # raw state. + self.bind(input_name, self.fold_encoder(state_id, n_x=n_x, state_keys=state_keys)) + + for idx in sorted(needed): + node = nodes[idx] + self.emit_node(node, _onnx_attrs(node)) + + if resolved_output not in self.tensors: + raise ValueError(f"ONNX output {resolved_output!r} was not produced by any reachable node") + + action_id = self.flatten_output(self.tensors[resolved_output]) + ports = [STATE_PORT] + epsilon_id = self._emit_epsilon(action_id, ports) if self.samples_actions else None + u = self.apply_action(action_id, epsilon_id) + self.emit("output", [u], name=OUTPUT_PORT) + return ComposedGraph(graph=self.g, inputs=ports, outputs=[OUTPUT_PORT]) + + def _resolve_ports(self, graph_proto: Any) -> tuple[str, list[int], int]: + """Resolve the input tensor name, observation indices, and ``n_x``. + + Validates the model's declared input against the observation config, so + a ``state_keys`` list that disagrees with the ONNX input (or reaches + past ``n_x``) fails here rather than producing a silently mis-wired + encoder. + + Args: + graph_proto: The model's ``GraphProto``. + + Returns: + ``(input_name, state_keys, n_x)``. + + Raises: + ValueError: On multiple graph inputs, an ``input_name`` override + that does not match, an unknown observation key, an + un-inferable observation dimension, or out-of-range + ``state_keys``. + """ + unknown = set(self.obs_cfg) - _OBS_KEYS + if unknown: + raise ValueError(f"observation has unknown key(s): {sorted(unknown)} — valid keys: {sorted(_OBS_KEYS)}") + real_inputs = [i for i in graph_proto.input if i.name not in self.initializers] + if len(real_inputs) != 1: + raise ValueError( + f"ONNX policy must declare exactly one non-initializer input " + f"(got {[i.name for i in real_inputs]}); multi-input policies are not supported" + ) + input_name = self.obs_cfg.get("input_name") or real_inputs[0].name + if input_name != real_inputs[0].name: + raise ValueError(f"observation.input_name {input_name!r} is not the model's input {real_inputs[0].name!r}") + + state_keys = self.obs_cfg.get("state_keys") + declared_obs = _declared_last_dim(real_inputs[0]) + if state_keys is None: + if declared_obs is None: + raise ValueError("cannot infer the observation dimension from the ONNX input shape — set observation.state_keys") + state_keys = list(range(declared_obs)) + state_keys = [int(k) for k in state_keys] + if declared_obs is not None and declared_obs != len(state_keys): + raise ValueError(f"observation.state_keys selects {len(state_keys)} entries but the ONNX input declares {declared_obs}") + + n_x = self.requested_n_x + if n_x is None: + n_x = max(state_keys) + 1 if state_keys else 0 + if n_x <= 0 or any(k < 0 or k >= n_x for k in state_keys): + raise ValueError(f"observation.state_keys {state_keys} out of range for n_x={n_x}") + return input_name, state_keys, n_x + + # ── graph emission ──────────────────────────────────────────────────── + + def emit(self, op: str, inputs: list[int], **attrs: Any) -> int: + """Evaluate ``op`` eagerly, then append it with the resulting shape. + + Args: + op: Registered shinro op name. + inputs: Node ids of the operands. + **attrs: Op-specific attributes (baked ``value``, ``target_shape``, + ``name``, ...). + + Returns: + The new node's id. + + Raises: + NotImplementedError: If ``op`` is not in the registry. + """ + if not has_op(op): + raise missing_op_error(op) + probe = Node(op=op, inputs=list(inputs), shape=(), attrs=dict(attrs)) + value = np.asarray(OP_HANDLERS[op](probe, self.values, self.inputs), dtype=np.float64) + node_id = self.g.emit(op, inputs, value.shape, **attrs) + self.values[node_id] = value + return node_id + + def const(self, value: Any) -> int: + """Emit a ``const`` node carrying ``value`` as an f64 array.""" + return self.emit("const", [], value=np.asarray(value, dtype=np.float64)) + + def tensor(self, name: str) -> int: + """Resolve an ONNX tensor name to its node id, baking initializers lazily. + + Raises: + ValueError: If ``name`` is neither already produced nor an initializer. + """ + node_id = self.tensors.get(name) + if node_id is None: + value = self.initializers.get(name) + if value is None: + raise ValueError(f"ONNX tensor {name!r} is neither a produced value nor an initializer") + node_id = self.const(value) + self.tensors[name] = node_id + return node_id + + def bind(self, name: str, node_id: int) -> None: + """Record that ONNX tensor ``name`` is now produced by ``node_id``.""" + self.tensors[name] = node_id + + def value_of(self, node_id: int) -> np.ndarray: + """Return the eagerly-computed value of ``node_id`` (the shape oracle).""" + return self.values[node_id] + + # ── observation encoder ─────────────────────────────────────────────── + + def fold_encoder(self, state_id: int, *, n_x: int, state_keys: list[int]) -> int: + """Fold the observation encoder into the graph as arithmetic nodes. + + Mirrors the old runtime encoder exactly, in order: integer index + selection (a baked 0/1 selection matrix, i.e. a matmul), mean/std + normalization (``sub`` / ``div`` with baked constants), then clipping + (``clip``). When the selection is the identity the matmul is skipped and + the state feeds straight through. + + Args: + state_id: Node id of the raw ``state`` input port. + n_x: Plant state dimension. + state_keys: Integer indices of the state used as observations. + + Returns: + Node id of the encoded observation vector. + + Raises: + ValueError: On a normalization request without mean/std, or a + constant whose length does not match the observation dimension. + """ + n_obs = len(state_keys) + obs = state_id + if state_keys != list(range(n_obs)) or n_x != n_obs: + selection = np.zeros((n_x, n_obs), dtype=np.float64) + selection[state_keys, np.arange(n_obs)] = 1.0 + obs = self.emit("matmul", [obs, self.const(selection)]) + + if self.obs_cfg.get("normalize", False): + mean = self.obs_cfg.get("obs_mean") + std = self.obs_cfg.get("obs_std") + if mean is None or std is None: + raise ValueError("observation.normalize requires both obs_mean and obs_std") + obs = self.emit("sub", [obs, self.const(_obs_vector(mean, n_obs, "obs_mean"))]) + obs = self.emit("div", [obs, self.const(_obs_vector(std, n_obs, "obs_std"))]) + + if "clip" in self.obs_cfg: + lo, hi = self.obs_cfg["clip"] + obs = self.emit("clip", [obs], lo=float(lo), hi=float(hi)) + return obs + + # ── ONNX op translation ─────────────────────────────────────────────── + + def emit_node(self, node: Any, attrs: dict[str, Any]) -> None: + """Translate one ONNX node into shinro node(s) and bind its output tensor. + + Args: + node: The ``onnx.NodeProto`` to translate. + attrs: Pre-extracted node attributes. + + Raises: + NotImplementedError: On an unsupported op, attribute, or arity. + """ + op_type = node.op_type + inputs = [n for n in node.input if n] + outputs = [n for n in node.output if n] + if op_type not in _SUPPORTED_OPS: + raise NotImplementedError( + f"ONNX op {op_type!r} is not supported by the policy importer. Supported ops: " + f"{sorted(_SUPPORTED_OPS)}. Decompose the policy to Gemm/MatMul/Add + Relu/Tanh/Sigmoid, " + f"or extend shinro.codegen.onnx_import." + ) + if len(outputs) != 1: + raise NotImplementedError(f"ONNX op {op_type!r} must have exactly one output (got {outputs})") + + if op_type == "Gemm": + result = self.emit_gemm(inputs, attrs) + elif op_type == "MatMul": + _require_no_attrs(op_type, attrs) + result = self.emit("matmul", [self.tensor(inputs[0]), self.tensor(inputs[1])]) + elif op_type == "Add": + _require_no_attrs(op_type, attrs) + result = self.emit("add", [self.tensor(inputs[0]), self.tensor(inputs[1])]) + elif op_type in _UNARY_OPS: + _require_no_attrs(op_type, attrs) + result = self.emit(_UNARY_OPS[op_type], [self.tensor(inputs[0])]) + else: # Sigmoid + _require_no_attrs(op_type, attrs) + result = self.emit_sigmoid(self.tensor(inputs[0])) + + self.bind(outputs[0], result) + + def emit_gemm(self, inputs: list[str], attrs: dict[str, Any]) -> int: + """Emit ``Y = alpha * A' * B' + beta * C`` from matmul/transpose/mul/add. + + ``B'`` is realized as a ``transpose`` node rather than by baking a + pre-transposed constant, so the non-square transpose path in the VM is + exercised by every torch-style export (``transB=1``). The ``alpha`` / + ``beta`` multipliers are emitted unconditionally, including the default + ``1.0``: a ``mul`` by one is cheap and keeps the Gemm lowering uniform + rather than branching on attribute values. + + Raises: + NotImplementedError: On unknown attributes, wrong arity, or ``transA=1``. + """ + unknown = set(attrs) - _GEMM_ATTRS + if unknown: + raise NotImplementedError(f"ONNX Gemm carries unsupported attribute(s): {sorted(unknown)}") + if len(inputs) not in (2, 3): + raise NotImplementedError(f"ONNX Gemm must have 2 or 3 inputs (got {len(inputs)})") + if int(attrs.get("transA", 0)): + raise NotImplementedError("ONNX Gemm transA=1 is not supported (transpose the activation upstream)") + + alpha = float(attrs.get("alpha", 1.0)) + beta = float(attrs.get("beta", 1.0)) + b = self.tensor(inputs[1]) + if int(attrs.get("transB", 0)): + b = self.emit("transpose", [b]) + y = self.emit("matmul", [self.tensor(inputs[0]), b]) + y = self.emit("mul", [y, self.const(alpha)]) + if len(inputs) == 3: + c = self.tensor(inputs[2]) + c = self.emit("mul", [c, self.const(beta)]) + y = self.emit("add", [y, c]) + return y + + def emit_sigmoid(self, x_id: int) -> int: + """Emit ``sigmoid(x) = 1 / (1 + exp(-x))`` from existing VM ops.""" + exp_neg = self.emit("exp", [self.emit("neg", [x_id])]) + return self.emit("div", [self.const(1.0), self.emit("add", [self.const(1.0), exp_neg])]) + + def flatten_output(self, node_id: int) -> int: + """Reduce a batch-1 output to a 1-D action vector. + + The graph's action port is ``(n_u,)`` like every classical controller. A + network whose last op produced ``(1, n_u)`` (a rank-2 initializer can + force that) is reshaped; a genuinely batched output is rejected. + + Args: + node_id: Node id producing the raw policy output. + + Returns: + Node id of the flattened ``(n_u,)`` action. + + Raises: + ValueError: If the output is not 1-D or batch-1 2-D. + """ + value = self.value_of(node_id) + if value.ndim == 1: + return node_id + if value.ndim == 2 and value.shape[0] == 1: + return self.emit("reshape", [node_id], target_shape=(value.shape[1],)) + raise ValueError( + f"policy output shape {value.shape} is not a single action vector — " + f"only batch-1 policies (a leading dimension of 1) are supported" + ) + + # ── action space ────────────────────────────────────────────────────── + + @property + def samples_actions(self) -> bool: + """Whether the graph consumes host noise instead of a deterministic action.""" + return self.action_space != "continuous" and not self.deterministic + + def _resolve_action_cfg(self) -> None: + """Validate the action config and resolve the constants it bakes. + + Raises: + ValueError: On an unknown action space, a single-sided clip, or a + non-finite clip bound. The lowerer writes floats as Zig hex + literals and ``inf`` is not a Zig identifier, so an ``±inf`` + bound would fail the build — rejecting it here keeps the error + at import time. + """ + space = self.action_cfg.get("action_space", "continuous") + if space not in _ACTION_SPACES: + raise ValueError(f"action_space must be one of {sorted(_ACTION_SPACES)}, got {space!r}") + self.action_space = space + self.deterministic = bool(self.action_cfg.get("deterministic", True)) + self.action_scale = np.asarray(self.action_cfg.get("action_scale", 1.0), dtype=np.float64) + self.action_bias = np.asarray(self.action_cfg.get("action_bias", 0.0), dtype=np.float64) + + has_low = "action_clip_low" in self.action_cfg + has_high = "action_clip_high" in self.action_cfg + if has_low != has_high: + raise ValueError( + "action_clip_low and action_clip_high must be given together: a missing bound " + "would default to ±inf, which the lowerer cannot emit (inf is not a Zig literal)" + ) + clip = None + if has_low: + clip = (float(self.action_cfg["action_clip_low"]), float(self.action_cfg["action_clip_high"])) + if not all(np.isfinite(clip)): + raise ValueError(f"action_clip_low/action_clip_high must be finite (got {clip})") + self.action_clip = clip + + def apply_action(self, raw_id: int, epsilon_id: int | None = None) -> int: + """Translate the raw policy output into the graph's action port. + + Mirrors the old runtime post-processing exactly: ``continuous`` applies + scale/bias; ``discrete`` emits an argmax one-hot (scale/bias do not + apply to a one-hot action, matching the adapter); ``stochastic`` splits + ``[mean; log_std]`` and either returns the mean or adds + ``exp(clip(log_std, -10, 2)) * epsilon``. The optional clip is applied + last in every space. + + Args: + raw_id: Node id of the flattened raw policy output. + epsilon_id: Node id of the host-noise port when the space samples, + else ``None``. + + Returns: + Node id of the final action. + """ + if self.action_space == "continuous": + u = self._scale_bias(raw_id) + elif self.action_space == "discrete": + logits = raw_id if epsilon_id is None else self.emit("add", [raw_id, epsilon_id]) + u = self.emit("one_hot", [self.emit("argmax", [logits])], depth=self.value_of(raw_id).size) + else: # stochastic + u = self._stochastic(raw_id, epsilon_id) + if self.action_clip is not None: + lo, hi = self.action_clip + u = self.emit("clip", [u], lo=lo, hi=hi) + return u + + def _emit_epsilon(self, raw_id: int, ports: list[str]) -> int: + """Emit the host-noise input port and return its node id. + + The noise kind is part of the deployment contract: ``stochastic`` + expects standard-normal draws of length ``n_u``; ``discrete`` expects + Gumbel noise of length ``n_actions`` (``g = -log(-log(u))`` from + ``u ~ U(0, 1)``), which turns ``argmax(logits + g)`` into exact + categorical sampling from ``softmax(logits)``. + """ + n_noise = self.value_of(raw_id).size if self.action_space == "discrete" else self._stochastic_half(raw_id) + self.inputs[EPSILON_PORT] = np.zeros(n_noise, dtype=np.float64) + ports.append(EPSILON_PORT) + return self.emit("input", [], name=EPSILON_PORT) + + def _scale_bias(self, x_id: int) -> int: + """Apply ``scale * x + bias``, emitting both nodes unconditionally. + + Like Gemm's ``alpha`` / ``beta`` multipliers, the default ``1.0`` / + ``0.0`` still produce their ``mul`` / ``add``: a no-op node is cheap and + it keeps the action lowering uniform instead of branching on configured + values. + """ + x = x_id + x = self.emit("mul", [x, self.const(self.action_scale)]) + x = self.emit("add", [x, self.const(self.action_bias)]) + return x + + def _stochastic_half(self, raw_id: int) -> int: + """Return ``n_u`` for a ``[mean; log_std]`` output, validating its size. + + Raises: + ValueError: If the output size is zero or odd. + """ + size = self.value_of(raw_id).size + if size == 0 or size % 2: + raise ValueError(f"stochastic policy output must be [mean; log_std] with an even, non-zero size (got {size})") + return size // 2 + + def _stochastic(self, raw_id: int, epsilon_id: int | None) -> int: + """Split ``[mean; log_std]`` and, when sampling, add the scaled noise.""" + half = self._stochastic_half(raw_id) + mean = self.emit("slice", [raw_id], start=0, stop=half) + if epsilon_id is None: + return self._scale_bias(mean) + log_std = self.emit("slice", [raw_id], start=half, stop=2 * half) + std = self.emit("exp", [self.emit("clip", [log_std], lo=-10.0, hi=2.0)]) + return self._scale_bias(self.emit("add", [mean, self.emit("mul", [std, epsilon_id])])) + + +def import_onnx_policy( + model_path: str, + *, + n_x: int | None = None, + obs_cfg: dict | None = None, + action_cfg: dict | None = None, + output_name: str | None = None, +) -> ComposedGraph: + """Translate an ONNX policy into a memoryless composed graph. + + The graph's declared output is resolved from the model (or ``output_name``), + only the nodes it depends on are imported, and the observation encoder is + folded in front of the network. The returned graph has one input port + (``state``), one output port (``u``), and no recurrent state. + + Args: + model_path: Path to the ``.onnx`` policy file. + n_x: Plant state dimension — the length of the ``state`` port. Defaults + to ``max(state_keys) + 1`` (or the model's observation dimension + when no ``state_keys`` are given). + obs_cfg: Observation-encoder config, mirroring the old adapter's + ``[observation]`` TOML table. Supported keys: ``input_name``, + ``state_keys``, ``normalize``, ``obs_mean``, ``obs_std``, ``clip``, + and (accepted but ignored in the compiled path) ``add_batch_dim``. + action_cfg: Action-space config, mirroring the old adapter's top-level + TOML fields: ``action_space``, ``deterministic``, ``action_scale``, + ``action_bias``, and ``action_clip_low`` / ``action_clip_high`` + (which must be given together — the lowerer cannot emit ``±inf``). + output_name: ONNX tensor to import as the action. Defaults to the + model's first declared output. + + Returns: + A :class:`ComposedGraph` ready for ``interpret`` or ``lower_zig``. Its + input ports are ``state`` and, for a sampling action space, + ``epsilon``. + + Raises: + ImportError: If the ``onnx`` package is not installed. + ValueError: On a malformed model/config or an unsupported layout. + NotImplementedError: On an unsupported ONNX op or attribute. + """ + return _OnnxImporter(model_path, n_x=n_x, obs_cfg=obs_cfg, action_cfg=action_cfg, output_name=output_name).build() + + +def _onnx_modules() -> tuple[Any, Any]: + """Import and return the lazily-required ``(onnx, numpy_helper)`` modules. + + Raises: + ImportError: If the ``onnx`` package (the ``onnx-rl`` extra) is missing. + """ + try: + import onnx + from onnx import numpy_helper + except ImportError as exc: # pragma: no cover - exercised only without the extra + raise ImportError("the ONNX policy importer needs the 'onnx' package — install with `pip install \"shinro[onnx-rl]\"`") from exc + return onnx, numpy_helper + + +def _needed_node_indices(nodes: list[Any], output_name: str, initializers: dict[str, np.ndarray], input_name: str) -> set[int]: + """Collect the node indices the declared output actually depends on. + + A backwards walk from ``output_name`` through tensor producers. Exporters + routinely leave nodes that do not feed the output (logging, unused + branches), and those must not fail the import — only the reachable subgraph + is translated, so an unsupported op is rejected exactly when it matters. + + Args: + nodes: The ONNX graph's nodes, in topological order. + output_name: ONNX tensor name of the requested output. + initializers: Initializer names (leaves of the walk). + input_name: The model's real input tensor name (leaf of the walk). + + Returns: + Indices into ``nodes`` of the reachable nodes. + + Raises: + ValueError: If the graph references a tensor nothing produces. + """ + producer: dict[str, int] = {} + for idx, node in enumerate(nodes): + for out in node.output: + if out: + producer[out] = idx + + needed: set[int] = set() + visited: set[str] = set() + stack = [output_name] + while stack: + name = stack.pop() + if name in visited: + continue + visited.add(name) + if name in initializers: + continue + idx = producer.get(name) + if idx is None: + if name == input_name: + continue + raise ValueError(f"ONNX graph references unknown tensor {name!r}") + needed.add(idx) + stack.extend(n for n in nodes[idx].input if n) + return needed + + +def _require_no_attrs(op_type: str, attrs: dict[str, Any]) -> None: + """Raise if a pointwise ONNX op carries attributes the importer ignores.""" + if attrs: + raise NotImplementedError(f"ONNX op {op_type!r} carries unsupported attribute(s): {sorted(attrs)}") + + +def _declared_last_dim(value_info: Any) -> int | None: + """Return the model input's declared last dimension, or None if symbolic.""" + dims = value_info.type.tensor_type.shape.dim + if not dims: + return None + last = dims[-1] + return int(last.dim_value) if last.dim_value and last.dim_value > 0 else None + + +def _obs_vector(values: Any, n_obs: int, field: str) -> np.ndarray: + """Validate and flatten an encoder constant to length ``n_obs``. + + Raises: + ValueError: If the constant does not have exactly ``n_obs`` entries. + """ + arr = np.asarray(values, dtype=np.float64).ravel() + if arr.size != n_obs: + raise ValueError(f"observation.{field} must have {n_obs} entries (got {arr.size})") + return arr + + +def _onnx_attrs(node: Any) -> dict[str, Any]: + """Extract an ONNX node's attributes as plain Python values.""" + from onnx import helper + + return {a.name: helper.get_attribute_value(a) for a in node.attribute} diff --git a/src/shinro/codegen/oracle.py b/src/shinro/codegen/oracle.py index 1ec622c..44b224b 100644 --- a/src/shinro/codegen/oracle.py +++ b/src/shinro/codegen/oracle.py @@ -87,11 +87,11 @@ def random_inputs(cg, rng: np.random.Generator) -> dict[str, np.ndarray]: return inputs -def load_so(prefix: str | Path): - """dlopen ``/lib/libbase.so`` and wire up the shinro_step C ABI.""" - so_path = Path(prefix) / "lib" / "libbase.so" +def load_so(prefix: str | Path, name: str = "libbase"): + """dlopen ``/lib/.so`` and wire up the shinro_step C ABI.""" + so_path = Path(prefix) / "lib" / f"{name}.so" if not so_path.exists(): - raise FileNotFoundError(f"zig build produced no libbase.so at {so_path}") + raise FileNotFoundError(f"zig build produced no {so_path.name} at {so_path}") lib = ctypes.CDLL(str(so_path)) lib.shinro_step.argtypes = [ctypes.POINTER(ctypes.c_double)] * 3 lib.shinro_step.restype = None diff --git a/src/shinro/codegen/scenario_build.py b/src/shinro/codegen/scenario_build.py index 228a28c..43b2275 100644 --- a/src/shinro/codegen/scenario_build.py +++ b/src/shinro/codegen/scenario_build.py @@ -71,13 +71,13 @@ class BuildError(RuntimeError): def _check_zig() -> bool: if shutil.which("zig") is not None: return True - print("ERROR: zig not on PATH — the e2e workflow needs it to compile libbase.so.", file=sys.stderr) + print("ERROR: zig not on PATH — the e2e workflow needs it to compile the kernel.", file=sys.stderr) print("Install: https://ziglang.org/download (or your package manager).", file=sys.stderr) print("Nothing was built.", file=sys.stderr) return False -def _build(graph_path: Path, prefix: Path, optimize: str, target: str, solver_dir: str | None) -> None: +def _build(graph_path: Path, prefix: Path, optimize: str, target: str, solver_dir: str | None, name: str = "libbase") -> None: """Compile the comptime VM against the given graph into an isolated prefix.""" cmd = [ "zig", @@ -87,6 +87,7 @@ def _build(graph_path: Path, prefix: Path, optimize: str, target: str, solver_di "--prefix", str(prefix), f"-Dgraph={graph_path}", + f"-Dname={name}", ] if optimize == "release": cmd += ["-Doptimize=ReleaseFast"] @@ -127,10 +128,11 @@ def build_scenario( optimize: str | None = None, target: str | None = None, solver_dir: str | None = None, + artifact_name: str | None = None, samples: int = 20, seed: int = 0, ) -> int: - """Compile a generated graph into a verified, stamped ``libbase.so``. + """Compile a generated graph into a verified, stamped kernel (``lib.so``). Args: graph_dir: Directory containing ``graph_data.zig`` + its manifest @@ -143,6 +145,8 @@ def build_scenario( target: Override ``[compile].target`` (zig triple, e.g. ``aarch64-linux-gnu``). solver_dir: Override ``[compile].solver_dir`` (baked OSQP solver dir). + artifact_name: Override ``[compile].artifact_name`` — the kernel is + installed as ``lib/.so`` (default ``libbase`` → ``libbase.so``). samples: Random inputs for the oracle (default 20). seed: RNG seed for the oracle (default 0). @@ -169,7 +173,7 @@ def build_scenario( return EXIT_USAGE # Build flags: CLI > [compile] TOML > defaults. - opt, tgt, sdir = "debug", "native", None + opt, tgt, sdir, name = "debug", "native", None, "libbase" if scenario: try: spec = load_scenario(scenario) @@ -179,12 +183,15 @@ def build_scenario( opt = spec["compile"]["optimize"] tgt = spec["compile"]["target"] sdir = spec["compile"]["solver_dir"] + name = spec["compile"]["artifact_name"] if optimize: opt = optimize if target: tgt = target if solver_dir: sdir = solver_dir + if artifact_name: + name = artifact_name if manifest["has_solve_qp"] and not sdir: print( @@ -198,7 +205,7 @@ def build_scenario( prefix_path = Path(prefix) if prefix else graph_dir try: - _build(graph_path, prefix_path, opt, tgt, sdir) + _build(graph_path, prefix_path, opt, tgt, sdir, name) except BuildError as e: print(f"BUILD FAILED: {e}", file=sys.stderr) return EXIT_BUILD @@ -216,7 +223,7 @@ def build_scenario( # [compile].oracle_tol overrides the tier default for QP graphs # whose settling at this problem size is coarser than 1e-3. tol = spec["compile"].get("oracle_tol") or tol_for(manifest) - lib = load_so(prefix_path) + lib = load_so(prefix_path, name) max_err = run_oracle(lib, fresh_cg, samples, seed) if max_err >= tol: print( @@ -238,8 +245,8 @@ def build_scenario( file=sys.stderr, ) - stamp(prefix_path, RUNTIME) - record = prefix_path / "lib" / "libbase.deployment.json" + stamp(prefix_path, RUNTIME, name) + record = prefix_path / "lib" / f"{name}.deployment.json" if verify(record, graph_path=graph_path) != 0: print("VERIFY FAILED: deployment record does not match artifacts", file=sys.stderr) return EXIT_BUILD @@ -254,6 +261,7 @@ def main() -> int: parser.add_argument("--optimize", choices=["debug", "release"], help="override [compile].optimize") parser.add_argument("--target", help="override [compile].target (zig triple, e.g. aarch64-linux-gnu)") parser.add_argument("--solver-dir", help="override [compile].solver_dir (baked OSQP solver dir)") + parser.add_argument("--artifact-name", help="override [compile].artifact_name (kernel installs as lib/.so)") parser.add_argument("--samples", type=int, default=20, help="random inputs for the oracle (default 20)") parser.add_argument("--seed", type=int, default=0, help="RNG seed for the oracle (default 0)") args = parser.parse_args() @@ -264,6 +272,7 @@ def main() -> int: optimize=args.optimize, target=args.target, solver_dir=args.solver_dir, + artifact_name=args.artifact_name, samples=args.samples, seed=args.seed, ) diff --git a/src/shinro/codegen/scenario_gen.py b/src/shinro/codegen/scenario_gen.py index 9ab3054..9524909 100644 --- a/src/shinro/codegen/scenario_gen.py +++ b/src/shinro/codegen/scenario_gen.py @@ -35,6 +35,7 @@ import tomllib from importlib.metadata import version from pathlib import Path +from typing import TYPE_CHECKING import numpy as np @@ -44,11 +45,14 @@ from shinro.utils.config_resolver import resolve_config_path from shinro.utils.linearization import derive_model +if TYPE_CHECKING: + from shinro.codegen.compose import ComposedGraph + EXIT_OK = 0 EXIT_UNTRACEABLE = 1 EXIT_USAGE = 2 -_COMPILE_KEYS = {"n_x", "n_u", "optimize", "target", "solver_dir", "oracle_tol"} +_COMPILE_KEYS = {"n_x", "n_u", "optimize", "target", "solver_dir", "oracle_tol", "artifact_name"} _ALLOWED_OPTIMIZE = {"debug", "release"} @@ -58,14 +62,24 @@ def _sha256(path: str) -> str: return hashlib.sha256(f.read()).hexdigest() +def _sha256_raw(path: str) -> str: + """Return the sha256 of a file by literal path (no config-dir resolution). + + Used for the ``.onnx`` weights, which live wherever the controller config + points and are not one of the resolved config locations. + """ + with open(path, "rb") as f: + return hashlib.sha256(f.read()).hexdigest() + + def _validate_compile(compile_cfg: dict | None, scenario_path: str) -> dict: """Parse and strictly validate the ``[compile]`` section. Returns a dict with ``n_x`` / ``n_u`` (optional — derived from ``[plant]`` - when absent) and ``optimize`` / ``target`` / ``solver_dir`` (optional, with - defaults). Unknown keys and invalid ``optimize`` values are loud errors — - the section is the build spec, so a typo must not silently change the - build. + when absent) and ``optimize`` / ``target`` / ``solver_dir`` / + ``artifact_name`` (optional, with defaults). Unknown keys and invalid + ``optimize`` values are loud errors — the section is the build spec, so a + typo must not silently change the build. Raises: ValueError: On a missing section, unknown keys, or an invalid @@ -83,12 +97,21 @@ def _validate_compile(compile_cfg: dict | None, scenario_path: str) -> dict: f"(got '{optimize}'). ReleaseSafe hangs in osqp_solve (Zig integration " f"bug) and ReleaseSmall is unvalidated — only ReleaseFast is shippable." ) + # The artifact stem becomes lib/.so + lib/.{manifest,deployment}.json; + # a stray separator or space would silently write outside the prefix. + artifact_name = compile_cfg.get("artifact_name", "libbase") + if not isinstance(artifact_name, str) or not artifact_name or any(c in artifact_name for c in "/\\ \t"): + raise ValueError( + f"{scenario_path}: [compile].artifact_name must be a simple file-name stem " + f"like 'libbase' or 'lib_neural_network' (got {artifact_name!r})" + ) return { "n_x": int(compile_cfg["n_x"]) if "n_x" in compile_cfg else None, "n_u": int(compile_cfg["n_u"]) if "n_u" in compile_cfg else None, "optimize": optimize, "target": compile_cfg.get("target", "native"), "solver_dir": compile_cfg.get("solver_dir"), + "artifact_name": artifact_name, # Optional oracle-B override for QP graphs whose realistic solution # settles more coarsely than the tier default at this problem size # (two independent OSQP runs — C-baked vs Python — agree to solver @@ -107,46 +130,60 @@ def load_scenario(scenario_path: str) -> dict: ``n_x``/``n_u`` and the ``A_dynamics``/``B_dynamics`` model from it. Raises: - ValueError: On a missing ``[controller]``/``[estimator]`` section or an - invalid ``[compile]`` section. + ValueError: On a missing ``[controller]`` section, an invalid + ``[compile]`` section, or a missing ``[estimator]`` on anything + other than a policy-only (``onnx_rl``) scenario. """ with open(resolve_config_path(scenario_path), "rb") as f: cfg = tomllib.load(f) - if "controller" not in cfg or "estimator" not in cfg: - raise ValueError( - f"{scenario_path}: scenario requires [controller] and [estimator] sections" - ) + if "controller" not in cfg: + raise ValueError(f"{scenario_path}: scenario requires a [controller] section") + controller = cfg["controller"] + estimator = cfg.get("estimator") + if estimator is None: + # A policy-only scenario: the ONNX policy is a standalone graph with no + # estimator to compose with. Any other controller needs [estimator]. + ctype = controller.get("type") or _type_from_component_config(controller["config"]) + if ctype != "onnx_rl": + raise ValueError( + f"{scenario_path}: missing [estimator] section — only a policy-only 'onnx_rl' " + f"scenario may omit it (got controller type {ctype!r})" + ) limits = None il = cfg.get("scenario", {}).get("input_limits") if il: limits = (np.array(il["min"], dtype=np.float64), np.array(il["max"], dtype=np.float64)) return { - "estimator_config": cfg["estimator"]["config"], - "controller_config": cfg["controller"]["config"], - "estimator_type": cfg["estimator"].get("type"), - "controller_type": cfg["controller"].get("type"), + "estimator_config": estimator["config"] if estimator else None, + "controller_config": controller["config"], + "estimator_type": estimator.get("type") if estimator else None, + "controller_type": controller.get("type"), "input_limits": limits, "compile": _validate_compile(cfg.get("compile"), scenario_path), "plant": cfg.get("plant"), } -def _provenance(scenario_path: str, spec: dict) -> dict: +def _provenance(scenario_path: str, spec: dict, model_path: str | None = None) -> dict: """Build the provenance dict for ``lower_zig``. Records the sha256 of the scenario TOML itself (pinning the whole build - spec, including ``[compile]``) plus the estimator/controller configs — and - the plant config when derivation was used, so the deployment record's - config slot commits to every file the artifact was built from. + spec, including ``[compile]``) plus the estimator/controller configs — the + plant config when derivation was used, and the ``.onnx`` weights for a + policy-only scenario — so the deployment record's config slot commits to + every file the artifact was built from. """ configs = { scenario_path: _sha256(scenario_path), - spec["estimator_config"]: _sha256(spec["estimator_config"]), spec["controller_config"]: _sha256(spec["controller_config"]), } + if spec["estimator_config"]: + configs[spec["estimator_config"]] = _sha256(spec["estimator_config"]) plant = spec.get("plant") if plant and "config" in plant: configs[plant["config"]] = _sha256(plant["config"]) + if model_path: + configs[model_path] = _sha256_raw(model_path) return { "configs": configs, "python_version": sys.version.split()[0], @@ -154,6 +191,53 @@ def _provenance(scenario_path: str, spec: dict) -> dict: } +def _policy_graph(spec: dict, scenario_path: str) -> ComposedGraph: + """Build the composed graph for a policy-only scenario (no estimator). + + The controller config is read directly rather than through the factory — + the importer wants the raw ONNX/observation tables, and the policy graph is + standalone (nothing to compose it with). The ``.onnx`` path is recorded on + ``spec`` so the provenance can pin the weights. + + Args: + spec: The loaded scenario spec (``estimator_config`` is None). + scenario_path: Scenario TOML path, for error messages. + + Returns: + The imported, memoryless :class:`ComposedGraph`. + + Raises: + ValueError: If the controller config has no ``model_path``. + """ + from shinro.codegen.onnx_import import import_onnx_policy + + with open(resolve_config_path(spec["controller_config"]), "rb") as f: + ctrl = tomllib.load(f) + model_path = ctrl.get("model_path") + if model_path is None: + raise ValueError( + f"{scenario_path}: a policy-only scenario needs the controller's model_path " + f"(artifact_dir is the *result* of this compile, not an input)" + ) + action_cfg = { + "action_space": ctrl.get("action_space", "continuous"), + "deterministic": ctrl.get("deterministic", True), + "action_scale": ctrl.get("action_scale", 1.0), + "action_bias": ctrl.get("action_bias", 0.0), + } + for key in ("action_clip_low", "action_clip_high"): + if key in ctrl: + action_cfg[key] = ctrl[key] + spec["policy_model_path"] = model_path + return import_onnx_policy( + model_path, + n_x=spec["compile"]["n_x"], + obs_cfg=ctrl.get("observation", {}), + action_cfg=action_cfg, + output_name=ctrl.get("output_name"), + ) + + def _type_from_component_config(config_path: str) -> str: """Read the component's registered name from its own config TOML. @@ -187,6 +271,22 @@ def gen_scenario(scenario_path: str, out_dir: str) -> tuple: NotImplementedError: If a component uses an untraceable op. """ spec = load_scenario(scenario_path) + + # A policy-only scenario (onnx_rl) is standalone: no estimator to compose + # with, so it bypasses the plant-derivation + build_composed_graph path and + # goes straight from the ONNX graph to the lowered table. + if spec["estimator_config"] is None: + cg = _policy_graph(spec, scenario_path) + out = Path(out_dir) + out.mkdir(parents=True, exist_ok=True) + graph_path = out / "graph_data.zig" + lower_zig( + cg, + str(graph_path), + provenance=_provenance(scenario_path, spec, model_path=spec["policy_model_path"]), + ) + return cg, graph_path + n_x = spec["compile"]["n_x"] n_u = spec["compile"]["n_u"] est_cfg = spec["estimator_config"] diff --git a/src/shinro/codegen/stamp.py b/src/shinro/codegen/stamp.py index 1bf2232..69113fd 100644 --- a/src/shinro/codegen/stamp.py +++ b/src/shinro/codegen/stamp.py @@ -1,9 +1,9 @@ """Stamp a deployment record for a built ``libbase.so``. -Post-compile step: after ``zig build`` produces ``/lib/libbase.so``, -this module reads the build manifest (``libbase.manifest.json``), hashes the +Post-compile step: after ``zig build`` produces ``/lib/lib.so``, +this module reads the build manifest (``lib.manifest.json``), hashes the binary and the baked solver tree, and writes a deterministic deployment -record (``libbase.deployment.json``) carrying a single **master hash** that +record (``lib.deployment.json``) carrying a single **master hash** that commits to the whole config -> graph -> solver -> binary chain. The master hash is a pure function of its inputs (no timestamps in the @@ -62,7 +62,7 @@ def _master(config_slot: str, graph_slot: str, solver_slot: str, binary_slot: st ) -def stamp(prefix: Path, build_root: Path) -> dict: +def stamp(prefix: Path, build_root: Path, name: str = "libbase") -> dict: """Compute and write the deployment record for a built prefix dir. Args: @@ -70,18 +70,20 @@ def stamp(prefix: Path, build_root: Path) -> dict: build_root: The zig build root, used to resolve a relative ``solver_dir`` recorded in the manifest (default: the packaged ``src/shinro/runtime/``). + name: Artifact stem — ``.so`` / ``.manifest.json`` + (default ``libbase``). Must match the ``-Dname`` the build used. Returns: The deployment record dict (also written to disk). """ lib_dir = prefix / "lib" - manifest_path = lib_dir / "libbase.manifest.json" - so_path = lib_dir / "libbase.so" + manifest_path = lib_dir / f"{name}.manifest.json" + so_path = lib_dir / f"{name}.so" if not manifest_path.exists(): raise FileNotFoundError(f"no build manifest at {manifest_path}") if not so_path.exists(): - raise FileNotFoundError(f"no libbase.so at {so_path}") + raise FileNotFoundError(f"no {so_path.name} at {so_path}") manifest = json.loads(manifest_path.read_text()) @@ -124,7 +126,7 @@ def stamp(prefix: Path, build_root: Path) -> dict: "binary": {"sha256": binary_slot, "path": str(so_path)}, } - record_path = lib_dir / "libbase.deployment.json" + record_path = lib_dir / f"{name}.deployment.json" record_path.write_text(json.dumps(record, indent=2, sort_keys=True) + "\n") # Archive copy: timestamp in filename only, so the record stays a pure @@ -141,15 +143,16 @@ def stamp(prefix: Path, build_root: Path) -> dict: def main() -> None: - parser = argparse.ArgumentParser(description="Stamp a deployment record for a built libbase.so.") + parser = argparse.ArgumentParser(description="Stamp a deployment record for a built kernel (lib.so).") parser.add_argument("--prefix", default="build", help="zig build prefix dir (default: build)") parser.add_argument( "--build-root", default=str(runtime_root()), help="zig build root for resolving a relative solver_dir (default: the packaged runtime)", ) + parser.add_argument("--name", default="libbase", help="artifact stem, e.g. libbase → libbase.so (default: libbase)") args = parser.parse_args() - stamp(Path(args.prefix), Path(args.build_root)) + stamp(Path(args.prefix), Path(args.build_root), args.name) if __name__ == "__main__": diff --git a/src/shinro/configs/controllers/onnx_rl.toml b/src/shinro/configs/controllers/onnx_rl.toml index 0b34a90..c3ebd0e 100644 --- a/src/shinro/configs/controllers/onnx_rl.toml +++ b/src/shinro/configs/controllers/onnx_rl.toml @@ -1,29 +1,49 @@ # FILE: configs/controllers/onnx_rl.toml -# ONNX RL policy adapter — wraps any ONNX-exported RL policy as a Controller. -# Requires: pip install onnxruntime +# ONNX RL policy adapter — run any ONNX-exported RL policy as a Controller. # -# Export the actor network from your RL framework (sb3, RLlib, CleanRL, -# custom torch, JAX, ...) to policy.onnx, then point model_path at it. +# Requires: pip install "shinro[onnx-rl]" (the `onnx` package, compile-time only; +# there is no onnxruntime dependency — the deployed kernel has no dependencies). # -# Usage: -# python -m demos.demo_base_tracking --controller onnx_rl +# Export the actor network from any RL stack (sb3, RLlib, CleanRL, custom torch, +# JAX, ...) to policy.onnx, then pick one of two interchangeable backends: +# +# eager — model_path points at the .onnx; the graph is imported once and run +# in-process by the numpy interpreter (no build step, no artifact). +# compiled — artifact_dir points at a `make compile --out` directory; the +# kernel `lib/.so` is dlopen'd and driven through the +# C ABI (the deployment path; the .onnx is not needed then). +# +# Both backends run the same graph, so they agree bit-for-bit. The observation +# encoder and the action space are baked into that graph, so one policy = one +# artifact: retraining means re-exporting and re-compiling. +# +# Compile a policy-only scenario (no [estimator] — only onnx_rl may omit it): +# make compile SCENARIO=tests/fixtures/configs/scenarios/toy_mlp_policy.toml OUT=build/policy +# then set artifact_dir = "build/policy" above. +# +# Sampling action spaces draw noise on the host and feed it to the kernel's +# epsilon port (Gumbel for discrete, standard normal for stochastic); `seed` +# seeds that generator, and reset() reseeds it. type = "onnx_rl" name = "ppo_policy" -model_path = "path/to/policy.onnx" +model_path = "path/to/policy.onnx" # eager backend (import + interpret) +# artifact_dir = "build/policy" # compiled backend; wins if both are set +# n_x = 6 # plant state dim when state_keys stop short of the last entry +# output_name = "action" # ONNX output to import (default: the model's first) action_space = "continuous" # continuous | discrete | stochastic -deterministic = true # false samples from the policy distribution -action_scale = 1.0 # u = scale * a + bias (tanh-squashed policies) +deterministic = true # false samples (discrete: Gumbel-max; stochastic: Gaussian) +action_scale = 1.0 # u = scale * u + bias (continuous / stochastic only) action_bias = 0.0 -# action_clip_low = -1.0 -# action_clip_high = 1.0 +# action_clip_low = -1.0 # both bounds are required together — a single +# action_clip_high = 1.0 # bound would be ±inf, which the lowerer cannot emit seed = 0 [observation] # input_name = "obs" # override ONNX input name if needed state_keys = [0, 1, 2] # integer indices into the plant state vector -normalize = false # apply (x - mean) / std before the model +normalize = false # apply (x - mean) / std before the model # obs_mean = [0.0, 0.0, 0.0] # obs_std = [1.0, 1.0, 1.0] -# clip = [-1.0, 1.0] # observation clipping -add_batch_dim = true # ONNX models expect [B, obs_dim] +# clip = [-1.0, 1.0] # observation clipping +# add_batch_dim = true # legacy: the compiled path treats the batch as 1 implicitly diff --git a/src/shinro/configs/scenarios/_template.toml b/src/shinro/configs/scenarios/_template.toml index cc9fc1b..511166a 100644 --- a/src/shinro/configs/scenarios/_template.toml +++ b/src/shinro/configs/scenarios/_template.toml @@ -43,6 +43,7 @@ optimize = "debug" # "debug" | "release" (release → -Doptimize=ReleaseF # target = "aarch64-linux-gnu" # cross-compile for the robot board; omit = native # solver_dir = "src/shinro/runtime/codegen/emosqp" # REQUIRED for MPC controllers only # oracle_tol = 1e-3 # override the numeric oracle tolerance (QP graphs default 1e-3) +# artifact_name = "libbase" # deployed artifact stem: lib/.so (default libbase → libbase.so) # ── SIMULATION-ONLY — delete this block for a compile-only robot ───────────── # [physics] diff --git a/src/shinro/controllers/onnx_rl_adapter.py b/src/shinro/controllers/onnx_rl_adapter.py index 5f20c81..79eca4e 100644 --- a/src/shinro/controllers/onnx_rl_adapter.py +++ b/src/shinro/controllers/onnx_rl_adapter.py @@ -1,37 +1,50 @@ -"""ONNX RL policy adapter — wraps any ONNX-exported reinforcement-learning policy as a Controller. - -Allows swapping between classical control (LQR, MPC) and policies trained in -*any* external RL stack (Stable-Baselines3, RLlib, CleanRL, custom PyTorch, -JAX/Flax, ...) provided the actor network is exported to ONNX. Inference runs -through ``onnxruntime`` — no torch / gym / framework dependency at deploy time. - -The adapter supports three action-space conventions: - -- ``continuous``: output is a real-valued action vector, optionally scaled and - biased for tanh-squashed policies (``u = scale * tanh(a) + bias``). -- ``discrete``: output is a logits vector; greedy ``argmax`` by default, or - sampled from ``softmax`` when ``deterministic = false``. -- ``stochastic``: output is ``[mean; log_std]``; the mean is used when - ``deterministic = true``, otherwise a Gaussian sample is drawn with the - configured seed. - -Observations are built from the flat plant state via integer index selection -plus optional per-dimension normalization and clipping. The adapter is -backend-agnostic: state may arrive as a numpy array or torch tensor, and the -action is returned in the same backend's native type (conversion happens at -the ONNX boundary, which requires numpy feeds). - -Usage: - # In configs/controllers/onnx_rl.toml: +"""ONNX RL policy adapter — run an ONNX-exported policy as a Controller. + +``onnxruntime`` is gone. The policy's ONNX graph is translated into a shinro +graph by :mod:`shinro.codegen.onnx_import` — observation encoding, the network, +and the action post-processing all become arithmetic on baked constants — and +that graph is executed through one of two interchangeable backends: + +- **eager** (``model_path``): the imported graph is run in-process by + :func:`shinro.codegen.interpreter.interpret` (pure numpy, f64). No build step + and no compiled artifact are needed; this is the testing/reference path. +- **compiled** (``artifact_dir``): the scenario's compiled kernel + (``lib/lib_neural_network.so``) is dlopen'd and driven through the + ``shinro_step`` C ABI, with the graph manifest next to it describing the port + layout. This is the deployment path: no Python array framework, no ONNX + runtime, no dependencies at all. + +Both backends see the same graph, so they agree bit-for-bit (the compile gate +checks exactly that). The only input port is the raw plant state; a sampling +action space adds an ``epsilon`` port that the host fills with noise each tick +(Gumbel for ``discrete``, standard normal for ``stochastic``) — the kernel does +the arithmetic, RNG stays on the host, matching MPPI's contract. + +Action spaces (baked at import time, mirroring the historical runtime): + +- ``continuous``: ``u = scale * a + bias`` (optionally clipped). +- ``discrete``: argmax one-hot (deterministic) or Gumbel-max sampling. +- ``stochastic``: ``[mean; log_std]`` — the mean, or + ``mean + exp(clip(log_std, -10, 2)) * epsilon``, then ``scale * u + bias``. + +The host-side noise is drawn from a seeded generator; :meth:`reset` reseeds it, +so a run is reproducible. + +Usage (configs/controllers/onnx_rl.toml):: + # type = "onnx_rl" - # model_path = "path/to/policy.onnx" + # model_path = "path/to/policy.onnx" # eager mode + # # artifact_dir = "build/my_policy" # compiled mode (make compile --out) # action_space = "continuous" - # - # python -m demos.demo_base_tracking --controller onnx_rl """ from __future__ import annotations +import ctypes +import json +import math +from dataclasses import dataclass, field +from pathlib import Path from typing import Any import numpy as np @@ -40,242 +53,300 @@ from shinro.factories.registry import register_controller from shinro.utils.array_backend import ArrayBackend, NumpyBackend +#: Clamp for the ``U(0, 1)`` draws feeding the Gumbel transform, so a draw of +#: exactly 0 cannot produce ``-inf`` noise. +_GUMBEL_EPS = 1e-12 + +#: Filename the compiled policy kernel must be installed as under +#: ``/lib/`` (the ``make compile`` output for this scenario). +KERNEL_FILENAME = "lib_neural_network.so" + -class _ObsEncoder: - """Config-driven observation encoder: plant state -> ONNX feed dict. +@dataclass(frozen=True) +class OnnxRLConfig: + """Strict TOML schema for :class:`OnnxRLAdapter`. - Accepts any backend-native state (numpy array or torch tensor), converts - it to numpy at the ONNX boundary via ``bk.to_numpy``, subselects integer - indices, optionally applies per-dimension normalization - ``(x - mean) / std`` and clipping, then packs the result into a - batch-ready float32 array for a single ONNX input. + Exactly one of ``model_path`` / ``artifact_dir`` must be given: the former + imports and interprets the ONNX graph in-process, the latter loads an + already-compiled kernel. ``artifact_dir`` wins if both are set, so a config + can keep the model path for provenance while deploying the ``.so``. + + The observation sub-table is left as a plain dict because its keys map + straight onto the importer's ``obs_cfg`` (which validates them); every other + field is the same action-space surface the old adapter exposed. """ - def __init__( - self, - input_name: str, - state_keys: list[int], - obs_mean: np.ndarray | None = None, - obs_std: np.ndarray | None = None, - clip: tuple[float, float] | None = None, - add_batch_dim: bool = True, - backend: ArrayBackend | None = None, - ) -> None: - self.input_name = input_name - self.state_keys = np.asarray(state_keys, dtype=int) - self.obs_mean = obs_mean - self.obs_std = obs_std - self.clip = clip - self.add_batch_dim = add_batch_dim - self.bk = backend or NumpyBackend() + model_path: str | None = None + artifact_dir: str | None = None + n_x: int | None = None + output_name: str | None = None + action_space: str = "continuous" + deterministic: bool = True + action_scale: Any = 1.0 + action_bias: Any = 0.0 + action_clip_low: float | None = None + action_clip_high: float | None = None + seed: int = 0 + observation: dict[str, Any] = field(default_factory=dict) + name: str = "onnx_rl" + + +class _GraphPolicy: + """Eager artifact: an imported shinro graph executed by the interpreter.""" + + def __init__(self, cg) -> None: + """Wrap a composed graph as a runnable policy. + + Args: + cg: The :class:`~shinro.codegen.compose.ComposedGraph` returned by + :func:`shinro.codegen.onnx_import.import_onnx_policy`. + """ + from shinro.codegen.onnx_import import EPSILON_PORT, STATE_PORT + + self._cg = cg + self.inputs = list(cg.inputs) + self.ops = frozenset(node.op for node in cg.graph.nodes) + self._u_port = cg.outputs[0] + self.state_port = STATE_PORT + self.state_size = _graph_port_size(cg.graph, STATE_PORT) + self.noise_port = EPSILON_PORT if EPSILON_PORT in self.inputs else None + self.noise_size = _graph_port_size(cg.graph, EPSILON_PORT) if self.noise_port else 0 + + @property + def gumbel(self) -> bool: + """True when the graph expects Gumbel noise (i.e. it samples discretely).""" + return "one_hot" in self.ops - def encode(self, state: Any) -> dict[str, np.ndarray]: - """Convert a backend-native plant state into an ``{input_name: tensor}`` feed dict.""" - s = self.bk.to_numpy(state) - obs = np.asarray(s, dtype=np.float32)[self.state_keys].astype(np.float32, copy=True) - if self.obs_mean is not None: - obs = obs - self.obs_mean - if self.obs_std is not None: - obs = obs / self.obs_std - if self.clip is not None: - obs = np.clip(obs, *self.clip) - if self.add_batch_dim: - obs = obs[None, :] - return {self.input_name: obs} + def step(self, feed: dict[str, np.ndarray]) -> np.ndarray: + """Run one tick through the interpreter.""" + from shinro.codegen.interpreter import interpret + + return interpret(self._cg.graph, feed)[self._u_port] + + +class _CompiledPolicy: + """Deployment artifact: a compiled policy kernel plus its graph manifest. + + The manifest (written by :func:`shinro.codegen.lower_zig.lower_zig` next to + the graph) is the artifact's self-description, so the loader reads the port + order, shapes, and op histogram from it rather than from the original ONNX + model — the ``.onnx`` file is not needed at deploy time. + """ + + def __init__(self, artifact_dir: str | Path) -> None: + """Load ``/lib/lib_neural_network.so`` and its manifest. + + Args: + artifact_dir: A ``make compile --out`` directory (contains + ``graph_data_manifest.json`` and ``lib/lib_neural_network.so``). + + Raises: + FileNotFoundError: If the manifest or the shared object is missing. + ValueError: If the artifact does not expose the expected ``state`` + input and ``u`` output ports. + """ + from shinro.codegen.onnx_import import EPSILON_PORT, OUTPUT_PORT, STATE_PORT + + root = Path(artifact_dir) + manifest_path = root / "graph_data_manifest.json" + so_path = root / "lib" / KERNEL_FILENAME + if not manifest_path.exists(): + raise FileNotFoundError(f"no graph manifest at {manifest_path} — run `make compile --out {root}` first") + + try: + self.manifest = json.loads(manifest_path.read_text()) + except (OSError, json.JSONDecodeError) as exc: + raise ValueError(f"graph manifest at {manifest_path} is unreadable or corrupt: {exc}") from exc + self.inputs = [port["name"] for port in self.manifest["inputs"]] + self.ops = frozenset(self.manifest["op_histogram"]) + if STATE_PORT not in self.inputs: + raise ValueError(f"artifact {root} has no '{STATE_PORT}' input port (inputs: {self.inputs})") + + out_names = [port["name"] for port in self.manifest["outputs"]] + if OUTPUT_PORT not in out_names: + raise ValueError(f"artifact {root} has no '{OUTPUT_PORT}' output port (outputs: {out_names})") + + self._in_sizes = [_flat_size(port["shape"]) for port in self.manifest["inputs"]] + out_sizes = [_flat_size(port["shape"]) for port in self.manifest["outputs"]] + self._n_out = sum(out_sizes) + self._n_state = sum(_flat_size(port["shape"]) for port in self.manifest["state_outputs"]) + u_index = out_names.index(OUTPUT_PORT) + u_start = sum(out_sizes[:u_index]) + self._u_slice = (u_start, u_start + out_sizes[u_index]) + + self.state_port = STATE_PORT + self.state_size = _flat_size(self.manifest["inputs"][self.inputs.index(STATE_PORT)]["shape"]) + self.noise_port = EPSILON_PORT if EPSILON_PORT in self.inputs else None + self.noise_size = _flat_size(self.manifest["inputs"][self.inputs.index(EPSILON_PORT)]["shape"]) if self.noise_port else 0 + + if not so_path.exists(): + raise FileNotFoundError(f"no compiled kernel at {so_path} — run `make compile --out {root}` first") + lib = ctypes.CDLL(str(so_path)) + lib.shinro_step.argtypes = [ctypes.POINTER(ctypes.c_double)] * 3 + lib.shinro_step.restype = None + self._lib = lib + + @property + def gumbel(self) -> bool: + """True when the compiled graph expects Gumbel noise (it samples discretely).""" + return "one_hot" in self.ops + + def step(self, feed: dict[str, np.ndarray]) -> np.ndarray: + """Pack the ports, call ``shinro_step``, and return the ``u`` slice.""" + packed = np.concatenate([np.asarray(feed[name], dtype=np.float64).ravel() for name in self.inputs]) + out = np.zeros(self._n_out, dtype=np.float64) + # A memoryless policy declares no state outputs; the C ABI still wants a + # non-null pointer, so give it a one-element scratch buffer. + state = np.zeros(max(self._n_state, 1), dtype=np.float64) + ptr = ctypes.POINTER(ctypes.c_double) + self._lib.shinro_step( + packed.ctypes.data_as(ptr), + out.ctypes.data_as(ptr), + state.ctypes.data_as(ptr), + ) + start, stop = self._u_slice + return out[start:stop].copy() @register_controller("onnx_rl") class OnnxRLAdapter(Controller): - """Wrap an ONNX-exported RL policy as a Controller. + """Run an ONNX-exported RL policy as a Controller. - The policy is loaded from a local ``.onnx`` file via ``onnxruntime``. - ``compute()`` encodes the plant state (normalization / clipping / index - selection), runs the model, post-processes the raw output into a control - action, and returns it as a numpy array. + The policy is prepared once (imported + interpreted, or a compiled ``.so`` + is loaded) and each :meth:`compute` call runs one tick on the plant state. Args: - session: Loaded ``onnxruntime.InferenceSession``. - obs_encoder: Encoder mapping plant state to the ONNX feed dict. - output_name: ONNX output tensor name. - action_space: ``"continuous"``, ``"discrete"``, or ``"stochastic"``. - deterministic: For discrete/stochastic policies, return the greedy - argmax / mean instead of sampling (default: ``True``). - action_scale: Post-policy per-action scaling (tanh-squash support). - action_bias: Post-policy per-action bias. - action_clip: Optional ``(low, high)`` tuple to clip the final action. - seed: RNG seed for sampling action spaces. - backend: Array backend for state input and action output. ONNX - inference itself always runs on numpy arrays, but the adapter - converts the backend-native state to numpy at the boundary and - the resulting action back to the backend's native type. + policy: A loaded policy artifact — :class:`_GraphPolicy` (eager) or + :class:`_CompiledPolicy` (compiled). Built by :meth:`from_config`. + seed: RNG seed for action sampling. Sampling action spaces (discrete + non-deterministic, stochastic non-deterministic) draw their noise + from a generator seeded here. + backend: Array backend for the state input and action output. The + kernel itself always works on numpy f64; the adapter converts at the + boundary, so a torch state yields a torch action. """ - def __init__( - self, - weights, - obs_encoder: _ObsEncoder, - output_name: str, - action_space: str = "continuous", - deterministic: bool = True, - action_scale: float | np.ndarray = 1.0, - action_bias: float | np.ndarray = 0.0, - action_clip: tuple[float, float] | None = None, - seed: int = 0, - backend: ArrayBackend | None = None, - ) -> None: - if action_space not in ("continuous", "discrete", "stochastic"): - raise ValueError(f"action_space must be continuous/discrete/stochastic, got {action_space!r}") - self.session = weights - self.obs_encoder = obs_encoder - self.output_name = output_name - self.action_space = action_space - self.deterministic = deterministic - self.action_scale = np.asarray(action_scale, dtype=np.float32) - self.action_bias = np.asarray(action_bias, dtype=np.float32) - self.action_clip = action_clip + Config = OnnxRLConfig + + def __init__(self, policy: _GraphPolicy | _CompiledPolicy, *, seed: int = 0, backend: ArrayBackend | None = None) -> None: + self.policy = policy self.seed = seed self.bk = backend or NumpyBackend() - self._rng = np.random.default_rng(seed) - - def _postprocess(self, raw: np.ndarray) -> np.ndarray: - """Turn raw network output into a control action.""" - raw = raw.astype(np.float32, copy=False) - - if self.action_space == "continuous": - action = raw * self.action_scale + self.action_bias - elif self.action_space == "discrete": - logits = raw.reshape(-1) - if self.deterministic: - action = np.zeros_like(logits) - action[np.argmax(logits)] = 1.0 - else: - probs = np.exp(logits - np.max(logits)) - probs = probs / probs.sum() - idx = self._rng.choice(len(logits), p=probs) - action = np.zeros_like(logits) - action[idx] = 1.0 - else: # stochastic: raw = [mean; log_std] - half = raw.size // 2 - mean = raw[:half] - if self.deterministic: - u = mean - else: - log_std = np.clip(raw[half:], -10.0, 2.0) - u = mean + np.exp(log_std) * self._rng.standard_normal(mean.size) - u = u * self.action_scale + self.action_bias - if self.action_clip is not None: - u = np.clip(u, *self.action_clip) - return u.astype(np.float32, copy=False) - - if self.action_clip is not None: - action = np.clip(action, *self.action_clip) - return action + self._rng = np.random.default_rng(self.seed) def compute(self, state, target=None): - """Run the ONNX policy on the current state. + """Run the policy on the current plant state. Args: - state: Plant state vector in the configured backend's native - type (numpy array or torch tensor). - target: Ignored for learned policies — they generate actions - from observation alone. + state: Plant state vector in the configured backend's native type + (numpy array, torch tensor, or a sequence). Must have the + compiled graph's ``state`` port length. + target: Ignored — learned policies act on the observation alone. Returns: - Action vector (n_u,) in the backend-native type. + The action vector (n_u,) in the backend's native type. """ - feed = self.obs_encoder.encode(state) - raw = self.session.run([self.output_name], feed)[0] - action = self._postprocess(raw).reshape(-1) - return self.bk.from_numpy(action) + x = np.asarray(self.bk.to_numpy(state), dtype=np.float64).ravel() + if x.size != self.policy.state_size: + raise ValueError(f"onnx_rl: expected a state of {self.policy.state_size} elements, got {x.size}") + feed = {self.policy.state_port: x} + if self.policy.noise_port is not None: + feed[self.policy.noise_port] = self._draw_noise(self.policy.noise_size) + return self.bk.from_numpy(self.policy.step(feed)) + + def _draw_noise(self, size: int) -> np.ndarray: + """Draw the host noise the epsilon port expects. + + For a discretely-sampling graph this is Gumbel noise, which makes + ``argmax(logits + g)`` an exact categorical draw from + ``softmax(logits)``; otherwise it is a standard normal. + """ + if self.policy.gumbel: + uniform = np.clip(self._rng.uniform(0.0, 1.0, size=size), _GUMBEL_EPS, 1.0 - _GUMBEL_EPS) + return -np.log(-np.log(uniform)) + return self._rng.standard_normal(size) def reset(self): - """Reset the policy RNG to the configured seed.""" + """Reseed the action-sampling RNG (a fresh run is reproducible).""" self._rng = np.random.default_rng(self.seed) @classmethod def from_config(cls, config, backend: ArrayBackend | None = None): - """Create an OnnxRLAdapter from a TOML config dict. + """Create an OnnxRLAdapter from a TOML config dict or :class:`OnnxRLConfig`. Config fields: - model_path: Path to the ``.onnx`` model file (required). + model_path: Path to the ``.onnx`` model (eager mode). + artifact_dir: A ``make compile --out`` directory (compiled mode). action_space: ``"continuous"``, ``"discrete"``, or ``"stochastic"`` - (default: ``"continuous"``). - deterministic: Whether to return argmax/mean instead of sampling - (default: ``true``). - action_scale: Post-policy action scale (default: 1.0). - action_bias: Post-policy action bias (default: 0.0). - action_clip_low / action_clip_high: Clip the final action - (default: no clipping). - seed: RNG seed for stochastic sampling (default: 0). - - ``[observation]`` subtable fields: - input_name: ONNX input tensor name (defaults to the model's first - input). - state_keys: Integer indices into the plant state to use as - observations (default: ``[0, 1, ..., n-1]``). - normalize: Apply mean/std normalization (default: false). - obs_mean / obs_std: Arrays for normalization. - clip: ``[low, high]`` observation clipping (default: none). - add_batch_dim: Prepend a batch axis (default: true). + (default: continuous). + deterministic: Return argmax/mean instead of sampling (default true). + action_scale / action_bias: Post-policy affine transform. + action_clip_low / action_clip_high: Clip the final action; both are + required together (a missing bound would be ``±inf``, which the + lowerer cannot represent). + seed: RNG seed for sampling action spaces. + n_x: Plant state dimension, when the observation sub-table does not + reach the last state entry. + output_name: ONNX tensor to import as the action. + observation: The importer's observation table (``state_keys``, + ``normalize``, ``obs_mean``, ``obs_std``, ``clip``, ...). Args: - config: TOML config dict. + config: TOML config dict or :class:`OnnxRLConfig`. backend: Array backend for state input and action output. - Defaults to NumpyBackend. Returns: OnnxRLAdapter instance. - """ - import onnxruntime # type: ignore + Raises: + ValueError: On a missing/invalid mode or an invalid action config. + FileNotFoundError: If a compiled artifact is missing. + """ + cfg = cls.parse_config(config) bk = backend or NumpyBackend() - session = onnxruntime.InferenceSession(config["model_path"], providers=["CPUExecutionProvider"]) - output_name = config.get("output_name") - if output_name is None: - output_name = session.get_outputs()[0].name - - obs_cfg = config.get("observation", {}) - input_name = obs_cfg.get("input_name") - if input_name is None: - input_name = session.get_inputs()[0].name - - state_keys = obs_cfg.get("state_keys") - if state_keys is None: - state_keys = list(range(session.get_inputs()[0].shape[1] or 0)) - - obs_mean = None - obs_std = None - if obs_cfg.get("normalize", False): - obs_mean = np.asarray(obs_cfg["obs_mean"], dtype=np.float32) - obs_std = np.asarray(obs_cfg["obs_std"], dtype=np.float32) - - obs_clip = None - if "clip" in obs_cfg: - obs_clip = (float(obs_cfg["clip"][0]), float(obs_cfg["clip"][1])) - - encoder = _ObsEncoder( - input_name=input_name, - state_keys=state_keys, - obs_mean=obs_mean, - obs_std=obs_std, - clip=obs_clip, - add_batch_dim=obs_cfg.get("add_batch_dim", True), - ) - - action_clip = None - if "action_clip_low" in config or "action_clip_high" in config: - action_clip = (float(config.get("action_clip_low", -np.inf)), float(config.get("action_clip_high", np.inf))) - - return cls( - weights=session, - obs_encoder=encoder, - output_name=output_name, - action_space=config.get("action_space", "continuous"), - deterministic=config.get("deterministic", True), - action_scale=config.get("action_scale", 1.0), - action_bias=config.get("action_bias", 0.0), - action_clip=action_clip, - seed=config.get("seed", 0), - backend=bk, - ) + if cfg.artifact_dir is not None: + policy: _GraphPolicy | _CompiledPolicy = _CompiledPolicy(cfg.artifact_dir) + elif cfg.model_path is not None: + from shinro.codegen.onnx_import import import_onnx_policy + + action_cfg: dict[str, Any] = { + "action_space": cfg.action_space, + "deterministic": cfg.deterministic, + "action_scale": cfg.action_scale, + "action_bias": cfg.action_bias, + } + if cfg.action_clip_low is not None: + action_cfg["action_clip_low"] = cfg.action_clip_low + if cfg.action_clip_high is not None: + action_cfg["action_clip_high"] = cfg.action_clip_high + cg = import_onnx_policy( + cfg.model_path, + n_x=cfg.n_x, + obs_cfg=cfg.observation, + action_cfg=action_cfg, + output_name=cfg.output_name, + ) + policy = _GraphPolicy(cg) + else: + raise ValueError("onnx_rl: config needs model_path (eager) or artifact_dir (compiled)") + + return cls(policy, seed=cfg.seed, backend=bk) + + +def _flat_size(shape: Any) -> int: + """Flat element count of a manifest shape (an empty shape is a scalar).""" + dims = list(shape) if shape is not None else [] + return math.prod(dims) if dims else 1 + + +def _graph_port_size(graph, name: str) -> int: + """Flat element count of a named input port of a shinro graph. + + Raises: + KeyError: If the graph declares no such input port. + """ + for node in graph.nodes: + if node.op == "input" and node.attrs["name"] == name: + return _flat_size(node.shape) + raise KeyError(f"input port '{name}' not found in graph") diff --git a/src/shinro/runtime/build.zig b/src/shinro/runtime/build.zig index dc6c620..6e3a1ee 100644 --- a/src/shinro/runtime/build.zig +++ b/src/shinro/runtime/build.zig @@ -18,7 +18,7 @@ pub fn build(b: *std.Build) void { // sources works; tracked as a Zig integration bug), so don't ship it: // zig build -Doptimize=ReleaseFast --build-file runtime/build.zig --prefix build/release/ // Cross targets need no sysroot, e.g. -Dtarget=aarch64-linux-gnu. - // The manifest (libbase.manifest.json) records optimize + strip mode. + // The manifest (lib.manifest.json) records optimize + strip mode. const optimize = b.standardOptimizeOption(.{}); // Build options: which generated graph and which baked OSQP solver to @@ -35,6 +35,14 @@ pub fn build(b: *std.Build) void { "graph", "Path to the generated graph_data.zig (default: graph_data.zig)", ) orelse "graph_data.zig"; + // Deployed artifact stem: `-Dname=lib_neural_network` installs + // lib/lib_neural_network.so (plus its .manifest.json). Defaults to the + // historical "libbase" → libbase.so, so existing builds are unchanged. + const lib_name = b.option( + []const u8, + "name", + "Installed artifact stem: /lib/.so (default: libbase)", + ) orelse "libbase"; const solver_dir_opt = b.option( []const u8, "solver_dir", @@ -127,11 +135,17 @@ pub fn build(b: *std.Build) void { } const lib = b.addLibrary(.{ - .name = "base", + .name = lib_name, .root_module = lib_mod, .linkage = .dynamic, }); - b.installArtifact(lib); + // Install under an explicit sub-path: Zig would otherwise prefix "lib", + // so `-Dname=lib_neural_network` would land as liblib_neural_network.so. + const install_lib = b.addInstallArtifact(lib, .{ + .dest_dir = .{ .override = .lib }, + .dest_sub_path = b.fmt("{s}.so", .{lib_name}), + }); + b.getInstallStep().dependOn(&install_lib.step); // Zig-side unit tests. Only linalg.zig for now; future test files slot // in as additional test modules under runtime/tests/. @@ -169,7 +183,7 @@ pub fn build(b: *std.Build) void { // contains, written next to the artifact after every build, plus a // timestamped archive copy under /manifests/ so teams can browse // which controller combinations were built and when. - writeManifest(b, target, optimize, graph_path, solver_dir); + writeManifest(b, target, optimize, graph_path, solver_dir, lib_name); } // ─── build manifest (audit trail) ───────────────────────────────────────── @@ -182,8 +196,13 @@ fn resolvePath(b: *std.Build, p: []const u8) []const u8 { } /// Read a file at build time; returns "" (with a warning) if unreadable. +/// +/// The cap must clear the generated ``graph_data.zig``, which embeds every +/// baked constant as a hex-float literal (~20 bytes per f64) — a 1M-parameter +/// policy is a ~20 MB source file, so a small cap silently yields "" and the +/// build panics as if the graph were stale. fn readFile(b: *std.Build, path: []const u8) []const u8 { - return std.Io.Dir.cwd().readFileAlloc(b.graph.io, path, b.allocator, .limited(1 << 20)) catch |err| { + return std.Io.Dir.cwd().readFileAlloc(b.graph.io, path, b.allocator, .limited(1 << 28)) catch |err| { std.debug.print("warning: could not read {s}: {s}\n", .{ path, @errorName(err) }); return ""; }; @@ -308,6 +327,7 @@ fn writeManifest( optimize: std.builtin.OptimizeMode, graph_path: []const u8, solver_dir: ?[]const u8, + lib_name: []const u8, ) void { const target_triple = target.result.zigTriple(b.allocator) catch @panic("OOM"); const optimize_name = @tagName(optimize); @@ -362,7 +382,8 @@ fn writeManifest( const cwd = std.Io.Dir.cwd(); const lib_dir = std.fs.path.join(b.allocator, &.{ b.install_prefix, "lib" }) catch @panic("OOM"); cwd.createDirPath(b.graph.io, lib_dir) catch {}; - const report_path = std.fs.path.join(b.allocator, &.{ lib_dir, "libbase.manifest.json" }) catch @panic("OOM"); + const report_name = std.fmt.allocPrint(b.allocator, "{s}.manifest.json", .{lib_name}) catch @panic("OOM"); + const report_path = std.fs.path.join(b.allocator, &.{ lib_dir, report_name }) catch @panic("OOM"); cwd.writeFile(b.graph.io, .{ .sub_path = report_path, .data = json_text }) catch |err| { std.debug.print("warning: could not write manifest {s}: {s}\n", .{ report_path, @errorName(err) }); }; diff --git a/src/shinro/runtime/lower.zig b/src/shinro/runtime/lower.zig index 3b123f5..5d06c4d 100644 --- a/src/shinro/runtime/lower.zig +++ b/src/shinro/runtime/lower.zig @@ -5,10 +5,16 @@ // (op enum, inputs, shapes) is generated by shinro.codegen.lower_zig into // graph_data.zig; the *code* here is handwritten once and never regenerated. // -// The VM is comptime-specialized: `inline for` over the node table unrolls the -// loop, so every node's shapes are comptime constants and every array is a -// fixed-size stack slice of one contiguous buffer. Nothing is allocated, no -// dispatch happens at runtime — the "fixed at compile time" guarantee. +// The VM is comptime-specialized along one axis only: the outer `inline for` +// over the node table unrolls every node, so each node's op and shape are +// comptime constants and its slot is a fixed-size slice of one contiguous +// file-scope buffer — no heap, no runtime op dispatch. The element loops *inside* +// each node are runtime `for` loops over those comptime-known sizes. Unrolling +// them too emits one statement per element (~`buf_len` of them) and forces the +// compiler to optimize a single function of hundreds of thousands of +// instructions: minutes of compile and a `.so` that is ~2x the buffer in pure +// `.text`. Runtime element loops keep compile time and code size proportional +// to the node count, not the element count. // // C ABI (all flat, row-major f64 buffers; layout known to the host from the // ComposedGraph port lists): @@ -30,6 +36,19 @@ const sm = if (g.has_solve_qp) @import("solver_meta") else struct {}; const la = @import("linalg.zig"); const qp = if (g.has_solve_qp) @import("qp.zig") else struct {}; +/// Per-tick workspace: every node's output slot is a slice of `workspace`, +/// addressed by `g.offsets`. Declared at file scope rather than as a local in +/// `shinro_step` so its size — `g.buf_len` f64, tens of MiB for large policies +/// — does not have to fit on the caller's stack (a stack-local copy overflows +/// the default 16 MiB stack past roughly 500k parameters). +/// +/// This is process-global mutable state, so `shinro_step` is **not reentrant +/// or thread-safe**: one call must finish before the next starts. That is the +/// deployment model (one control loop, one tick at a time) and matches the QP +/// path's statically-allocated `solver` global. The graph writes every slot +/// before it is read, so the buffer needs no initialization. +var workspace: [g.buf_len]f64 align(16) = undefined; + /// Run one tick of the closed-loop step through the generated node table. /// /// The C-ABI entry point (exported as `shinro_step`) that the host calls once @@ -42,8 +61,14 @@ const qp = if (g.has_solve_qp) @import("qp.zig") else struct {}; /// - `state_out`: packed in `cg.state_outputs` order (recurrent → next tick) /// /// The `inline for` over `g.nodes` unrolls the whole node table at compile -/// time, so every node's shape is a comptime constant and every array is a -/// fixed-size slice of the single stack buffer — no heap, no runtime dispatch. +/// time, so every node's op and shape are comptime constants and every array +/// is a fixed-size slice of the single file-scope `workspace` — no heap, no +/// runtime op dispatch. The element loops *within* each node are runtime loops over +/// those comptime-known sizes, so code size and compile time scale with the +/// node count rather than with `buf_len`. +/// +/// Not reentrant or thread-safe: writes go to the shared file-scope +/// `workspace` (see its declaration). Callers run one tick at a time. /// /// Args: /// inputs: Flat buffer of this tick's host inputs. @@ -52,101 +77,100 @@ const qp = if (g.has_solve_qp) @import("qp.zig") else struct {}; /// written (fed back as state inputs next tick). export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) void { @setEvalBranchQuota(1_000_000); - var buf: [g.buf_len]f64 align(16) = undefined; inline for (g.nodes, 0..) |node, i| { - const out = buf[g.offsets[i]..][0 .. node.rows * node.cols]; + const out = workspace[g.offsets[i]..][0 .. node.rows * node.cols]; switch (node.op) { .cst => { - inline for (0..node.rows * node.cols) |j| out[j] = g.const_blob[node.aux + j]; + for (0..node.rows * node.cols) |j| out[j] = g.const_blob[node.aux + j]; }, .inp => { - inline for (0..node.rows * node.cols) |j| out[j] = inputs[node.aux + j]; + for (0..node.rows * node.cols) |j| out[j] = inputs[node.aux + j]; }, .out => { - const src = node_input(g.nodes[0..], node, &buf); + const src = node_input(g.nodes[0..], node, &workspace); if (node.aux < g.n_outputs) { - inline for (0..node.rows * node.cols) |j| outputs[g.output_offsets[node.aux] + j] = src[j]; + for (0..node.rows * node.cols) |j| outputs[g.output_offsets[node.aux] + j] = src[j]; } else { - inline for (0..node.rows * node.cols) |j| state_out[g.state_offsets[node.aux - g.n_outputs] + j] = src[j]; + for (0..node.rows * node.cols) |j| state_out[g.state_offsets[node.aux - g.n_outputs] + j] = src[j]; } }, .matmul => { - const a = node_input(g.nodes[0..], node, &buf); - const b = node_input_at(g.nodes[0..], node.inputs[1], &buf); + const a = node_input(g.nodes[0..], node, &workspace); + const b = node_input_at(g.nodes[0..], node.inputs[1], &workspace); const left = g.nodes[node.inputs[0]]; const right = g.nodes[node.inputs[1]]; if (left.vec) { // vecmat: (k,) @ (k, n) -> (n,) — a genuinely 1-D left operand const r = la.vecmat(left.rows, node.rows, a, b); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; } else if (right.vec) { // matvec: (m, k) @ (k,) -> (m,) const r = la.matvec(node.rows, right.rows, a, b); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; } else { // matmul: (m, k) @ (k, n) -> (m, n); also covers (m,1)@(1,n) (k=1) const r = la.matmul(node.rows, left.cols, node.cols, a, b); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; } }, - .add => ew2(g.nodes[0..], node, i, &buf, .add), - .sub => ew2(g.nodes[0..], node, i, &buf, .sub), - .mul => ew2(g.nodes[0..], node, i, &buf, .mul), - .div => ew2(g.nodes[0..], node, i, &buf, .div), - .ne => ew2(g.nodes[0..], node, i, &buf, .ne), - .lt => ew2(g.nodes[0..], node, i, &buf, .lt), - .pow => ew2(g.nodes[0..], node, i, &buf, .pow), + .add => ew2(g.nodes[0..], node, i, &workspace, .add), + .sub => ew2(g.nodes[0..], node, i, &workspace, .sub), + .mul => ew2(g.nodes[0..], node, i, &workspace, .mul), + .div => ew2(g.nodes[0..], node, i, &workspace, .div), + .ne => ew2(g.nodes[0..], node, i, &workspace, .ne), + .lt => ew2(g.nodes[0..], node, i, &workspace, .lt), + .pow => ew2(g.nodes[0..], node, i, &workspace, .pow), .neg => { - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| out[j] = -s[j]; + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| out[j] = -s[j]; }, .abs => { - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| out[j] = @abs(s[j]); + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| out[j] = @abs(s[j]); }, .sign => { // Matches np.sign: -1 / 0 / +1 (0 maps to 0, not +1). - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| { + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| { out[j] = if (s[j] > 0.0) 1.0 else if (s[j] < 0.0) -1.0 else 0.0; } }, .transpose => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); // True 2-D transpose: out (node.rows, node.cols) = src.T, so // out[i][j] = src[j][i] — flat out[i*node.cols + j] = // s[j*src.cols + i]. (The old form baked in the square case's // stride symmetry and silently scrambled non-square inputs.) - inline for (0..node.rows) |oi| { - inline for (0..node.cols) |oj| out[oi * node.cols + oj] = s[oj * g.nodes[node.inputs[0]].cols + oi]; + for (0..node.rows) |oi| { + for (0..node.cols) |oj| out[oi * node.cols + oj] = s[oj * g.nodes[node.inputs[0]].cols + oi]; } }, .inv => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.inv(node.rows, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .reshape => { - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| out[j] = s[j]; + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| out[j] = s[j]; }, .clip => { - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| { + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| { out[j] = std.math.clamp(s[j], g.clip_lo[node.aux + j], g.clip_hi[node.aux + j]); } }, .where_op => { - const cond = node_input(g.nodes[0..], node, &buf); - const a = node_input_at(g.nodes[0..], node.inputs[1], &buf); - const b = node_input_at(g.nodes[0..], node.inputs[2], &buf); + const cond = node_input(g.nodes[0..], node, &workspace); + const a = node_input_at(g.nodes[0..], node.inputs[1], &workspace); + const b = node_input_at(g.nodes[0..], node.inputs[2], &workspace); const cond_n = g.nodes[node.inputs[0]]; const a_n = g.nodes[node.inputs[1]]; const b_n = g.nodes[node.inputs[2]]; - inline for (0..node.rows) |oi| { - inline for (0..node.cols) |oj| { + for (0..node.rows) |oi| { + for (0..node.cols) |oj| { const f = oi * node.cols + oj; const c = cond[if (cond_n.rows * cond_n.cols == 1) 0 else f]; const av = a[bcast_flat(a_n, a_n.vec, node.rows, node.cols, oi, oj)]; @@ -156,9 +180,9 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) } }, .any => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); var found = false; - inline for (0..g.nodes[node.inputs[0]].rows * g.nodes[node.inputs[0]].cols) |j| { + for (0..g.nodes[node.inputs[0]].rows * g.nodes[node.inputs[0]].cols) |j| { if (s[j] != 0.0) found = true; } out[0] = if (found) 1.0 else 0.0; @@ -171,36 +195,36 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // scalar index to a depth-row; copy/slice are flat copies with // slice's `aux` holding the input offset. .copy => { - const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| out[j] = s[j]; + const s = node_input(g.nodes[0..], node, &workspace); + for (0..node.rows * node.cols) |j| out[j] = s[j]; }, .tanh => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.tanh(node.rows * node.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .relu => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.relu(node.rows * node.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .exp => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.elementwise_exponential(node.rows * node.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .sin => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.sin_vec(node.rows * node.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .cos => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const r = la.cos_vec(node.rows * node.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .argmax => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const n_in = g.nodes[node.inputs[0]].rows * g.nodes[node.inputs[0]].cols; const idx = la.argmax(n_in, s); out[0] = @floatFromInt(idx); @@ -210,36 +234,36 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // 1 = axis 0 down columns, 2 = axis 1 across rows). The output // shape was fixed at trace time, so rows*cols is the exact // element count the chosen reduction produces. - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const src = g.nodes[node.inputs[0]]; if (node.aux == 0) { const r = la.min_all(src.rows * src.cols, s); out[0] = r[0]; } else if (node.aux == 1) { const r = la.min_axis0(src.rows, src.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; } else { const r = la.min_axis1(src.rows, src.cols, s); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; } }, .one_hot => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const idx: usize = @intFromFloat(s[0]); const r = la.onehot(node.rows, idx); - inline for (0..node.rows * node.cols) |j| out[j] = r[j]; + for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .slice => { - const s = node_input(g.nodes[0..], node, &buf); + const s = node_input(g.nodes[0..], node, &workspace); const src = g.nodes[node.inputs[0]]; if (src.vec) { // 1-D source: aux is the flat element offset. - inline for (0..node.rows * node.cols) |j| out[j] = s[node.aux + j]; + for (0..node.rows * node.cols) |j| out[j] = s[node.aux + j]; } else { // 2-D source: the interpreter slices ROWS (x[start:stop] // along axis 0), so out[i][j] = src[start + i][j]. - inline for (0..node.rows) |oi| { - inline for (0..node.cols) |oj| out[oi * node.cols + oj] = s[(node.aux + oi) * src.cols + oj]; + for (0..node.rows) |oi| { + for (0..node.cols) |oj| out[oi * node.cols + oj] = s[(node.aux + oi) * src.cols + oj]; } } }, @@ -248,10 +272,10 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // inputs share the same shape (enforced at trace time), so the // output flat length is n_inputs * in_len == node.rows * node.cols. .stack => { - inline for (node.inputs, 0..) |inp_idx, row| { - const src = node_input_at(g.nodes[0..], inp_idx, &buf); + for (node.inputs, 0..) |inp_idx, row| { + const src = node_input_at(g.nodes[0..], inp_idx, &workspace); const in_len = g.nodes[inp_idx].rows * g.nodes[inp_idx].cols; - inline for (0..in_len) |j| out[row * in_len + j] = src[j]; + for (0..in_len) |j| out[row * in_len + j] = src[j]; } }, // .solve_qp — the convergence-iterative MPC op. The problem data @@ -275,7 +299,7 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) ") — rebuild with the matching -Dsolver_dir"); } } - const q_vec = node_input(g.nodes[0..], node, &buf); + const q_vec = node_input(g.nodes[0..], node, &workspace); qp.solve_qp(q_vec, out); }, } @@ -307,8 +331,9 @@ inline fn bcast_flat(op: g.Node, op_vec: bool, out_r: usize, out_c: usize, i: us /// /// Each operand is indexed through `bcast_flat`: same shape, scalar, /// (1, m) row, (n, 1) column, or a 1-D (`vec`) operand right-aligned to the -/// output's columns. `inline` so the comptime-known node bounds reach the -/// `inline for`. +/// output's columns. `inline` so the comptime-known node shapes and the `op` +/// tag constant-fold into the runtime element loop (the bounds stay comptime +/// constants, but the loop itself is not unrolled). /// /// Args: /// nodes: The full generated node table (for shape lookup). @@ -322,8 +347,8 @@ inline fn ew2(nodes: []const g.Node, node: g.Node, self_idx: usize, buf: *[g.buf const a_n = nodes[node.inputs[0]]; const b_n = nodes[node.inputs[1]]; var out = buf.*[g.offsets[self_idx]..][0 .. node.rows * node.cols]; - inline for (0..node.rows) |i| { - inline for (0..node.cols) |j| { + for (0..node.rows) |i| { + for (0..node.cols) |j| { const av = a[bcast_flat(a_n, a_n.vec, node.rows, node.cols, i, j)]; const bv = b[bcast_flat(b_n, b_n.vec, node.rows, node.cols, i, j)]; out[i * node.cols + j] = switch (op) { diff --git a/tests/fixtures/configs/controllers/onnx_toy.toml b/tests/fixtures/configs/controllers/onnx_toy.toml new file mode 100644 index 0000000..0e61970 --- /dev/null +++ b/tests/fixtures/configs/controllers/onnx_toy.toml @@ -0,0 +1,13 @@ +# Toy ONNX MLP policy fixture (3 -> 4 -> 2 tanh), used by the ONNX importer / +# adapter unit tests and the policy-only compile scenario. +# +# The .onnx is generated by scripts/gen_toy_onnx.py; model_path is CWD-relative, +# so run from the repo root. +type = "onnx_rl" +name = "toy_mlp" +model_path = "tests/fixtures/models/toy_mlp.onnx" +action_space = "continuous" +seed = 0 + +[observation] +state_keys = [0, 1, 2] diff --git a/tests/fixtures/configs/scenarios/toy_mlp_policy.toml b/tests/fixtures/configs/scenarios/toy_mlp_policy.toml new file mode 100644 index 0000000..c1d8cb1 --- /dev/null +++ b/tests/fixtures/configs/scenarios/toy_mlp_policy.toml @@ -0,0 +1,17 @@ +# Policy-only compile scenario: a committed toy ONNX MLP -> lib_neural_network.so. +# +# make compile SCENARIO=tests/fixtures/configs/scenarios/toy_mlp_policy.toml OUT=build/toy_mlp +# +# There is no [estimator]: the ONNX policy is a standalone graph (only onnx_rl +# may omit it). The controller config's model_path is CWD-relative, so run from +# the repo root. artifact_name makes the deployed kernel lib_neural_network.so, +# the name the onnx_rl adapter loads. + +[controller] +type = "onnx_rl" +config = "tests/fixtures/configs/controllers/onnx_toy.toml" + +[compile] +n_x = 3 +n_u = 2 +artifact_name = "lib_neural_network" diff --git a/tests/fixtures/models/toy_mlp.onnx b/tests/fixtures/models/toy_mlp.onnx new file mode 100644 index 0000000..864e6b9 Binary files /dev/null and b/tests/fixtures/models/toy_mlp.onnx differ diff --git a/tests/test_compile_scenario.py b/tests/test_compile_scenario.py index 46663eb..2778142 100644 --- a/tests/test_compile_scenario.py +++ b/tests/test_compile_scenario.py @@ -45,8 +45,7 @@ def _scenario_toml(tmp_path, compile_section: str) -> Path: cfg = tmp_path / "s.toml" cfg.write_text( '[controller]\nconfig = "configs/controllers/lqr_base.toml"\n' - '[estimator]\nconfig = "configs/estimators/kalman_base.toml"\n' - + compile_section + '[estimator]\nconfig = "configs/estimators/kalman_base.toml"\n' + compile_section ) return cfg @@ -104,9 +103,7 @@ def test_scenario_template_stays_in_sync_with_compile_schema(): cfg = tomllib.loads(text) compile_cfg = cfg.get("compile", {}) - assert set(compile_cfg) <= _COMPILE_KEYS, ( - f"template [compile] has keys outside the schema: {set(compile_cfg) - _COMPILE_KEYS}" - ) + assert set(compile_cfg) <= _COMPILE_KEYS, f"template [compile] has keys outside the schema: {set(compile_cfg) - _COMPILE_KEYS}" for key in _COMPILE_KEYS: assert key in text, f"template does not document [compile] key '{key}'" @@ -178,3 +175,106 @@ def test_e2e_shared_graph_untouched(tmp_path): assert _run(GEN, str(SCENARIO), "--out", str(out)).returncode == 0 assert _run(BUILD, str(out), "--scenario", str(SCENARIO)).returncode == 0 assert shipped.read_bytes() == before + + +# ─── policy-only (onnx_rl) scenarios — committed toy fixture ───────────────── + +# A checked-in 3 -> 4 -> 2 tanh MLP (scripts/gen_toy_onnx.py) plus its controller +# config and a policy-only compile scenario, so the whole ONNX -> .so path runs +# against a stable fixture with no model synthesis in the test. +TOY_SCENARIO = REPO_ROOT / "tests" / "fixtures" / "configs" / "scenarios" / "toy_mlp_policy.toml" +TOY_ONNX = REPO_ROOT / "tests" / "fixtures" / "models" / "toy_mlp.onnx" + + +def _policy_scenario(tmp_path, *, controller_config, compile_section): + """A minimal policy-only scenario around a given controller config.""" + cfg = tmp_path / "policy_scenario.toml" + cfg.write_text(f'[controller]\nconfig = "{controller_config}"\n' + compile_section) + return cfg + + +def test_toy_onnx_fixture_matches_its_generator(tmp_path): + """The committed .onnx matches scripts/gen_toy_onnx.py (regenerate-to-check).""" + import numpy as np + from onnx import numpy_helper + + onnx = pytest.importorskip("onnx") + fresh = tmp_path / "toy_mlp.onnx" + result = _run(REPO_ROOT / "scripts" / "gen_toy_onnx.py", "--out", str(fresh)) + assert result.returncode == 0, result.stderr + + committed = onnx.load(str(TOY_ONNX)).graph + regenerated = onnx.load(str(fresh)).graph + assert [n.op_type for n in committed.node] == [n.op_type for n in regenerated.node] + assert [t.name for t in committed.initializer] == [t.name for t in regenerated.initializer] + for a, b in zip(committed.initializer, regenerated.initializer, strict=True): + np.testing.assert_array_equal(numpy_helper.to_array(a), numpy_helper.to_array(b)) + + +def test_policy_scenario_omits_estimator(tmp_path): + """A policy-only scenario (no [estimator]) loads and lowers standalone.""" + from shinro.codegen.scenario_gen import gen_scenario, load_scenario + + pytest.importorskip("onnx") + spec = load_scenario(str(TOY_SCENARIO)) + assert spec["estimator_config"] is None + + out = tmp_path / "g" + cg, graph_path = gen_scenario(str(TOY_SCENARIO), str(out)) + assert cg.inputs == ["state"] + assert cg.outputs == ["u"] + assert cg.state_outputs == [] + assert graph_path.exists() + + +def test_policy_scenario_provenance_pins_onnx_weights(tmp_path): + """The deployment record's config slot commits to the exact .onnx file.""" + from shinro.codegen.scenario_gen import gen_scenario + + pytest.importorskip("onnx") + out = tmp_path / "g" + gen_scenario(str(TOY_SCENARIO), str(out)) + configs = json.loads((out / "graph_data_manifest.json").read_text())["provenance"]["configs"] + key = next(k for k in configs if k.endswith("toy_mlp.onnx")) + assert configs[key] == hashlib.sha256(TOY_ONNX.read_bytes()).hexdigest() + + +def test_non_policy_scenario_still_requires_estimator(tmp_path): + """Only onnx_rl may drop [estimator]; a classical controller still needs one.""" + from shinro.codegen.scenario_gen import load_scenario + + scenario = _policy_scenario( + tmp_path, + controller_config="configs/controllers/lqr_base.toml", + compile_section="[compile]\nn_x = 3\nn_u = 3\n", + ) + with pytest.raises(ValueError, match=r"missing \[estimator\]"): + load_scenario(str(scenario)) + + +def test_compile_rejects_unsafe_artifact_name(tmp_path): + """artifact_name is a file-name stem: separators must be rejected.""" + from shinro.codegen.scenario_gen import load_scenario + + bad = _scenario_toml(tmp_path, '[compile]\nn_x = 3\nn_u = 3\nartifact_name = "../evil"\n') + with pytest.raises(ValueError, match="artifact_name"): + load_scenario(str(bad)) + + +@pytest.mark.skipif(shutil.which("zig") is None, reason="zig not on PATH") +def test_e2e_policy_named_artifact(tmp_path): + """The committed policy scenario builds lib_neural_network.so and verifies.""" + out = tmp_path / "scenario" + gen = _run(GEN, str(TOY_SCENARIO), "--out", str(out)) + assert gen.returncode == 0, gen.stderr + build = _run(BUILD, str(out), "--scenario", str(TOY_SCENARIO)) + assert build.returncode == 0, build.stderr + assert "oracle B" in build.stdout + + so = out / "lib" / "lib_neural_network.so" + record = out / "lib" / "lib_neural_network.deployment.json" + assert so.exists() + assert (out / "lib" / "lib_neural_network.manifest.json").exists() + assert record.exists() + assert not (out / "lib" / "libbase.so").exists() + assert json.loads(record.read_text())["slots"]["binary"] == hashlib.sha256(so.read_bytes()).hexdigest() diff --git a/tests/test_zig_lowering.py b/tests/test_zig_lowering.py index 6b9b204..f417b6d 100644 --- a/tests/test_zig_lowering.py +++ b/tests/test_zig_lowering.py @@ -2347,3 +2347,131 @@ def test_graph_provenance_recorded(self, tmp_path): manifest = json.loads((tmp_path / "graph_data_manifest.json").read_text()) assert manifest["provenance"]["configs"]["configs/controllers/lqr_base.toml"] == "abc123" assert manifest["provenance"]["python_version"] == "3.12" + + +# ─── ONNX policy oracle (imported graph -> .so) ───────────────────────────── + +#: The committed toy policy (scripts/gen_toy_onnx.py): 3 -> 4 -> 2 tanh MLP with +#: the closed form action = [tanh(x0) + 0.5, tanh(x1) - 0.5]. +TOY_ONNX = REPO_ROOT / "tests" / "fixtures" / "models" / "toy_mlp.onnx" + + +def _onnx_policy_graph(action_cfg: dict): + """Import the toy ONNX policy with the given (baked) action-space config.""" + pytest.importorskip("onnx") + from shinro.codegen.onnx_import import import_onnx_policy + + return import_onnx_policy(str(TOY_ONNX), obs_cfg={"state_keys": [0, 1, 2]}, action_cfg=action_cfg) + + +@pytest.fixture(scope="session") +def onnx_continuous_so(tmp_path_factory): + """The continuous (no epsilon port) baked policy kernel.""" + d = tmp_path_factory.mktemp("zig-build-onnx-continuous") + return _build_so(_onnx_policy_graph({"action_space": "continuous"}), d, graph_path=d / "graph_data.zig") + + +@pytest.fixture(scope="session") +def onnx_discrete_so(tmp_path_factory): + """The deterministic discrete kernel (argmax + one_hot, no epsilon port).""" + d = tmp_path_factory.mktemp("zig-build-onnx-discrete") + return _build_so(_onnx_policy_graph({"action_space": "discrete"}), d, graph_path=d / "graph_data.zig") + + +@pytest.fixture(scope="session") +def onnx_discrete_eps_so(tmp_path_factory): + """The sampling discrete kernel: it consumes host Gumbel noise.""" + d = tmp_path_factory.mktemp("zig-build-onnx-discrete-eps") + cfg = {"action_space": "discrete", "deterministic": False} + return _build_so(_onnx_policy_graph(cfg), d, graph_path=d / "graph_data.zig") + + +@pytest.fixture(scope="session") +def onnx_stochastic_eps_so(tmp_path_factory): + """The sampling stochastic kernel; the toy's 2 outputs read as [mean; log_std].""" + d = tmp_path_factory.mktemp("zig-build-onnx-stochastic-eps") + cfg = {"action_space": "stochastic", "deterministic": False} + return _build_so(_onnx_policy_graph(cfg), d, graph_path=d / "graph_data.zig") + + +class TestOnnxPolicyOracle: + """An imported ONNX policy lowers to a .so that matches the interpreter. + + The graph comes from the committed toy fixture, so this is the only oracle + whose subject is a *learned* policy rather than a hand-written control law: + it proves the importer's output (baked encoder, transposed Gemms, composed + activations, and the action post-processing) compiles bit-for-bit. Each + action space gets its own kernel because the space — and whether an + ``epsilon`` port exists — is baked at import time. Graphs lower to tmp + paths, never the shared ``src/shinro/runtime/graph_data.zig``. + """ + + def test_continuous_matches_interpreter_and_closed_form(self, onnx_continuous_so): + lib, cg = onnx_continuous_so + assert cg.inputs == ["state"] # deterministic: no noise port + assert cg.state_outputs == [] + n_out, n_state = output_split(cg) + assert (n_out, n_state) == (2, 0) + + rng = np.random.default_rng(7) + for _ in range(50): + state = rng.normal(0.0, 1.0, 3) + out, _ = step_so(lib, pack_arrays(cg, {"state": state}), n_out, n_state) + traced = interpret(cg.graph, {"state": state})["u"] + closed = np.array([np.tanh(state[0]) + 0.5, np.tanh(state[1]) - 0.5]) + np.testing.assert_allclose(out, traced, rtol=1e-14, atol=1e-14) + np.testing.assert_allclose(out, closed, rtol=1e-6, atol=1e-7) + + def test_discrete_deterministic_one_hot(self, onnx_discrete_so): + lib, cg = onnx_discrete_so + assert cg.inputs == ["state"] + n_out, n_state = output_split(cg) + + rng = np.random.default_rng(8) + for _ in range(25): + state = rng.normal(0.0, 1.0, 3) + out, _ = step_so(lib, pack_arrays(cg, {"state": state}), n_out, n_state) + traced = interpret(cg.graph, {"state": state})["u"] + logits = np.array([np.tanh(state[0]) + 0.5, np.tanh(state[1]) - 0.5]) + want = np.zeros(2) + want[int(np.argmax(logits))] = 1.0 + np.testing.assert_allclose(out, traced, rtol=0, atol=0) + np.testing.assert_allclose(out, want, rtol=0, atol=0) + + def test_discrete_sampling_consumes_gumbel_noise(self, onnx_discrete_eps_so): + lib, cg = onnx_discrete_eps_so + assert cg.inputs == ["state", "epsilon"] + assert input_shape(cg.graph, "epsilon") == (2,) + n_out, n_state = output_split(cg) + + rng = np.random.default_rng(9) + for _ in range(25): + state = rng.normal(0.0, 1.0, 3) + gumbel = -np.log(-np.log(rng.uniform(size=2))) # the host's Gumbel noise + arrays = {"state": state, "epsilon": gumbel} + out, _ = step_so(lib, pack_arrays(cg, arrays), n_out, n_state) + traced = interpret(cg.graph, arrays)["u"] + logits = np.array([np.tanh(state[0]) + 0.5, np.tanh(state[1]) - 0.5]) + want = np.zeros(2) + want[int(np.argmax(logits + gumbel))] = 1.0 + np.testing.assert_allclose(out, traced, rtol=0, atol=0) + np.testing.assert_allclose(out, want, rtol=0, atol=0) + + def test_stochastic_sampling_matches_interpreter_and_formula(self, onnx_stochastic_eps_so): + lib, cg = onnx_stochastic_eps_so + assert cg.inputs == ["state", "epsilon"] + assert input_shape(cg.graph, "epsilon") == (1,) # the toy's 2 outputs -> n_u = 1 + n_out, n_state = output_split(cg) + + rng = np.random.default_rng(10) + for _ in range(25): + state = rng.normal(0.0, 1.0, 3) + eps = rng.normal(size=1) + arrays = {"state": state, "epsilon": eps} + out, _ = step_so(lib, pack_arrays(cg, arrays), n_out, n_state) + traced = interpret(cg.graph, arrays)["u"] + mean = np.tanh(state[0]) + 0.5 + log_std = np.clip(np.tanh(state[1]) - 0.5, -10.0, 2.0) + want = mean + np.exp(log_std) * eps + np.testing.assert_allclose(out, traced, rtol=1e-14, atol=1e-14) + np.testing.assert_allclose(out, want, rtol=1e-12, atol=1e-12) diff --git a/tests/unit/test_onnx_import.py b/tests/unit/test_onnx_import.py new file mode 100644 index 0000000..cdb972e --- /dev/null +++ b/tests/unit/test_onnx_import.py @@ -0,0 +1,608 @@ +"""Tests for the ONNX policy → shinro graph importer. + +The importer is the only path into the codegen machinery that does not use the +tracer: ``onnx.load(path).graph`` is already a dataflow graph, so these tests +build tiny models with ``onnx.helper`` and assert the *translated* graph runs +``interpret()`` to the hand-computed values. Zig parity for the same graphs is +covered separately in ``tests/test_zig_lowering.py``. +""" + +import numpy as np +import pytest + +from shinro.codegen.interpreter import interpret +from shinro.codegen.onnx_import import EPSILON_PORT, OUTPUT_PORT, STATE_PORT, import_onnx_policy + +onnx = pytest.importorskip("onnx") + + +def _vi(name, shape): + """Build a float ValueInfoProto.""" + from onnx import TensorProto, helper + + return helper.make_tensor_value_info(name, TensorProto.FLOAT, shape) + + +def _init(name, array): + """Build a float initializer from a numpy array.""" + from onnx import TensorProto, helper + + a = np.asarray(array, dtype=np.float32) + return helper.make_tensor(name, TensorProto.FLOAT, a.shape, a.flatten().tolist()) + + +def _save(nodes, inputs, outputs, initializers, tmp_path, name="policy.onnx"): + """Serialize an ONNX graph to a temp file and return its path.""" + from onnx import helper + + graph = helper.make_graph(nodes, "g", inputs, outputs, initializers) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + model.ir_version = 8 + path = tmp_path / name + onnx.save(model, str(path)) + return str(path) + + +def _run(cg, x): + """Interpret the imported graph on a state vector.""" + return interpret(cg.graph, {STATE_PORT: np.asarray(x, dtype=np.float64)})[OUTPUT_PORT] + + +def _sample(cg, x, epsilon): + """Interpret the imported graph feeding both the state and the noise port.""" + feed = { + STATE_PORT: np.asarray(x, dtype=np.float64), + EPSILON_PORT: np.asarray(epsilon, dtype=np.float64), + } + return interpret(cg.graph, feed)[OUTPUT_PORT] + + +def _op_names(cg): + return [n.op for n in cg.graph.nodes] + + +class TestPortLayout: + def test_graph_is_memoryless_single_port(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Gemm", ["state", "w", "b"], ["y"], transB=1)], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2, 3)), _init("b", [0.0, 0.0])], + tmp_path, + ) + cg = import_onnx_policy(path) + assert cg.inputs == [STATE_PORT] + assert cg.outputs == [OUTPUT_PORT] + assert cg.state_inputs == [] + assert cg.state_outputs == [] + + def test_all_nodes_rank_at_most_two(self, tmp_path): + """The lowered VM is 2-D only — the importer must not emit rank-3 nodes.""" + from onnx import helper + + path = _save( + [ + helper.make_node("Gemm", ["state", "w1", "b1"], ["h"], transB=1), + helper.make_node("Tanh", ["h"], ["a"]), + helper.make_node("Gemm", ["a", "w2", "b2"], ["y"], transB=1), + ], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [ + _init("w1", np.ones((4, 3))), + _init("b1", np.zeros(4)), + _init("w2", np.ones((2, 4))), + _init("b2", np.zeros(2)), + ], + tmp_path, + ) + cg = import_onnx_policy(path) + assert all(len(n.shape) <= 2 for n in cg.graph.nodes) + + +class TestGemm: + def test_torch_layout_transB(self, tmp_path): + from onnx import helper + + w = np.array([[1.0, 0.0, 0.0], [0.0, 2.0, 0.0]], dtype=np.float32) # (2, 3) + b = np.array([0.5, -0.5], dtype=np.float32) + path = _save( + [helper.make_node("Gemm", ["state", "w", "b"], ["y"], transB=1)], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w", w), _init("b", b)], + tmp_path, + ) + cg = import_onnx_policy(path) + x = np.array([1.0, 2.0, 3.0]) + np.testing.assert_allclose(_run(cg, x), x @ w.T + b, rtol=1e-6) + # transB is realized with a real transpose node, not a baked transposed const. + assert {"transpose", "matmul", "add"} <= set(_op_names(cg)) + assert _op_names(cg).count("add") == 2 # Gemm bias + the action bias + + def test_default_layout_transB_off(self, tmp_path): + from onnx import helper + + w = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32) # (3, 2) + path = _save( + [helper.make_node("Gemm", ["state", "w"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w", w)], + tmp_path, + ) + cg = import_onnx_policy(path) + x = np.array([1.0, 1.0, 1.0]) + np.testing.assert_allclose(_run(cg, x), x @ w, rtol=1e-6) + assert "transpose" not in _op_names(cg) + + def test_alpha_beta(self, tmp_path): + from onnx import helper + + w = np.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], dtype=np.float32) # (2, 3) for transB + c = np.array([2.0, 4.0], dtype=np.float32) + path = _save( + [helper.make_node("Gemm", ["state", "w", "c"], ["y"], transB=1, alpha=0.5, beta=2.0)], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w", w), _init("c", c)], + tmp_path, + ) + cg = import_onnx_policy(path) + x = np.array([2.0, 3.0, 0.0]) + np.testing.assert_allclose(_run(cg, x), 0.5 * (x @ w.T) + 2.0 * c, rtol=1e-6) + # Gemm's alpha + beta multipliers, plus the action-surface scale. + assert _op_names(cg).count("mul") == 3 + + def test_no_bias(self, tmp_path): + from onnx import helper + + w = np.array([[1.0, 0.0], [0.0, 1.0]], dtype=np.float32) + path = _save( + [helper.make_node("Gemm", ["state", "w"], ["y"], transB=1)], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", w)], + tmp_path, + ) + cg = import_onnx_policy(path) + np.testing.assert_allclose(_run(cg, [3.0, 4.0]), [3.0, 4.0], rtol=1e-6) + assert _op_names(cg).count("add") == 1 # the action bias only; the Gemm has none + + def test_transA_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Gemm", ["state", "w"], ["y"], transA=1)], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [_init("w", np.eye(3))], + tmp_path, + ) + with pytest.raises(NotImplementedError, match="transA"): + import_onnx_policy(path) + + def test_unknown_attribute_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Gemm", ["state", "w"], ["y"], broadcast=1)], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [_init("w", np.eye(3))], + tmp_path, + ) + with pytest.raises(NotImplementedError, match="broadcast"): + import_onnx_policy(path) + + +class TestActivations: + def test_mlp_tanh(self, tmp_path): + from onnx import helper + + w1 = np.arange(12, dtype=np.float32).reshape(4, 3) / 10.0 + b1 = np.array([0.1, -0.2, 0.3, -0.4], dtype=np.float32) + w2 = np.arange(8, dtype=np.float32).reshape(2, 4) / 5.0 + b2 = np.array([0.5, -0.5], dtype=np.float32) + path = _save( + [ + helper.make_node("Gemm", ["state", "w1", "b1"], ["h"], transB=1), + helper.make_node("Tanh", ["h"], ["a"]), + helper.make_node("Gemm", ["a", "w2", "b2"], ["y"], transB=1), + ], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w1", w1), _init("b1", b1), _init("w2", w2), _init("b2", b2)], + tmp_path, + ) + cg = import_onnx_policy(path) + x = np.array([0.5, -1.5, 2.0]) + expected = np.tanh(x @ w1.T + b1) @ w2.T + b2 + np.testing.assert_allclose(_run(cg, x), expected, rtol=1e-6, atol=1e-6) + + def test_relu(self, tmp_path): + from onnx import helper + + path = _save( + [ + helper.make_node("Gemm", ["state", "w"], ["h"], transB=1), + helper.make_node("Relu", ["h"], ["y"]), + ], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + cg = import_onnx_policy(path) + np.testing.assert_allclose(_run(cg, [3.0, -4.0]), [3.0, 0.0]) + assert "relu" in _op_names(cg) + + def test_sigmoid_composed_from_existing_ops(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Sigmoid", ["state"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [], + tmp_path, + ) + cg = import_onnx_policy(path) + # No sigmoid op exists in the VM; the importer composes it. + assert {"neg", "exp", "add", "div"} <= set(_op_names(cg)) + x = np.array([-1.0, 0.0, 2.0]) + np.testing.assert_allclose(_run(cg, x), 1.0 / (1.0 + np.exp(-x)), rtol=1e-12) + + +class TestPointwiseGraphs: + def test_matmul_add(self, tmp_path): + from onnx import helper + + w = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32) # (3, 2) + b = np.array([1.0, -1.0], dtype=np.float32) + path = _save( + [ + helper.make_node("MatMul", ["state", "w"], ["m"]), + helper.make_node("Add", ["m", "b"], ["y"]), + ], + [_vi("state", [None, 3])], + [_vi("y", [None, 2])], + [_init("w", w), _init("b", b)], + tmp_path, + ) + cg = import_onnx_policy(path) + x = np.array([1.0, 2.0, 3.0]) + np.testing.assert_allclose(_run(cg, x), x @ w + b, rtol=1e-6) + + +class TestObservationEncoder: + def test_selection_normalize_clip(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Gemm", ["state", "w"], ["y"], transB=1)], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + cg = import_onnx_policy( + path, + n_x=3, + obs_cfg={ + "state_keys": [2, 0], + "normalize": True, + "obs_mean": [1.0, 2.0], + "obs_std": [2.0, 4.0], + "clip": [-1.0, 1.0], + }, + ) + # state [10, 0, 5] -> obs [5, 10] -> ([4, 8])/[2,4] = [2,2] -> clip [1,1] + np.testing.assert_allclose(_run(cg, [10.0, 0.0, 5.0]), [1.0, 1.0], rtol=1e-6) + + def test_identity_selection_skips_matmul(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + cg = import_onnx_policy(path, n_x=2, obs_cfg={"state_keys": [0, 1]}) + assert _op_names(cg).count("matmul") == 1 # only the policy's own matmul + + def test_state_keys_infer_n_x(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + cg = import_onnx_policy(path, obs_cfg={"state_keys": [2, 0]}) + # n_x = max(state_keys)+1 = 3; state [7, 8, 9] -> obs [9, 7] + np.testing.assert_allclose(_run(cg, [7.0, 8.0, 9.0]), [9.0, 7.0], rtol=1e-6) + + def test_normalize_without_stats_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + with pytest.raises(ValueError, match="obs_mean"): + import_onnx_policy(path, obs_cfg={"normalize": True}) + + def test_obs_dim_mismatch_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [_init("w", np.eye(3))], + tmp_path, + ) + with pytest.raises(ValueError, match="state_keys"): + import_onnx_policy(path, obs_cfg={"state_keys": [0, 1]}) + + def test_state_keys_out_of_range_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + with pytest.raises(ValueError, match="out of range"): + import_onnx_policy(path, n_x=2, obs_cfg={"state_keys": [0, 5]}) + + def test_unknown_observation_key_rejected(self, tmp_path): + """A typo like `obs_means` must not silently drop normalization.""" + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + with pytest.raises(ValueError, match="unknown key"): + import_onnx_policy(path, obs_cfg={"obs_means": [0.0, 0.0]}) + + +class TestRejections: + def test_unsupported_op_names_itself(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Softmax", ["state"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [], + tmp_path, + ) + with pytest.raises(NotImplementedError, match="Softmax"): + import_onnx_policy(path) + + def test_unreachable_unsupported_node_ignored(self, tmp_path): + from onnx import helper + + path = _save( + [ + helper.make_node("Softmax", ["state"], ["junk"]), + helper.make_node("MatMul", ["state", "w"], ["y"]), + ], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2))], + tmp_path, + ) + cg = import_onnx_policy(path) + assert "Softmax" not in _op_names(cg) + np.testing.assert_allclose(_run(cg, [1.0, 2.0]), [1.0, 2.0], rtol=1e-6) + + def test_multi_input_policy_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Add", ["a", "b"], ["y"])], + [_vi("a", [None, 2]), _vi("b", [None, 2])], + [_vi("y", [None, 2])], + [], + tmp_path, + ) + with pytest.raises(ValueError, match="exactly one"): + import_onnx_policy(path) + + def test_rank2_batch1_output_flattened(self, tmp_path): + from onnx import helper + + # (3,) + (1, 3) broadcasts to (1, 3), a rank-2 batch-1 result the + # importer must reshape down to the (n_u,) action port. + path = _save( + [helper.make_node("Add", ["state", "c"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [_init("c", np.array([[1.0, 2.0, 3.0]], dtype=np.float32))], + tmp_path, + ) + cg = import_onnx_policy(path) + u = _run(cg, [1.0, 1.0, 1.0]) + assert u.shape == (3,) + np.testing.assert_allclose(u, [2.0, 3.0, 4.0], rtol=1e-6) + assert "reshape" in _op_names(cg) + + def test_batched_output_rejected(self, tmp_path): + from onnx import helper + + # (3,) + (2, 1) broadcasts to (2, 3) — a genuine batch, not a single + # action vector, so the importer must refuse it rather than drop a row. + path = _save( + [helper.make_node("Add", ["state", "c"], ["y"])], + [_vi("state", [None, 3])], + [_vi("y", [None, 3])], + [_init("c", np.array([[1.0], [2.0]], dtype=np.float32))], + tmp_path, + ) + with pytest.raises(ValueError, match="action vector"): + import_onnx_policy(path) + + def test_unknown_tensor_rejected(self, tmp_path): + from onnx import helper + + path = _save( + [helper.make_node("Add", ["state", "missing"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [], + tmp_path, + ) + with pytest.raises(ValueError, match="unknown tensor"): + import_onnx_policy(path) + + +def _gemm_policy(tmp_path, w, b=None, *, name="policy.onnx"): + """Build a single-Gemm policy in torch layout (transB=1).""" + from onnx import helper + + w = np.asarray(w, dtype=np.float32) + inputs = ["state", "w"] + (["b"] if b is not None else []) + inits = [_init("w", w)] + ([_init("b", np.asarray(b, dtype=np.float32))] if b is not None else []) + return _save( + [helper.make_node("Gemm", inputs, ["y"], transB=1)], + [_vi("state", [None, w.shape[1]])], + [_vi("y", [None, w.shape[0]])], + inits, + tmp_path, + name, + ) + + +class TestActionSurface: + """The baked post-processing must mirror the old runtime `_postprocess`.""" + + def _tiny(self, tmp_path): + """2-output policy: raw(x) = [x0 + 0.5, 2*x1 - 0.5].""" + return _gemm_policy(tmp_path, [[1.0, 0.0, 0.0], [0.0, 2.0, 0.0]], [0.5, -0.5]) + + def _stochastic(self, tmp_path): + """4-output policy: raw(x) = [x0+1, x1+2, x2+3, 4] = [mean; log_std].""" + w = [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0], [0.0, 0.0, 0.0]] + return _gemm_policy(tmp_path, w, [1.0, 2.0, 3.0, 4.0], name="stochastic.onnx") + + def test_continuous_default_passthrough(self, tmp_path): + cg = import_onnx_policy(self._tiny(tmp_path)) + assert cg.inputs == [STATE_PORT] + assert "clip" not in _op_names(cg) # no clip configured + np.testing.assert_allclose(_run(cg, [1.0, 2.0, 3.0]), [1.5, 3.5], rtol=1e-6) + + def test_continuous_default_still_emits_scale_bias(self, tmp_path): + """Uniform lowering: the default 1.0 / 0.0 still produce the mul/add nodes.""" + from onnx import helper + + path = _save( + [helper.make_node("MatMul", ["state", "w"], ["y"])], + [_vi("state", [None, 2])], + [_vi("y", [None, 2])], + [_init("w", np.eye(2, dtype=np.float32))], + tmp_path, + ) + cg = import_onnx_policy(path) # default continuous action config + ops = _op_names(cg) + assert ops.count("mul") == 1 # action scale, emitted even though it is 1.0 + assert ops.count("add") == 1 # action bias, emitted even though it is 0.0 + + def test_continuous_scale_bias_clip(self, tmp_path): + cg = import_onnx_policy( + self._tiny(tmp_path), + action_cfg={"action_scale": 2.0, "action_bias": 1.0, "action_clip_low": -3.0, "action_clip_high": 3.0}, + ) + # raw [1.5, 1.5] -> *2+1 = [4, 4] -> clipped to 3 + np.testing.assert_allclose(_run(cg, [1.0, 1.0, 0.0]), [3.0, 3.0], rtol=1e-6) + + def test_continuous_vector_scale_bias(self, tmp_path): + path = _gemm_policy(tmp_path, [[1.0, 0.0, 0.0], [0.0, 2.0, 0.0]], name="vec.onnx") + cg = import_onnx_policy(path, action_cfg={"action_scale": [2.0, 3.0], "action_bias": [1.0, -1.0]}) + # raw [1, 2] -> [1*2+1, 2*3-1] = [3, 5] + np.testing.assert_allclose(_run(cg, [1.0, 1.0, 0.0]), [3.0, 5.0], rtol=1e-6) + + def test_discrete_deterministic_one_hot(self, tmp_path): + cg = import_onnx_policy(self._tiny(tmp_path), action_cfg={"action_space": "discrete"}) + assert cg.inputs == [STATE_PORT] # deterministic: no noise port + assert {"argmax", "one_hot"} <= set(_op_names(cg)) + # equal logits [1.5, 1.5] -> first-max wins, matching numpy argmax + np.testing.assert_allclose(_run(cg, [1.0, 1.0, 1.0]), [1.0, 0.0], rtol=1e-6) + + def test_discrete_ignores_scale_bias(self, tmp_path): + cg = import_onnx_policy( + self._tiny(tmp_path), + action_cfg={"action_space": "discrete", "action_scale": 5.0, "action_bias": 1.0}, + ) + np.testing.assert_allclose(_run(cg, [1.0, 1.0, 1.0]), [1.0, 0.0], rtol=1e-6) + + def test_discrete_gumbel_max_uses_epsilon(self, tmp_path): + cg = import_onnx_policy(self._tiny(tmp_path), action_cfg={"action_space": "discrete", "deterministic": False}) + assert cg.inputs == [STATE_PORT, EPSILON_PORT] + state = np.array([1.0, 1.0, 1.0]) + # A huge positive Gumbel draw flips the argmax; the kernel only adds. + np.testing.assert_allclose(_sample(cg, state, [0.0, 100.0]), [0.0, 1.0], rtol=1e-6) + np.testing.assert_allclose(_sample(cg, state, [100.0, 0.0]), [1.0, 0.0], rtol=1e-6) + + def test_discrete_gumbel_max_matches_softmax(self, tmp_path): + # logits = [0, ln 3] -> softmax p(1) = 0.75; Gumbel-max must reproduce it. + path = _gemm_policy(tmp_path, np.zeros((2, 3), dtype=np.float32), [0.0, np.log(3.0)], name="dist.onnx") + cg = import_onnx_policy(path, action_cfg={"action_space": "discrete", "deterministic": False}) + rng = np.random.default_rng(0) + n = 2000 + draws = -np.log(-np.log(rng.uniform(size=(n, 2)))) + state = np.zeros(3) + hits = sum(int(np.argmax(_sample(cg, state, eps))) for eps in draws) + assert abs(hits / n - 0.75) < 0.06 + + def test_stochastic_deterministic_returns_mean(self, tmp_path): + cg = import_onnx_policy(self._stochastic(tmp_path), action_cfg={"action_space": "stochastic"}) + assert cg.inputs == [STATE_PORT] + # raw = [11, 22, 33, 4] -> mean = [11, 22] + np.testing.assert_allclose(_run(cg, [10.0, 20.0, 30.0]), [11.0, 22.0], rtol=1e-6) + + def test_stochastic_epsilon_formula(self, tmp_path): + cg = import_onnx_policy(self._stochastic(tmp_path), action_cfg={"action_space": "stochastic", "deterministic": False}) + assert cg.inputs == [STATE_PORT, EPSILON_PORT] + # raw = [11, 22, 33, 4]; log_std clipped to 2 -> std = e^2 + u = _sample(cg, [10.0, 20.0, 30.0], [0.5, -1.0]) + std = np.exp(2.0) + np.testing.assert_allclose(u, [11.0 + std * 0.5, 22.0 - std], rtol=1e-6) + + def test_stochastic_odd_output_rejected(self, tmp_path): + path = _gemm_policy(tmp_path, [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], name="odd.onnx") + with pytest.raises(ValueError, match="even"): + import_onnx_policy(path, action_cfg={"action_space": "stochastic"}) + + def test_invalid_action_space_rejected(self, tmp_path): + with pytest.raises(ValueError, match="action_space"): + import_onnx_policy(self._tiny(tmp_path), action_cfg={"action_space": "bogus"}) + + def test_single_sided_clip_rejected(self, tmp_path): + # The missing bound would be ±inf, which Zig cannot represent as a literal. + with pytest.raises(ValueError, match="together"): + import_onnx_policy(self._tiny(tmp_path), action_cfg={"action_clip_low": -1.0}) + + def test_infinite_clip_rejected(self, tmp_path): + with pytest.raises(ValueError, match="finite"): + import_onnx_policy( + self._tiny(tmp_path), + action_cfg={"action_clip_low": -np.inf, "action_clip_high": np.inf}, + ) diff --git a/tests/unit/test_onnx_rl_adapter.py b/tests/unit/test_onnx_rl_adapter.py index f3a5130..e330275 100644 --- a/tests/unit/test_onnx_rl_adapter.py +++ b/tests/unit/test_onnx_rl_adapter.py @@ -1,345 +1,354 @@ -"""Tests for the ONNX RL policy adapter.""" +"""Tests for the ONNX RL policy adapter (eager mode). + +The adapter no longer calls ``onnxruntime``: ``from_config`` imports the ONNX +graph into a shinro graph and runs it with the interpreter. The graph-level +behavior (encoder folding, action spaces, epsilon ports) is pinned in +``tests/unit/test_onnx_import.py``; these tests cover the *adapter* contract — +strict config parsing, backend conversion, RNG seeding, and error surfaces. +Compiled-artifact mode is exercised end-to-end by the Zig oracle suite. +""" + +import dataclasses +import json +from pathlib import Path import numpy as np import pytest -from shinro.controllers.onnx_rl_adapter import OnnxRLAdapter, _ObsEncoder +from shinro.controllers.onnx_rl_adapter import KERNEL_FILENAME, OnnxRLAdapter, OnnxRLConfig, _CompiledPolicy onnx = pytest.importorskip("onnx") -onnxruntime = pytest.importorskip("onnxruntime") +#: The committed toy policy (scripts/gen_toy_onnx.py): action = [tanh(x0)+0.5, tanh(x1)-0.5]. +_TOY = Path(__file__).resolve().parents[1] / "fixtures" / "models" / "toy_mlp.onnx" -def _build_model(input_name: str = "obs", output_name: str = "output"): - """Build a tiny ONNX linear model: obs -> Gemm -> output.""" - import tempfile +def _save_model(w, b, tmp_path, *, input_name="obs", output_name="output", name="policy.onnx"): + """Write a single-Gemm (torch layout, transB=1) policy and return its path.""" from onnx import TensorProto, helper - w = np.array([[1.0, 0.0], [0.0, 2.0], [0.0, 0.0]], dtype=np.float32) - b = np.array([0.5, -0.5], dtype=np.float32) - X = helper.make_tensor_value_info(input_name, TensorProto.FLOAT, [None, 3]) - Y = helper.make_tensor_value_info(output_name, TensorProto.FLOAT, [None, 2]) + w = np.asarray(w, dtype=np.float32) + b = np.asarray(b, dtype=np.float32) + assert w.shape[0] == b.shape[0], (w.shape, b.shape) + x = helper.make_tensor_value_info(input_name, TensorProto.FLOAT, [None, w.shape[1]]) + y = helper.make_tensor_value_info(output_name, TensorProto.FLOAT, [None, w.shape[0]]) + node = helper.make_node("Gemm", [input_name, "w", "b"], [output_name], transB=1) w_init = helper.make_tensor("w", TensorProto.FLOAT, w.shape, w.flatten().tolist()) b_init = helper.make_tensor("b", TensorProto.FLOAT, b.shape, b.flatten().tolist()) - node = helper.make_node("Gemm", [input_name, "w", "b"], [output_name]) - graph = helper.make_graph([node], "g", [X], [Y], [w_init, b_init]) + graph = helper.make_graph([node], "g", [x], [y], [w_init, b_init]) model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) model.ir_version = 8 - path = tempfile.NamedTemporaryFile(suffix=".onnx", delete=False).name - onnx.save(model, path) - return path + path = tmp_path / name + onnx.save(model, str(path)) + return str(path) -@pytest.fixture(scope="module") -def model_path(): - return _build_model() +# raw(x) = [x0 + 0.5, 2*x1 - 0.5] +_W_2OUT = [[1.0, 0.0, 0.0], [0.0, 2.0, 0.0]] +_B_2OUT = [0.5, -0.5] +# raw(x) = [x0 + 1, x1 + 2, x2 + 3, 4] = [mean; log_std] for a 2-action policy +_W_4OUT = [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0], [0.0, 0.0, 0.0]] +_B_4OUT = [1.0, 2.0, 3.0, 4.0] -def _make_encoder( - *, - input_name: str = "obs", - state_keys: list[int] | None = None, - obs_mean: np.ndarray | None = None, - obs_std: np.ndarray | None = None, - clip: tuple[float, float] | None = None, - add_batch_dim: bool = True, -) -> _ObsEncoder: - if state_keys is None: - state_keys = [0, 1, 2] - return _ObsEncoder(input_name, state_keys, obs_mean, obs_std, clip, add_batch_dim) +@pytest.fixture(scope="module") +def model_path(tmp_path_factory): + return _save_model(_W_2OUT, _B_2OUT, tmp_path_factory.mktemp("onnx")) -def test_obs_encoder_index_selection(): - enc = _make_encoder(state_keys=[2, 0]) - feed = enc.encode(np.array([10.0, 20.0, 30.0])) - np.testing.assert_allclose(feed["obs"], [[30.0, 10.0]]) +@pytest.fixture(scope="module") +def stochastic_path(tmp_path_factory): + return _save_model(_W_4OUT, _B_4OUT, tmp_path_factory.mktemp("onnx"), name="stochastic.onnx") -def test_obs_encoder_normalize_and_clip(): - enc = _make_encoder(obs_mean=np.array([1.0, 2.0, 3.0]), obs_std=np.array([2.0, 2.0, 2.0]), clip=(-1.0, 1.0)) - feed = enc.encode(np.array([10.0, 10.0, 10.0])) - # (10-1)/2=4.5 clipped to 1.0 - np.testing.assert_allclose(feed["obs"], [[1.0, 1.0, 1.0]]) +def _ctrl(model_path, **overrides): + cfg = {"model_path": str(model_path), "action_space": "continuous"} + cfg.update(overrides) + return OnnxRLAdapter.from_config(cfg) -def test_obs_encoder_no_batch_dim(): - enc = _make_encoder(add_batch_dim=False) - feed = enc.encode(np.array([1.0, 2.0, 3.0])) - assert feed["obs"].shape == (3,) +class TestConfigSurface: + def test_declares_a_frozen_config_dataclass(self): + """The registry checks this (and it becomes a hard error eventually).""" + assert dataclasses.is_dataclass(OnnxRLConfig) + assert OnnxRLAdapter.Config is OnnxRLConfig + with pytest.raises(dataclasses.FrozenInstanceError): + setattr(OnnxRLConfig(), "name", "mutated") + def test_unknown_key_rejected(self, model_path): + with pytest.raises(ValueError, match="unknown key"): + OnnxRLAdapter.from_config({"model_path": str(model_path), "bogus": 1}) -class TestOnnxRLAdapter: - def test_from_config_continuous(self, model_path, tmp_path): - config = tmp_path / "rl.toml" - config.write_text(f'type = "onnx_rl"\nmodel_path = "{model_path}"\naction_space = "continuous"\n') - cfg = {"type": "onnx_rl", "model_path": str(model_path), "action_space": "continuous"} - ctrl = OnnxRLAdapter.from_config(cfg) - action = ctrl.compute(np.array([1.0, 2.0, 3.0])) - assert action.shape == (2,) - expected = np.array([1.0 * 1.0 + 0.5, 2.0 * 2.0 - 0.5]) - np.testing.assert_allclose(action, expected) - - def test_continuous_action_scale_bias_clip(self, model_path): - cfg = { - "model_path": str(model_path), - "action_space": "continuous", - "action_scale": 2.0, - "action_bias": 1.0, - "action_clip_low": -3.0, - "action_clip_high": 3.0, - } - ctrl = OnnxRLAdapter.from_config(cfg) - action = ctrl.compute(np.array([1.0, 1.0, 0.0])) - # w[:,0] = [1,0,0], bias 0.5 -> (1*2+1)=3, clipped to 3.0 - np.testing.assert_allclose(action, [3.0, 3.0]) + def test_missing_mode_rejected(self): + with pytest.raises(ValueError, match="model_path"): + OnnxRLAdapter.from_config({"action_space": "continuous"}) def test_action_space_invalid(self, model_path): with pytest.raises(ValueError, match="action_space"): OnnxRLAdapter.from_config({"model_path": str(model_path), "action_space": "bogus"}) - def test_reset_reseeds(self, model_path): - cfg = {"model_path": str(model_path), "action_space": "stochastic", "deterministic": False, "seed": 7} - ctrl = OnnxRLAdapter.from_config(cfg) - ctrl.compute(np.array([0.0, 0.0, 0.0])) - ctrl.reset() - ctrl2 = OnnxRLAdapter.from_config(cfg) - a1 = ctrl.compute(np.array([0.0, 0.0, 0.0])) - a2 = ctrl2.compute(np.array([0.0, 0.0, 0.0])) - np.testing.assert_allclose(a1, a2) + def test_default_action_space_is_continuous(self, model_path): + assert _ctrl(model_path).compute(np.array([1.0, 2.0, 3.0])).shape == (2,) + def test_single_sided_action_clip_rejected(self, model_path): + # A missing bound would be ±inf, which the lowerer cannot emit. + with pytest.raises(ValueError, match="together"): + OnnxRLAdapter.from_config({"model_path": str(model_path), "action_clip_low": 2.0}) -class TestDiscreteActionSpace: - def test_deterministic_argmax(self, model_path): - cfg = {"model_path": str(model_path), "action_space": "discrete", "deterministic": True} - ctrl = OnnxRLAdapter.from_config(cfg) - # output for obs [1,1,1]: [1.5, 1.5] -> argmax=0 (first max) - action = ctrl.compute(np.array([1.0, 1.0, 1.0])) - assert action.shape == (2,) - assert action.dtype == np.float32 - assert action[0] == 1.0 and action[1] == 0.0 + def test_normalize_without_stats_rejected(self, model_path): + with pytest.raises(ValueError, match="obs_mean"): + OnnxRLAdapter.from_config({"model_path": str(model_path), "observation": {"normalize": True}}) - def test_stochastic_samples_one_hot(self, model_path): - cfg = {"model_path": str(model_path), "action_space": "discrete", "deterministic": False, "seed": 1} - ctrl = OnnxRLAdapter.from_config(cfg) - action = ctrl.compute(np.array([1.0, 1.0, 1.0])) - assert set(np.unique(action)) <= {0.0, 1.0} - assert action.sum() == 1.0 + def test_config_file_round_trip(self, model_path, tmp_path): + """The TOML shape the shipped config uses parses strictly.""" + config = tmp_path / "rl.toml" + config.write_text( + f'type = "onnx_rl"\nname = "ppo_policy"\nmodel_path = "{model_path}"\n' + 'action_space = "continuous"\ndeterministic = true\naction_scale = 1.0\n' + "action_bias = 0.0\nseed = 0\n\n[observation]\nstate_keys = [0, 1, 2]\nnormalize = false\n" + ) + from shinro.factories.controller_factory import ControllerFactory + ctrl = ControllerFactory(str(config)).create() + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 2.0, 3.0])), [1.5, 3.5], rtol=1e-9) -class TestStochasticActionSpace: - def test_deterministic_returns_mean(self, model_path): - cfg = {"model_path": str(model_path), "action_space": "stochastic", "deterministic": True} - ctrl = OnnxRLAdapter.from_config(cfg) - action = ctrl.compute(np.array([1.0, 2.0, 3.0])) - # mean part = same as continuous output - expected = np.array([1.0 * 1.0 + 0.5, 2.0 * 2.0 - 0.5]) - np.testing.assert_allclose(action, expected) - def test_sample_reproducible(self, model_path): - cfg = {"model_path": str(model_path), "action_space": "stochastic", "deterministic": False, "seed": 5} - ctrl = OnnxRLAdapter.from_config(cfg) - ctrl2 = OnnxRLAdapter.from_config(cfg) - a1 = ctrl.compute(np.array([0.0, 0.0, 0.0])) - a2 = ctrl2.compute(np.array([0.0, 0.0, 0.0])) - np.testing.assert_allclose(a1, a2) - - -class TestFromConfigSurface: - """Cover the config-parsing branches in from_config().""" - - def test_default_action_space_continuous(self, model_path): - """No action_space field -> defaults to continuous.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path)}) - assert ctrl.action_space == "continuous" - - def test_default_state_keys_from_input_shape(self, model_path): - """No state_keys -> defaults to all input dims, so compute works.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path)}) - action = ctrl.compute(np.array([1.0, 2.0, 3.0])) - expected = np.array([1.5, 3.5]) - np.testing.assert_allclose(action, expected) - - def test_custom_input_output_names(self, tmp_path): - """Non-default ONNX I/O tensor names are honored.""" - path = _build_model(input_name="policy_in", output_name="policy_out") - ctrl = OnnxRLAdapter.from_config({"model_path": path, "action_space": "continuous"}) - action = ctrl.compute(np.array([1.0, 0.0, 0.0])) - # obs[0]=1 -> w[0]=1, bias 0.5 -> [1.5, -0.5] - np.testing.assert_allclose(action, [1.5, -0.5]) - - def test_custom_input_name_override(self, model_path): - """observation.input_name overrides the session's input name.""" - ctrl = OnnxRLAdapter.from_config( - {"model_path": str(model_path), "observation": {"input_name": "obs", "state_keys": [1, 2, 0]}} - ) - action = ctrl.compute(np.array([0.0, 1.0, 0.0])) - # obs = [1,0,0] -> [1.5, -0.5] - np.testing.assert_allclose(action, [1.5, -0.5]) +class TestContinuous: + def test_values(self, model_path): + np.testing.assert_allclose(_ctrl(model_path).compute(np.array([1.0, 2.0, 3.0])), [1.5, 3.5], rtol=1e-9) + + def test_scale_bias_clip(self, model_path): + ctrl = _ctrl(model_path, action_scale=2.0, action_bias=1.0, action_clip_low=-3.0, action_clip_high=3.0) + # raw [1.5, 1.5] -> *2+1 = [4, 4] -> clipped to 3 + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 1.0, 0.0])), [3.0, 3.0], rtol=1e-9) - def test_output_name_override(self, model_path): - """output_name field selects a different ONNX output.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path), "output_name": "output"}) - action = ctrl.compute(np.array([1.0, 0.0, 0.0])) - np.testing.assert_allclose(action, [1.5, -0.5]) + def test_vector_scale_bias(self, model_path): + ctrl = _ctrl(model_path, action_scale=[2.0, 3.0], action_bias=[1.0, -1.0]) + # raw [1.5, 3.5] -> [1.5*2 + 1, 3.5*3 - 1] = [4, 9.5] + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 2.0, 0.0])), [4.0, 9.5], rtol=1e-9) def test_obs_normalization_from_config(self, model_path): - """observation.normalize applies mean/std from config.""" - ctrl = OnnxRLAdapter.from_config( - { - "model_path": str(model_path), - "observation": {"normalize": True, "obs_mean": [1.0, 1.0, 1.0], "obs_std": [2.0, 2.0, 2.0]}, - } - ) - action = ctrl.compute(np.array([3.0, 5.0, 1.0])) - # obs = [(3-1)/2, (5-1)/2, 0] = [1,2,0] -> [1.5, 3.5] - np.testing.assert_allclose(action, [1.5, 3.5]) + ctrl = _ctrl(model_path, observation={"normalize": True, "obs_mean": [1.0, 1.0, 1.0], "obs_std": [2.0, 2.0, 2.0]}) + # obs = [(3-1)/2, (5-1)/2, 0] = [1, 2, 0] -> [1.5, 3.5] + np.testing.assert_allclose(ctrl.compute(np.array([3.0, 5.0, 1.0])), [1.5, 3.5], rtol=1e-9) def test_obs_clip_from_config(self, model_path): - """observation.clip clamps observations.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path), "observation": {"clip": [-1.0, 1.0]}}) - action = ctrl.compute(np.array([5.0, 0.0, 0.0])) - # obs[0] clipped to 1.0 -> [1.5, -0.5] - np.testing.assert_allclose(action, [1.5, -0.5]) - - def test_action_clip_one_sided(self, model_path): - """action_clip_low alone -> clip at inf high bound.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path), "action_clip_low": 2.0}) - action = ctrl.compute(np.array([1.0, 0.0, 0.0])) - # raw = [1.5, -0.5]; low-clip to 2.0 - np.testing.assert_allclose(action, [2.0, 2.0]) - - def test_action_clip_high_only(self, model_path): - """action_clip_high alone -> clip at -inf low bound.""" - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path), "action_clip_high": -1.0}) - action = ctrl.compute(np.array([1.0, 0.0, 0.0])) - np.testing.assert_allclose(action, [-1.0, -1.0]) - - def test_missing_stats_with_normalize_raises(self, model_path): - """normalize=true without obs_mean/obs_std raises KeyError.""" - with pytest.raises(KeyError): - OnnxRLAdapter.from_config({"model_path": str(model_path), "observation": {"normalize": True}}) + ctrl = _ctrl(model_path, observation={"clip": [-1.0, 1.0]}) + # obs[0] = min(5, 1) = 1 -> [1.5, -0.5] + np.testing.assert_allclose(ctrl.compute(np.array([5.0, 0.0, 0.0])), [1.5, -0.5], rtol=1e-9) + + def test_state_keys_override(self, model_path): + ctrl = _ctrl(model_path, observation={"state_keys": [1, 2, 0]}) + # obs = [x1, x2, x0] = [1, 0, 0] -> [1.5, -0.5] + np.testing.assert_allclose(ctrl.compute(np.array([0.0, 1.0, 0.0])), [1.5, -0.5], rtol=1e-9) + + def test_n_x_beyond_observation_reach(self, model_path): + """A state larger than the observed entries needs only n_x + state_keys.""" + ctrl = _ctrl(model_path, n_x=4, observation={"state_keys": [0, 1, 2]}) + # state is 4 long, obs reads the first three + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 2.0, 3.0, 9.0])), [1.5, 3.5], rtol=1e-9) + + def test_state_size_mismatch_raises(self, model_path): + with pytest.raises(ValueError, match="expected a state of 3"): + _ctrl(model_path).compute(np.array([1.0, 2.0])) + + def test_output_name_override(self, tmp_path): + path = _save_model(_W_2OUT, _B_2OUT, tmp_path, output_name="action") + ctrl = OnnxRLAdapter.from_config({"model_path": path, "output_name": "action"}) + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 0.0, 0.0])), [1.5, -0.5], rtol=1e-9) + + def test_custom_input_name(self, tmp_path): + path = _save_model(_W_2OUT, _B_2OUT, tmp_path, input_name="policy_in") + ctrl = OnnxRLAdapter.from_config({"model_path": path, "observation": {"input_name": "policy_in"}}) + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 0.0, 0.0])), [1.5, -0.5], rtol=1e-9) + + def test_wrong_input_name_override_rejected(self, model_path): + with pytest.raises(ValueError, match="input_name"): + OnnxRLAdapter.from_config({"model_path": str(model_path), "observation": {"input_name": "nope"}}) + + def test_compute_accepts_list(self, model_path): + np.testing.assert_allclose(_ctrl(model_path).compute([1.0, 2.0, 3.0]), [1.5, 3.5], rtol=1e-9) + + def test_compute_ignores_target(self, model_path): + ctrl = _ctrl(model_path) + a1 = ctrl.compute(np.array([1.0, 0.0, 0.0])) + a2 = ctrl.compute(np.array([1.0, 0.0, 0.0]), target=np.array([9.0, 9.0])) + np.testing.assert_allclose(a1, a2, rtol=0, atol=0) + def test_action_dtype_is_float64(self, model_path): + """The graph is f64 throughout, so the action is too (no f32 downcast).""" + action = _ctrl(model_path).compute(np.array([1.0, 2.0, 3.0])) + assert action.dtype == np.float64 + assert action.shape == (2,) -class TestPostprocessDirect: - """Exercise _postprocess branches that the stub model cannot reach.""" + def test_deterministic_action_has_no_noise_port(self, model_path): + ctrl = _ctrl(model_path) + assert ctrl.policy.noise_port is None + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 1.0, 1.0])), ctrl.compute(np.array([1.0, 1.0, 1.0]))) - def _ctrl(self, **overrides): - cfg = {"model_path": str(_build_model()), "action_space": "continuous"} - cfg.update(overrides) - return OnnxRLAdapter.from_config(cfg) - def test_discrete_reshape_flat(self): - """Multi-dimensional logits are flattened before argmax.""" - ctrl = self._ctrl(action_space="discrete", deterministic=True) - action = ctrl._postprocess(np.array([[5.0], [2.0]])) - assert action.shape == (2,) - assert action[0] == 1.0 and action[1] == 0.0 - - def test_discrete_extreme_logits(self): - """Stochastic discrete with extreme logits doesn't overflow (max subtraction).""" - ctrl = self._ctrl(action_space="discrete", deterministic=False, seed=0) - action = ctrl._postprocess(np.array([1e6, 0.0])) - assert action[0] == 1.0 and action[1] == 0.0 - - def test_discrete_deterministic_tie_picks_first(self): - ctrl = self._ctrl(action_space="discrete", deterministic=True) - action = ctrl._postprocess(np.array([1.0, 1.0])) - assert action[0] == 1.0 and action[1] == 0.0 - - def test_stochastic_logstd_clamped(self): - """Out-of-range log_std is clamped to [-10, 2].""" - ctrl = self._ctrl(action_space="stochastic", deterministic=False, seed=3) - # mean=0, log_std=100 -> clamped to 2 -> sigma ~7.39 -> sample stays ~O(20), - # not exp(100) ~ 2.7e43 - raw = np.array([0.0, 0.0, 100.0, 100.0]) - u = ctrl._postprocess(raw) - assert np.all(np.abs(u) <= 40.0) - # log_std=-100 -> clamped to -10 -> sigma ~4.5e-5 -> sample ~mean - u2 = ctrl._postprocess(np.array([5.0, 5.0, -100.0, -100.0])) - np.testing.assert_allclose(u2, [5.0, 5.0], atol=1e-3) - - def test_stochastic_mean_passthrough_astype(self): - """Deterministic stochastic postprocess returns mean as float32.""" - ctrl = self._ctrl(action_space="stochastic", deterministic=True, action_scale=2.0, action_bias=1.0) - u = ctrl._postprocess(np.array([1.0, 2.0, -3.0, -3.0])) - assert u.dtype == np.float32 - np.testing.assert_allclose(u, [3.0, 5.0]) - - def test_continuous_astype_float32(self): - ctrl = self._ctrl() - action = ctrl._postprocess(np.array([1.0, 2.0], dtype=np.float64)) - assert action.dtype == np.float32 - - def test_batched_compute_shapes(self): - """Raw ONNX output with batch dim is flattened back to (n_u,).""" - ctrl = self._ctrl(action_space="continuous") - action = ctrl.compute(np.array([1.0, 2.0, 3.0])) - assert action.shape == (2,) - np.testing.assert_allclose(action, [1.5, 3.5]) +class TestDiscrete: + def test_deterministic_one_hot(self, model_path): + ctrl = _ctrl(model_path, action_space="discrete", deterministic=True) + action = ctrl.compute(np.array([1.0, 1.0, 1.0])) # logits [1.5, 1.5] -> first max + np.testing.assert_allclose(action, [1.0, 0.0], rtol=0, atol=0) - def test_compute_list_input(self): - """compute accepts a plain Python list as state.""" - ctrl = self._ctrl(action_space="continuous") - action = ctrl.compute([1.0, 2.0, 3.0]) - np.testing.assert_allclose(action, [1.5, 3.5]) + def test_deterministic_has_no_noise_port(self, model_path): + assert _ctrl(model_path, action_space="discrete", deterministic=True).policy.noise_port is None - def test_compute_ignores_target(self): - """target is ignored for learned policies.""" - ctrl = self._ctrl(action_space="continuous") - a1 = ctrl.compute(np.array([1.0, 0.0, 0.0])) - a2 = ctrl.compute(np.array([1.0, 0.0, 0.0]), target=np.array([9.0, 9.0])) - np.testing.assert_allclose(a1, a2) + def test_sampling_uses_epsilon_and_is_one_hot(self, model_path): + ctrl = _ctrl(model_path, action_space="discrete", deterministic=False, seed=1) + assert ctrl.policy.noise_port == "epsilon" + assert ctrl.policy.gumbel + action = ctrl.compute(np.array([1.0, 1.0, 1.0])) + assert set(np.unique(action)) <= {0.0, 1.0} + assert action.sum() == 1.0 + def test_sampling_reproducible_with_seed(self, model_path): + cfg = {"model_path": str(model_path), "action_space": "discrete", "deterministic": False, "seed": 5} + a1 = OnnxRLAdapter.from_config(cfg).compute(np.array([1.0, 1.0, 1.0])) + a2 = OnnxRLAdapter.from_config(cfg).compute(np.array([1.0, 1.0, 1.0])) + np.testing.assert_allclose(a1, a2, rtol=0, atol=0) + + +class TestStochastic: + def test_deterministic_returns_mean(self, stochastic_path): + ctrl = _ctrl(stochastic_path, action_space="stochastic") + # raw = [11, 22, 33, 4] -> mean = [11, 22] + np.testing.assert_allclose(ctrl.compute(np.array([10.0, 20.0, 30.0])), [11.0, 22.0], rtol=1e-9) + assert ctrl.policy.noise_port is None + + def test_deterministic_mean_gets_scale_bias(self, stochastic_path): + ctrl = _ctrl(stochastic_path, action_space="stochastic", action_scale=2.0, action_bias=1.0) + np.testing.assert_allclose(ctrl.compute(np.array([10.0, 20.0, 30.0])), [23.0, 45.0], rtol=1e-9) + + def test_sampling_reproducible_with_seed(self, stochastic_path): + cfg = {"model_path": str(stochastic_path), "action_space": "stochastic", "deterministic": False, "seed": 3} + a1 = OnnxRLAdapter.from_config(cfg).compute(np.array([0.0, 0.0, 0.0])) + a2 = OnnxRLAdapter.from_config(cfg).compute(np.array([0.0, 0.0, 0.0])) + np.testing.assert_allclose(a1, a2, rtol=0, atol=0) + assert a1.shape == (2,) + + def test_reset_reseeds_and_repeats(self, stochastic_path): + """reset() must restore the stream, so the first draw repeats.""" + cfg = {"model_path": str(stochastic_path), "action_space": "stochastic", "deterministic": False, "seed": 7} + ctrl = OnnxRLAdapter.from_config(cfg) + first = ctrl.compute(np.array([0.0, 0.0, 0.0])) + ctrl.compute(np.array([0.0, 0.0, 0.0])) # advance the stream + ctrl.reset() + again = ctrl.compute(np.array([0.0, 0.0, 0.0])) + np.testing.assert_allclose(first, again, rtol=0, atol=0) -class TestBackendAgnostic: - """Verify the adapter converts at the ONNX boundary, not in the framework.""" + def test_sampling_is_not_standard_normal_scale(self, stochastic_path): + """The sampled action is mean + exp(clip(log_std, -10, 2)) * noise.""" + ctrl = _ctrl(stochastic_path, action_space="stochastic", deterministic=False, seed=0) + # raw = [1, 2, 3, 4]; log_std clipped to 2 -> sigma = e^2; mean = [1, 2] + u = ctrl.compute(np.array([0.0, 0.0, 0.0])) + sigma = np.exp(2.0) + assert np.all(np.abs(u - [1.0, 2.0]) <= 4.0 * sigma) + +class TestBackendAgnostic: def test_torch_backend_returns_tensor(self, model_path): torch = pytest.importorskip("torch") from shinro.utils.array_backend import TorchBackend - bk = TorchBackend(device="cpu") - ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path), "action_space": "continuous"}, backend=bk) + ctrl = OnnxRLAdapter.from_config({"model_path": str(model_path)}, backend=TorchBackend(device="cpu")) action = ctrl.compute(torch.tensor([1.0, 2.0, 3.0])) assert isinstance(action, torch.Tensor) - expected = torch.tensor([1.5, 3.5]) - torch.testing.assert_close(action, expected) - - def test_torch_backend_observations_normalized(self, model_path): - torch = pytest.importorskip("torch") - from shinro.utils.array_backend import TorchBackend - - bk = TorchBackend(device="cpu") - ctrl = OnnxRLAdapter.from_config( - {"model_path": str(model_path), "observation": {"normalize": True, "obs_mean": [1.0, 1.0, 1.0], "obs_std": [2.0, 2.0, 2.0]}}, - backend=bk, - ) - action = ctrl.compute(torch.tensor([3.0, 5.0, 1.0])) - # obs = [(3-1)/2, (5-1)/2, 0] = [1,2,0] -> [1.5, 3.5] - torch.testing.assert_close(action, torch.tensor([1.5, 3.5])) + torch.testing.assert_close(action, torch.tensor([1.5, 3.5], dtype=torch.float64)) def test_torch_backend_discrete(self, model_path): torch = pytest.importorskip("torch") from shinro.utils.array_backend import TorchBackend - bk = TorchBackend(device="cpu") - ctrl = OnnxRLAdapter.from_config( - {"model_path": str(model_path), "action_space": "discrete", "deterministic": True}, backend=bk - ) - action = ctrl.compute(torch.tensor([1.0, 1.0, 1.0])) - assert isinstance(action, torch.Tensor) - torch.testing.assert_close(action, torch.tensor([1.0, 0.0])) + cfg = {"model_path": str(model_path), "action_space": "discrete", "deterministic": True} + ctrl = OnnxRLAdapter.from_config(cfg, backend=TorchBackend(device="cpu")) + torch.testing.assert_close(ctrl.compute(torch.tensor([1.0, 1.0, 1.0])), torch.tensor([1.0, 0.0], dtype=torch.float64)) - def test_factory_passthrough_backend(self, model_path, tmp_path): - """ControllerFactory passes the backend through from_config.""" + def test_factory_passes_backend_through(self, model_path, tmp_path): torch = pytest.importorskip("torch") from shinro.factories.controller_factory import ControllerFactory from shinro.utils.array_backend import TorchBackend config = tmp_path / "rl.toml" config.write_text(f'type = "onnx_rl"\nmodel_path = "{model_path}"\naction_space = "continuous"\n') - factory = ControllerFactory(str(config)) - bk = TorchBackend(device="cpu") - ctrl = factory.create(backend=bk) - action = ctrl.compute(torch.tensor([1.0, 0.0, 0.0])) - assert isinstance(action, torch.Tensor) + ctrl = ControllerFactory(str(config)).create(backend=TorchBackend(device="cpu")) + assert isinstance(ctrl.compute(torch.tensor([1.0, 0.0, 0.0])), torch.Tensor) + + +class TestCommittedFixture: + """The checked-in toy MLP (scripts/gen_toy_onnx.py) drives the adapter.""" + + def test_continuous_matches_closed_form(self): + ctrl = OnnxRLAdapter.from_config({"model_path": str(_TOY)}) + u = ctrl.compute(np.array([1.0, 2.0, 3.0])) + np.testing.assert_allclose(u, [np.tanh(1.0) + 0.5, np.tanh(2.0) - 0.5], rtol=1e-9) + + def test_discrete_argmax(self): + ctrl = OnnxRLAdapter.from_config({"model_path": str(_TOY), "action_space": "discrete"}) + # logits [tanh(1)+0.5, tanh(2)-0.5] = [1.26, 0.46] -> the first index wins + np.testing.assert_allclose(ctrl.compute(np.array([1.0, 2.0, 3.0])), [1.0, 0.0], rtol=0, atol=0) + + def test_committed_controller_config_loads(self): + """The fixture controller TOML (CWD-relative model_path) goes through the factory.""" + from shinro.factories.controller_factory import ControllerFactory + + ctrl = ControllerFactory("tests/fixtures/configs/controllers/onnx_toy.toml").create() + u = ctrl.compute(np.array([1.0, 2.0, 3.0])) + np.testing.assert_allclose(u, [np.tanh(1.0) + 0.5, np.tanh(2.0) - 0.5], rtol=1e-9) + + +class TestCompiledModeSurface: + """Artifact-mode failures that need no compiled binary (the Zig oracle + suite covers a real ``.so`` end-to-end).""" + + def _manifest(self, artifact_dir, *, inputs=None, outputs=None): + artifact_dir.mkdir(parents=True, exist_ok=True) + manifest = { + "inputs": inputs if inputs is not None else [{"name": "state", "shape": [3], "bytes": 24}], + "outputs": outputs if outputs is not None else [{"name": "u", "shape": [2], "bytes": 16}], + "state_outputs": [], + "op_histogram": {"matmul": 1}, + "buf_len": 16, + } + (artifact_dir / "graph_data_manifest.json").write_text(json.dumps(manifest)) + return artifact_dir + + def test_expected_kernel_filename(self): + """The adapter looks for the renamed NN kernel, not the generic libbase.so.""" + assert KERNEL_FILENAME == "lib_neural_network.so" + + def test_missing_manifest(self, tmp_path): + with pytest.raises(FileNotFoundError, match="no graph manifest"): + _CompiledPolicy(tmp_path / "nope") + + def test_missing_state_port(self, tmp_path): + d = self._manifest(tmp_path / "art", inputs=[{"name": "y", "shape": [3], "bytes": 24}]) + with pytest.raises(ValueError, match="no 'state' input port"): + _CompiledPolicy(d) + + def test_missing_u_port(self, tmp_path): + d = self._manifest(tmp_path / "art", outputs=[{"name": "logits", "shape": [2], "bytes": 16}]) + with pytest.raises(ValueError, match="no 'u' output port"): + _CompiledPolicy(d) + + def test_manifest_ok_but_no_binary(self, tmp_path): + d = self._manifest(tmp_path / "art") + with pytest.raises(FileNotFoundError, match="no compiled kernel"): + _CompiledPolicy(d) + + def test_finds_kernel_under_expected_filename(self, tmp_path): + """A file at lib/ must satisfy the existence check.""" + d = self._manifest(tmp_path / "art") + (d / "lib").mkdir() + (d / "lib" / KERNEL_FILENAME).write_bytes(b"") # not a loadable .so + with pytest.raises(OSError): # got past the existence check to dlopen + _CompiledPolicy(d) + + def test_corrupt_manifest(self, tmp_path): + d = tmp_path / "art" + d.mkdir() + (d / "graph_data_manifest.json").write_text("{ not json") + with pytest.raises(ValueError, match="corrupt"): + _CompiledPolicy(d)