From 81a86a50e0ebfd73b14e6ad03cb242c3621024cb Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 15:25:44 -0400 Subject: [PATCH 01/11] chore(deps): declare onnx (not onnxruntime) for the onnx-rl extra MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The ONNX→Zig policy path reads the graph with `onnx` at compile time and deploys a dependency-free `.so`, so the extra no longer needs the onnxruntime inference engine. `onnx` also joins the dev extra, which fixes the ONNX adapter tests silently skipping under a stock `.[onnx-rl]` install (they guarded on both packages, but only onnxruntime was declared). The checked-in adapter still lazy-imports onnxruntime on its eager path; that path is replaced in a following commit. --- pyproject.toml | 3 ++- src/shinro/configs/controllers/onnx_rl.toml | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) 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/src/shinro/configs/controllers/onnx_rl.toml b/src/shinro/configs/controllers/onnx_rl.toml index 0b34a90..ce75d7c 100644 --- a/src/shinro/configs/controllers/onnx_rl.toml +++ b/src/shinro/configs/controllers/onnx_rl.toml @@ -1,6 +1,6 @@ # FILE: configs/controllers/onnx_rl.toml # ONNX RL policy adapter — wraps any ONNX-exported RL policy as a Controller. -# Requires: pip install onnxruntime +# Requires: pip install "shinro[onnx-rl]" (the `onnx` package; no runtime engine) # # Export the actor network from your RL framework (sb3, RLlib, CleanRL, # custom torch, JAX, ...) to policy.onnx, then point model_path at it. From c87701ca3c72e018bf0a9e7ddd6b86720c8370a9 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 15:25:48 -0400 Subject: [PATCH 02/11] feat(codegen): import ONNX policies as shinro graphs Add shinro.codegen.onnx_import: translate onnx.load(...).graph directly into a memoryless ComposedGraph, bypassing the tracer entirely (ONNX is already a dataflow graph, so there is nothing to intercept). Supports Gemm (alpha/beta/ transB), MatMul and Add, Relu/Tanh, and composes Sigmoid from exp/neg/add/div so no new VM op is needed. The observation encoder (integer index selection, mean/std normalization, clipping) is folded in as arithmetic on baked constants, leaving the raw plant state as the only input port. Every emitted node is evaluated eagerly with the ops registry's numpy handler and that result's shape becomes the node's declared shape, so interpreter and lowering semantics cannot drift. Nodes that do not feed the declared output are ignored; unsupported ops/attributes, multi-input models, and batched outputs fail loudly with actionable messages. Verified: 24 unit tests (pytest tests/unit/test_onnx_import.py), ruff clean, pyrefly 0 errors. --- src/shinro/codegen/onnx_import.py | 537 ++++++++++++++++++++++++++++++ tests/unit/test_onnx_import.py | 448 +++++++++++++++++++++++++ 2 files changed, 985 insertions(+) create mode 100644 src/shinro/codegen/onnx_import.py create mode 100644 tests/unit/test_onnx_import.py diff --git a/src/shinro/codegen/onnx_import.py b/src/shinro/codegen/onnx_import.py new file mode 100644 index 0000000..60bbeef --- /dev/null +++ b/src/shinro/codegen/onnx_import.py @@ -0,0 +1,537 @@ +"""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``. + +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" + +#: 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, + 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. + 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.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] = {} + + # ── orchestration ───────────────────────────────────────────────────── + + def build(self) -> ComposedGraph: + """Parse the model, emit the graph, and return the composed result. + + Returns: + A :class:`ComposedGraph` with one input port (``state``), one output + port (``u``), and no recurrent state. + + Raises: + ValueError: On a malformed model/config, a multi-input policy, an + unresolvable tensor, or a batched (leading dim ≠ 1) output. + NotImplementedError: On an unsupported ONNX op or attribute. + """ + 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]) + self.emit("output", [action_id], name=OUTPUT_PORT) + return ComposedGraph(graph=self.g, inputs=[STATE_PORT], 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 un-inferable observation dimension, or + out-of-range ``state_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``). + + 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" + ) + + +def import_onnx_policy( + model_path: str, + *, + n_x: int | None = None, + obs_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``. + 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``. + + 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, 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/tests/unit/test_onnx_import.py b/tests/unit/test_onnx_import.py new file mode 100644 index 0000000..5ca34e2 --- /dev/null +++ b/tests/unit/test_onnx_import.py @@ -0,0 +1,448 @@ +"""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 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 _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)) + + 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) + + 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 "add" not in _op_names(cg) + + 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]}) + + +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) From 33580c12ad3890e779a67f3da37395657b253d19 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 15:46:34 -0400 Subject: [PATCH 03/11] feat(codegen): bake the ONNX policy action space with an epsilon port MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add the action-space surface to the ONNX importer: continuous applies scale/bias, discrete emits an argmax one-hot, and stochastic splits [mean; log_std] and either returns the mean or adds exp(clip(log_std, -10, 2)) * epsilon — mirroring the old runtime _postprocess. A non-deterministic discrete/stochastic policy gains an `epsilon` C-ABI input port: the host supplies the noise (Gumbel for discrete, standard normal for stochastic) and the kernel does only the arithmetic, matching MPPI's port and keeping RNG on the host. Discrete sampling uses the Gumbel-max trick, argmax(logits + g) + one_hot, so no softmax/max VM op is needed. Single-sided action clips are rejected: the lowerer writes floats as Zig hex literals and `inf` is not a Zig identifier, so a bound defaulted to ±inf would fail the build. Gemm's alpha/beta multipliers and the action scale/bias are emitted unconditionally (defaults included) for uniform lowering. Verified: 38 importer unit tests, including a 2000-draw Gumbel-max vs softmax distribution check; all six action-space variants lower via lower_zig into isolated paths; ruff + pyrefly clean. --- src/shinro/codegen/onnx_import.py | 173 ++++++++++++++++++++++++++++-- tests/unit/test_onnx_import.py | 150 +++++++++++++++++++++++++- 2 files changed, 313 insertions(+), 10 deletions(-) diff --git a/src/shinro/codegen/onnx_import.py b/src/shinro/codegen/onnx_import.py index 60bbeef..5988a83 100644 --- a/src/shinro/codegen/onnx_import.py +++ b/src/shinro/codegen/onnx_import.py @@ -36,6 +36,14 @@ 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 @@ -59,7 +67,11 @@ 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"}) #: 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. @@ -95,6 +107,7 @@ def __init__( *, 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`. @@ -106,12 +119,15 @@ def __init__( 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() @@ -120,20 +136,31 @@ def __init__( 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 one input port (``state``), one output - port (``u``), and no recurrent state. + 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, or a batched (leading dim ≠ 1) output. + 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} @@ -158,8 +185,11 @@ def build(self) -> ComposedGraph: raise ValueError(f"ONNX output {resolved_output!r} was not produced by any reachable node") action_id = self.flatten_output(self.tensors[resolved_output]) - self.emit("output", [action_id], name=OUTPUT_PORT) - return ComposedGraph(graph=self.g, inputs=[STATE_PORT], outputs=[OUTPUT_PORT]) + 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``. @@ -348,7 +378,10 @@ def emit_gemm(self, inputs: list[str], attrs: dict[str, Any]) -> int: ``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``). + 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``. @@ -405,12 +438,130 @@ def flatten_output(self, node_id: int) -> int: 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. @@ -429,18 +580,24 @@ def import_onnx_policy( ``[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``. + 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, output_name=output_name).build() + 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]: diff --git a/tests/unit/test_onnx_import.py b/tests/unit/test_onnx_import.py index 5ca34e2..19b0ca4 100644 --- a/tests/unit/test_onnx_import.py +++ b/tests/unit/test_onnx_import.py @@ -11,7 +11,7 @@ import pytest from shinro.codegen.interpreter import interpret -from shinro.codegen.onnx_import import OUTPUT_PORT, STATE_PORT, import_onnx_policy +from shinro.codegen.onnx_import import EPSILON_PORT, OUTPUT_PORT, STATE_PORT, import_onnx_policy onnx = pytest.importorskip("onnx") @@ -48,6 +48,15 @@ def _run(cg, x): 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] @@ -111,6 +120,7 @@ def test_torch_layout_transB(self, tmp_path): 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 @@ -143,6 +153,8 @@ def test_alpha_beta(self, 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 @@ -157,7 +169,7 @@ def test_no_bias(self, 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 "add" not in _op_names(cg) + 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 @@ -446,3 +458,137 @@ def test_unknown_tensor_rejected(self, 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}, + ) From ad81d05794a3c8797199ba336ef22bd692b4c1ce Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 17:21:31 -0400 Subject: [PATCH 04/11] feat(controllers): run ONNX policies from a graph or a compiled kernel MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Rewrite the onnx_rl adapter on top of the ONNX→shinro importer and drop onnxruntime entirely. Two interchangeable backends: `model_path` imports the ONNX graph and runs it with the pure-numpy interpreter (eager, no build step), `artifact_dir` dlopens the scenario's compiled kernel under lib/lib_neural_network.so and drives it through the shinro_step C ABI, reading the port layout from the graph manifest so the .onnx file is not needed at deploy time. Sampling noise (Gumbel for discrete, standard normal for stochastic) is drawn host-side from a seeded RNG and fed to the epsilon port, keeping the kernel pure arithmetic; reset() makes a run reproducible. Add the frozen OnnxRLConfig dataclass (silences the registry's missing-Config warning for onnx_rl), strict-parse the config, and reject unknown observation keys so a typo cannot silently drop normalization. Behavior changes vs the runtime implementation: actions are float64 (the graph is f64 throughout), single-sided action_clip is rejected (a ±inf bound cannot be lowered), normalize without stats raises ValueError, and observation.add_batch_dim is a no-op (batch-1 is implicit in the 1-D ports). Verified: 81 unit tests across tests/unit/test_onnx_{rl_adapter,import}.py; full suite 1225 passed / 5 skipped; ruff + pyrefly clean. --- src/shinro/codegen/onnx_import.py | 11 +- src/shinro/controllers/onnx_rl_adapter.py | 523 ++++++++++++--------- tests/unit/test_onnx_import.py | 14 + tests/unit/test_onnx_rl_adapter.py | 535 +++++++++++----------- 4 files changed, 579 insertions(+), 504 deletions(-) diff --git a/src/shinro/codegen/onnx_import.py b/src/shinro/codegen/onnx_import.py index 5988a83..810f4ce 100644 --- a/src/shinro/codegen/onnx_import.py +++ b/src/shinro/codegen/onnx_import.py @@ -72,6 +72,9 @@ #: 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. @@ -207,9 +210,13 @@ def _resolve_ports(self, graph_proto: Any) -> tuple[str, list[int], int]: Raises: ValueError: On multiple graph inputs, an ``input_name`` override - that does not match, an un-inferable observation dimension, or - out-of-range ``state_keys``. + 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( 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/tests/unit/test_onnx_import.py b/tests/unit/test_onnx_import.py index 19b0ca4..cdb972e 100644 --- a/tests/unit/test_onnx_import.py +++ b/tests/unit/test_onnx_import.py @@ -368,6 +368,20 @@ def test_state_keys_out_of_range_rejected(self, 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): diff --git a/tests/unit/test_onnx_rl_adapter.py b/tests/unit/test_onnx_rl_adapter.py index f3a5130..a048ed0 100644 --- a/tests/unit/test_onnx_rl_adapter.py +++ b/tests/unit/test_onnx_rl_adapter.py @@ -1,345 +1,328 @@ -"""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 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") -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_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_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_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) + 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: - """Verify the adapter converts at the ONNX boundary, not in the framework.""" +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 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) From 7a1f09cbec2adbb5afe9a3235bae46aa473230cd Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 19:00:30 -0400 Subject: [PATCH 05/11] feat(codegen): compile standalone ONNX policies and name the artifact MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Make `make compile` reach a policy-only scenario end to end: - scenario_gen: a scenario may omit [estimator] when its controller is onnx_rl (any other type still requires one). Such a scenario bypasses build_composed_graph and lowers the importer's graph directly, with the .onnx file's sha256 pinned in the manifest provenance. - [compile].artifact_name selects the installed kernel stem. build.zig installs through an explicit sub-path (Zig's addLibrary would otherwise prefix "lib", turning lib_neural_network into liblib_neural_network); oracle.load_so, stamp, scenario_build and cli thread the name through. The default "libbase" keeps every existing artifact and test byte-identical. - build.zig's build-time readFile cap was 1 MiB, so any graph past ~50k baked constants failed with a misleading "graph_data.zig has no has_solve_qp flag; regenerate it" panic. Raised to 256 MiB. Add a committed toy policy so the ONNX path has a stable subject: tests/fixtures/models/toy_mlp.onnx (26-param 3->4->2 tanh MLP, generated by scripts/gen_toy_onnx.py), its controller config, and a policy-only scenario. The adapter/import tests and the compile e2e use it, a drift guard regenerates it through the script's CLI and compares the graph, and the scenario template documents the new [compile] key (the template drift guard forces that). Add scripts/measure_onnx_policy_scale.py — the cold-build sweep of parameter count against artifact size, graph source size, VM buffers and compile time. Its numbers and the comptime scaling wall (the VM is a ~10^5-parameter design; >=830k params trips @setEvalBranchQuota) are written up in the lab note. Verified: 100 onnx/compile tests including every zig-gated e2e; make test-zig 77 passed / 2 skipped; make test 1233 passed / 5 skipped; ruff + pyrefly clean. --- lab-notes/daily/2026-09-17.md | 103 ++++++++ scripts/gen_toy_onnx.py | 82 +++++++ scripts/measure_onnx_policy_scale.py | 229 ++++++++++++++++++ src/shinro/codegen/cli.py | 2 + src/shinro/codegen/oracle.py | 8 +- src/shinro/codegen/scenario_build.py | 25 +- src/shinro/codegen/scenario_gen.py | 140 +++++++++-- src/shinro/codegen/stamp.py | 23 +- src/shinro/configs/scenarios/_template.toml | 1 + src/shinro/runtime/build.zig | 33 ++- .../configs/controllers/onnx_toy.toml | 13 + .../configs/scenarios/toy_mlp_policy.toml | 17 ++ tests/fixtures/models/toy_mlp.onnx | Bin 0 -> 316 bytes tests/test_compile_scenario.py | 110 ++++++++- tests/unit/test_onnx_rl_adapter.py | 26 ++ 15 files changed, 759 insertions(+), 53 deletions(-) create mode 100644 scripts/gen_toy_onnx.py create mode 100644 scripts/measure_onnx_policy_scale.py create mode 100644 tests/fixtures/configs/controllers/onnx_toy.toml create mode 100644 tests/fixtures/configs/scenarios/toy_mlp_policy.toml create mode 100644 tests/fixtures/models/toy_mlp.onnx diff --git a/lab-notes/daily/2026-09-17.md b/lab-notes/daily/2026-09-17.md index d11c28c..b813464 100644 --- a/lab-notes/daily/2026-09-17.md +++ b/lab-notes/daily/2026-09-17.md @@ -42,3 +42,106 @@ 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`. + +**Not yet done.** Step 6 — a committed Zig-oracle test for the ONNX policy graph +(including the epsilon/discrete path); the compiled-vs-eager parity above was +verified manually, and steps 1–5 are still uncommitted. 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/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/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/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/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 0000000000000000000000000000000000000000..864e6b9e574d64b39df64dbcb3be588ae4181e55 GIT binary patch literal 316 zcmdHxn4!eto|>Dh#mmK3Qk0li>?FasfRTxdhl?>o zh%r%#B_uH~gG+;pF%hW62&lzKh%GU>Br`t`t3Dz2lKjf}+?)a}8x9r@W&uVe0|tf$ zdj=GMjSmuoh&utDrX|I}0@N#jq7BG#0yNgO~m rAXn@MnqU+q3=RV!9xg5pMj 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/unit/test_onnx_rl_adapter.py b/tests/unit/test_onnx_rl_adapter.py index a048ed0..e330275 100644 --- a/tests/unit/test_onnx_rl_adapter.py +++ b/tests/unit/test_onnx_rl_adapter.py @@ -10,6 +10,7 @@ import dataclasses import json +from pathlib import Path import numpy as np import pytest @@ -18,6 +19,9 @@ onnx = pytest.importorskip("onnx") +#: 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 _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.""" @@ -273,6 +277,28 @@ def test_factory_passes_backend_through(self, model_path, tmp_path): 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).""" From 26c928c3097afecd6a6ffc3c4f9a7eeb4987f979 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 19:14:59 -0400 Subject: [PATCH 06/11] test(zig): add the ONNX policy lowering oracle MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit TestOnnxPolicyOracle imports the committed toy policy (tests/fixtures/models/ toy_mlp.onnx) and compiles its four baked variants — continuous, discrete-deterministic, discrete-sampling and stochastic-sampling — into tmp-path kernels, then asserts .so == interpret() on seeded random states (exact for the continuous/one-hot paths, <=1e-14 for the noise paths) plus the closed form, 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, and it proves the importer's output — baked encoder, transposed Gemms, composed activations, action post-processing — compiles bit-for-bit. The shared src/shinro/runtime/graph_data.zig is never touched. Verified: make test-zig 81 passed / 2 skipped; make test 1241 passed / 5 skipped; make lint clean. --- tests/test_zig_lowering.py | 128 +++++++++++++++++++++++++++++++++++++ 1 file changed, 128 insertions(+) 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) From 0a3c84274cc0060c3b3473eaa08d2e2988138d59 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 19:14:59 -0400 Subject: [PATCH 07/11] docs: document the ONNX policy backends configs/controllers/onnx_rl.toml now explains the eager (model_path) and compiled (artifact_dir) backends, the policy-only `make compile` command, the epsilon / host-noise contract, and that add_batch_dim is legacy; the stale `--controller onnx_rl` demo line is removed (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. The lab note carries the step 6-7 write-up and the ONNX shift's scale-sweep results. Verified: the updated config strict-parses into OnnxRLConfig; ruff + pyrefly clean. --- docs/components.md | 5 +++ lab-notes/daily/2026-09-17.md | 20 +++++++-- src/shinro/configs/controllers/onnx_rl.toml | 48 +++++++++++++++------ 3 files changed, 56 insertions(+), 17 deletions(-) diff --git a/docs/components.md b/docs/components.md index 82b629f..54aabea 100644 --- a/docs/components.md +++ b/docs/components.md @@ -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 diff --git a/lab-notes/daily/2026-09-17.md b/lab-notes/daily/2026-09-17.md index b813464..de1194e 100644 --- a/lab-notes/daily/2026-09-17.md +++ b/lab-notes/daily/2026-09-17.md @@ -74,6 +74,7 @@ host. Single-sided action clip is rejected: a `±inf` bound cannot be emitted 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. @@ -142,6 +143,19 @@ agree to 1.1e-16. `tests/fixtures/configs/{controllers/onnx_toy.toml,scenarios/toy_mlp_policy.toml}`; tooling `scripts/gen_toy_onnx.py`, `scripts/measure_onnx_policy_scale.py`. -**Not yet done.** Step 6 — a committed Zig-oracle test for the ONNX policy graph -(including the epsilon/discrete path); the compiled-vs-eager parity above was -verified manually, and steps 1–5 are still uncommitted. +**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. diff --git a/src/shinro/configs/controllers/onnx_rl.toml b/src/shinro/configs/controllers/onnx_rl.toml index ce75d7c..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 "shinro[onnx-rl]" (the `onnx` package; no runtime engine) +# 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 From 875181a3db9257fe77e2e95c41c5abfbefcf5fd6 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 22:24:52 -0400 Subject: [PATCH 08/11] perf(runtime): de-inline the VM's element loops runtime/lower.zig unrolled both the outer node table and every per-element loop with `inline for` (38 sites, 0 runtime `for`). The outer unroll is cheap (~1300 nodes); the inner ones emit one statement per element, so shinro_step became a single function with ~buf_len statements (100k-340k). Zig's comptime evaluator and LLVM then processed that giant straight-line function: MPPI production builds at N=200/K=15 took 411-1306 s and produced 1.9-6.1 MiB .so files whose .text was ~2x buf_bytes. Keep the outer `inline for (g.nodes)` (op tag and shapes stay comptime, buffer offsets stay comptime constants) and make all 34 inner element loops runtime `for` loops over those comptime-known sizes. ew2/bcast_flat stay inline so shapes and the op tag constant-fold. Result (MPPI, N=200/K=15, ReleaseFast, real plants): - 5.2-11.5x smaller .so - 35-106x faster builds (8-kernel matrix ~90 s total) - 1.4-2.1x faster ticks (the unrolled function was thrashing I-cache) ONNX MLP scale sweep: the >=830k-param comptime branch-quota wall is gone; all six archs (up to 1.88M params) now compile and oracle-verify. Verified: make test 1241 passed/5 skipped (unchanged), make test-zig 81/2, make lint clean, all SMC/MPPI/ONNX and closed-loop oracles pass. Shipped graph_data.zig untouched. --- lab-notes/daily/2026-09-17.md | 110 ++++++++++++++++++++++++++++++++++ src/shinro/runtime/lower.zig | 94 ++++++++++++++++------------- 2 files changed, 162 insertions(+), 42 deletions(-) diff --git a/lab-notes/daily/2026-09-17.md b/lab-notes/daily/2026-09-17.md index de1194e..ea1085f 100644 --- a/lab-notes/daily/2026-09-17.md +++ b/lab-notes/daily/2026-09-17.md @@ -159,3 +159,113 @@ 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. diff --git a/src/shinro/runtime/lower.zig b/src/shinro/runtime/lower.zig index 3b123f5..f2bc1ea 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 +// stack 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): @@ -42,8 +48,11 @@ 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 stack buffer — 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`. /// /// Args: /// inputs: Flat buffer of this tick's host inputs. @@ -59,17 +68,17 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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); 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 => { @@ -80,15 +89,15 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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), @@ -100,16 +109,16 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) .pow => ew2(g.nodes[0..], node, i, &buf, .pow), .neg => { const s = node_input(g.nodes[0..], node, &buf); - inline for (0..node.rows * node.cols) |j| out[j] = -s[j]; + 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]); + 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| { + 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; } }, @@ -119,22 +128,22 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // 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 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]; + 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| { + 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]); } }, @@ -145,8 +154,8 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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)]; @@ -158,7 +167,7 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) .any => { const s = node_input(g.nodes[0..], node, &buf); 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; @@ -172,32 +181,32 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // 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]; + for (0..node.rows * node.cols) |j| out[j] = s[j]; }, .tanh => { const s = node_input(g.nodes[0..], node, &buf); 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 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 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 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 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); @@ -217,29 +226,29 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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 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 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 +257,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| { + for (node.inputs, 0..) |inp_idx, row| { const src = node_input_at(g.nodes[0..], inp_idx, &buf); 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 @@ -307,8 +316,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 +332,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) { From 394d36e42dac50f5d056d9abe16e45c3aa1839e0 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 23:11:22 -0400 Subject: [PATCH 09/11] perf(runtime): move the VM step workspace off the stack MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit shinro_step declared its per-node scratch as a stack local, `var buf: [g.buf_len]f64`. For large graphs that reserves buf_bytes on the caller's stack — 20.8 MiB at ~1.35M parameters — overflowing the default 16 MiB stack when the kernel is called. The compiler was fine; the call segfaulted. Move it to a file-scope `workspace` in .bss, sized at compile time from g.buf_len and addressed with the same comptime offsets (the helpers keep their `buf` pointer parameter). The graph writes every slot before it is read, so no initialisation is needed. Documented as not reentrant / thread-safe — one caller, one tick at a time — matching the QP path's statically-allocated solver global. Result: a 1.35M-param policy (20.8 MiB workspace) builds and oracle-verifies under the default 16 MiB stack; its .so is ~15 KB of .text plus the .bss workspace, still with no dynamic allocation. Tick latency unchanged within measurement noise. Verified: make test-zig 81 passed / 2 skipped, make test 1241 passed / 5 skipped. --- lab-notes/daily/2026-09-17.md | 35 +++++++++++++ src/shinro/runtime/lower.zig | 95 ++++++++++++++++++++--------------- 2 files changed, 90 insertions(+), 40 deletions(-) diff --git a/lab-notes/daily/2026-09-17.md b/lab-notes/daily/2026-09-17.md index ea1085f..a1424a1 100644 --- a/lab-notes/daily/2026-09-17.md +++ b/lab-notes/daily/2026-09-17.md @@ -269,3 +269,38 @@ 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/src/shinro/runtime/lower.zig b/src/shinro/runtime/lower.zig index f2bc1ea..5d06c4d 100644 --- a/src/shinro/runtime/lower.zig +++ b/src/shinro/runtime/lower.zig @@ -8,7 +8,7 @@ // 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 -// stack buffer — no heap, no runtime op dispatch. The element loops *inside* +// 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 @@ -36,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 @@ -49,10 +62,13 @@ const qp = if (g.has_solve_qp) @import("qp.zig") else struct {}; /// /// The `inline for` over `g.nodes` unrolls the whole node table at compile /// time, so every node's op and shape are comptime constants and every array -/// is a fixed-size slice of the single stack buffer — 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`. +/// 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. @@ -61,10 +77,9 @@ 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 => { @@ -74,7 +89,7 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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) { for (0..node.rows * node.cols) |j| outputs[g.output_offsets[node.aux] + j] = src[j]; } else { @@ -82,8 +97,8 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) } }, .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) { @@ -100,30 +115,30 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) 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); + 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); + 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); + 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 @@ -133,24 +148,24 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) } }, .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); for (0..node.rows * node.cols) |j| out[j] = r[j]; }, .reshape => { - const s = node_input(g.nodes[0..], node, &buf); + 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); + 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]]; @@ -165,7 +180,7 @@ 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; for (0..g.nodes[node.inputs[0]].rows * g.nodes[node.inputs[0]].cols) |j| { if (s[j] != 0.0) found = true; @@ -180,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); + 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); 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); 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); 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); 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); 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); @@ -219,7 +234,7 @@ 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); @@ -233,13 +248,13 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) } }, .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); 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. @@ -258,7 +273,7 @@ export fn shinro_step(inputs: [*]const f64, outputs: [*]f64, state_out: [*]f64) // output flat length is n_inputs * in_len == node.rows * node.cols. .stack => { for (node.inputs, 0..) |inp_idx, row| { - const src = node_input_at(g.nodes[0..], inp_idx, &buf); + const src = node_input_at(g.nodes[0..], inp_idx, &workspace); const in_len = g.nodes[inp_idx].rows * g.nodes[inp_idx].cols; for (0..in_len) |j| out[row * in_len + j] = src[j]; } @@ -284,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); }, } From 75c6cfa4006fb423e6400717c2ce01ad663799c9 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 23:12:22 -0400 Subject: [PATCH 10/11] chore(tools): add tick-latency benchmark harness scripts/bench_tick.py dlopens a compiled kernel, drives shinro_step with seeded inputs and recurrent state feedback, and reports ns/tick min/median/p99. It is the A/B yardstick for the VM lowering changes measured in the lab notes. --- scripts/bench_tick.py | 226 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 226 insertions(+) create mode 100644 scripts/bench_tick.py 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()) From 22355ea3564cf2afbaa09cdf1adcccc88aa6bda8 Mon Sep 17 00:00:00 2001 From: Adil Faisal Date: Thu, 17 Sep 2026 23:12:22 -0400 Subject: [PATCH 11/11] style(docs): normalize markdown table separators MD060 table-separator spacing; content unchanged. --- docs/components.md | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/docs/components.md b/docs/components.md index 54aabea..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` | @@ -43,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` | @@ -71,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` | @@ -92,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 | @@ -105,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 |