Source code for boundlab.interp

"""Operator-dispatch abstract interpreter for ONNX graphs."""

from __future__ import annotations

from abc import abstractmethod
from collections import Counter
from collections.abc import Callable, Iterable
from copy import deepcopy
from dataclasses import dataclass
from pathlib import Path
from typing import Any

import onnx_ir as ir
import torch

from boundlab import Expr
from boundlab.utils import TyDict
from .export import onnx_export, InterpreterModule


[docs] class OpHandler: """One implementation of an ONNX-level operator. A handler declares which operator it implements (``op``), when it applies (:meth:`condition`, typically a check on the operand classes), and the transformation itself (:meth:`handle`). For every dispatched call exactly one registered handler must be ready; ``overrides`` breaks declared ties. """ op: str overrides: list[type["OpHandler"]] = []
[docs] def override_handler(self, other: type["OpHandler"]) -> "OpHandler": """A copy of this handler that takes precedence over ``other`` when both are ready for the same call.""" if other in self.overrides: return self result = deepcopy(self) result.overrides = [*self.overrides, other] return result
[docs] def condition(self, *args: Any, **kwargs: Any) -> bool: """Whether this handler applies to these operands (default: always).""" return True
[docs] @abstractmethod def handle(self, interp: "Interpreter", *args: Any, **kwargs: Any) -> Any: """Transform the operands; sub-operations go through ``interp.<op>`` so the enclosing domain's handlers apply to them too.""" ...
[docs] def is_abstract(value: Any) -> bool: """Whether ``value`` is an abstract expression (vs. a concrete tensor).""" return isinstance(value, Expr)
def _select_handler( handlers: Iterable[OpHandler], op: str, *args: Any, **kwargs: Any ) -> OpHandler: ready = [ handler for handler in handlers if handler.op == op and handler.condition(*args, **kwargs) ] if not ready: types = tuple(type(value).__name__ for value in args) raise TypeError(f"No matching handler found for {op} with {types}, {kwargs}.") selected = [ handler for handler in ready if not any( other is not handler and any(isinstance(handler, cls) for cls in other.overrides) for other in ready ) ] if len(selected) != 1: raise RuntimeError( f"Multiple non-overridden handlers found for {op}: {selected}." ) return selected[0] def _tensor(array: Any) -> torch.Tensor: # ``from_numpy`` always lands on the host, so the default device has to be # requested explicitly here. return torch.from_numpy(array.copy()).to(torch.get_default_device()) def _attribute_value(attribute: ir.Attr) -> Any: kind = attribute.type if kind == ir.AttributeType.FLOAT: return attribute.as_float() if kind == ir.AttributeType.INT: return attribute.as_int() if kind == ir.AttributeType.STRING: return attribute.as_string() if kind == ir.AttributeType.TENSOR: return _tensor(attribute.as_tensor().numpy()) if kind == ir.AttributeType.FLOATS: return list(attribute.as_floats()) if kind == ir.AttributeType.INTS: return list(attribute.as_ints()) if kind == ir.AttributeType.STRINGS: return list(attribute.as_strings()) raise NotImplementedError(f"Unsupported ONNX attribute type {kind}.") def _constant_tensor(value: ir.Value) -> torch.Tensor: if value.const_value is None: raise ValueError(f"ONNX initializer {value.name!r} has no constant value.") return _tensor(value.const_value.numpy()) _ONNX_TO_HANDLER = { "Add": "add", "Sub": "sub", "Mul": "mul", "Div": "div", "Neg": "neg", "MatMul": "matmul", "Gemm": "gemm", "Relu": "relu", "Exp": "exp", "Tanh": "tanh", "Reciprocal": "reciprocal", "Softmax": "softmax", "Max": "max", "Reshape": "reshape", "Transpose": "transpose", "Expand": "broadcast_to", "ReduceSum": "reduce_sum", "ReduceMean": "reduce_mean", "Unsqueeze": "unsqueeze", "Squeeze": "squeeze", "Identity": "identity", "Cast": "cast", "MarkedIdentity": "marked_identity", }
[docs] class Interpreter(TyDict[OpHandler]): """Evaluate an ONNX graph over concrete tensors and abstract expressions. A collection of :class:`OpHandler`\\ s keyed by handler class. Calling the interpreter on a model (module, callable, ONNX IR, or path) returns a function that replays the graph node by node: each node's operator name dispatches to the unique handler whose ``condition`` accepts the operands, and ``after_each`` post-processes every node result (domains use it to normalize components). ``interpreter.<op>(...)`` exposes the same dispatch directly, which is how handlers compose. """ handler_called: list[OpHandler] | None after_each: Callable[[Any, str], Any] = staticmethod(lambda value, name: value)
[docs] def __init__( self, *handlers: OpHandler | "Interpreter", after_each: Callable[[Any, str], Any] | None = None, ): flattened = [ item for handler in handlers for item in (handler.values() if isinstance(handler, Interpreter) else (handler,)) ] super().__init__(*flattened) self.handler_called = None if after_each is not None: self.after_each = after_each
def __getattr__(self, name: str) -> Callable[..., Any]: if name.startswith("_"): raise AttributeError(name) def handle(*args: Any, **kwargs: Any) -> Any: handler = _select_handler(self.values(), name, *args, **kwargs) if self.handler_called is not None: self.handler_called.append(handler) return handler.handle(self, *args, **kwargs) return handle def __enter__(self): self.handler_called = [] return self.handler_called def __exit__(self, exc_type, exc_value, traceback): self.handler_called = None def _dispatch_node( self, node: ir.Node, args: list[Any], prepared_attributes: dict[str, Any] | None = None, ) -> Any: attributes = ( dict(prepared_attributes) if prepared_attributes is not None else { name: _attribute_value(attribute) for name, attribute in node.attributes.items() } ) op = _ONNX_TO_HANDLER.get(node.op_type) if op is None: raise NotImplementedError( f"ONNX operator {node.op_type!r} ({node.name!r}) is not supported." ) if node.op_type == "Reshape": shape = args.pop(1) if "new_sizes" not in attributes: attributes["new_sizes"] = tuple(int(x) for x in shape.tolist()) elif node.op_type == "Transpose": if "permutation" not in attributes: attributes["permutation"] = tuple( attributes.pop("perm", range(args[0].ndim - 1, -1, -1)) ) elif node.op_type == "Expand": shape = args.pop(1) if "shape" not in attributes: attributes["shape"] = tuple(int(x) for x in shape.tolist()) elif node.op_type in ("ReduceSum", "ReduceMean"): axes = args.pop(1) if len(args) > 1 and args[1] is not None else None if "axes" not in attributes: axes = axes if axes is not None else attributes.pop("axes", None) attributes["axes"] = ( None if axes is None else tuple(int(x) for x in axes.tolist()) ) attributes["keepdims"] = bool(attributes.get("keepdims", 1)) elif node.op_type in ("Unsqueeze", "Squeeze"): axes = args.pop(1) if len(args) > 1 else None if "axes" not in attributes: axes = axes if axes is not None else attributes.pop("axes", None) attributes["axes"] = ( None if axes is None else tuple(int(x) for x in axes.tolist()) ) elif node.op_type == "Cast": attributes["to"] = int(attributes["to"]) return self.__getattr__(op)(*args, **attributes)
[docs] def __call__( self, model: ir.Model | torch.onnx.ONNXProgram | Callable[..., Any] | str | Path, verbose: bool = False, ) -> Callable[..., Any]: """Build a callable that evaluates ``model`` over abstract inputs.""" if callable(model) and not isinstance(model, torch.onnx.ONNXProgram): function = model def export_then_interpret(*inputs: Any) -> Any: examples = tuple( torch.zeros(value.shape, dtype=value.dtype) if isinstance(value, Expr) else value for value in inputs ) graph = onnx_export(function, examples) return self(graph, verbose=verbose)(*inputs) return export_then_interpret if isinstance(model, torch.onnx.ONNXProgram): model = model.model if isinstance(model, (str, Path)): model = ir.load(str(model)) if not isinstance(model, ir.Model): raise TypeError(f"model must be ONNX IR, got {type(model)}.") initializers = { value.name: _constant_tensor(value) for value in model.graph.initializers.values() } input_names = [ value.name for value in model.graph.inputs if value.name not in initializers ] output_names = [value.name for value in model.graph.outputs] def constant_input(node: ir.Node, index: int) -> torch.Tensor | None: if index >= len(node.inputs) or node.inputs[index] is None: return None value = node.inputs[index] assert value is not None if value.name in initializers: return initializers[value.name] if value.const_value is not None: return _constant_tensor(value) return None prepared_nodes: list[tuple[ir.Node, dict[str, Any]]] = [] for node in model.graph: attributes = { name: _attribute_value(attribute) for name, attribute in node.attributes.items() } if node.op_type == "Reshape": shape = constant_input(node, 1) if shape is not None: attributes["new_sizes"] = tuple(int(x) for x in shape.tolist()) elif node.op_type == "Transpose": attributes["permutation"] = tuple( attributes.pop( "perm", range(len(node.inputs[0].shape) - 1, -1, -1), ) ) elif node.op_type == "Expand": shape = constant_input(node, 1) if shape is not None: attributes["shape"] = tuple(int(x) for x in shape.tolist()) elif node.op_type in ("ReduceSum", "ReduceMean", "Unsqueeze", "Squeeze"): axes = constant_input(node, 1) if axes is not None: attributes["axes"] = tuple(int(x) for x in axes.tolist()) elif "axes" in attributes: attributes["axes"] = tuple(attributes["axes"]) elif node.op_type == "Cast": attributes["to"] = int(attributes["to"]) prepared_nodes.append((node, attributes)) def interpret(*inputs: Any) -> Any: if len(inputs) != len(input_names): raise ValueError( f"Expected {len(input_names)} graph inputs, got {len(inputs)}." ) env = { name: self.after_each(value, f"input_{index}") for index, (name, value) in enumerate(zip(input_names, inputs)) } counters: Counter[str] = Counter() for node, prepared_attributes in prepared_nodes: args: list[Any] = [] for value in node.inputs: if value is None or not value.name: args.append(None) elif value.name in env: args.append(env[value.name]) elif value.name in initializers: args.append(initializers[value.name]) elif value.const_value is not None: args.append(_constant_tensor(value)) else: raise KeyError( f"Missing input {value.name!r} for {node.op_type} node." ) with self as called: result = self._dispatch_node(node, args, prepared_attributes) name = f"{node.op_type}_{counters[node.op_type]}" counters[node.op_type] += 1 result = self.after_each(result, name) results = result if isinstance(result, (tuple, list)) else (result,) outputs = [value for value in node.outputs if value is not None and value.name] if len(results) != len(outputs): raise RuntimeError( f"{node.op_type} returned {len(results)} outputs for {len(outputs)} values." ) for output, value in zip(outputs, results): env[output.name] = value if verbose: handlers = " -> ".join( f"{handler.__class__.__module__}.{handler.__class__.__name__}" for handler in called ) detail = result.torch_print().render() if isinstance(result, Expr) else "" print(f"{name}: {node.op_type} -> {handlers}{detail}") outputs = [env[name] for name in output_names] return outputs[0] if len(outputs) == 1 else tuple(outputs) return interpret
[docs] def handler_fn( op: str, condition: Callable[..., bool] | None = None, ) -> Callable[[Callable[..., Any]], type[OpHandler]]: """Wrap a plain function as an :class:`OpHandler` class for operator ``op``. The stateless common case: ``function(*args, **attrs)`` becomes the handler's ``handle`` (without the interpreter argument), and an optional ``condition`` predicate gates it. """ op_name = op def decorator(function: Callable[..., Any]) -> type[OpHandler]: @dataclass class Handler(OpHandler): op: str = op_name def condition(self, *args: Any, **kwargs: Any) -> bool: return condition(*args, **kwargs) if condition else True def handle(self, interp: Interpreter, *args: Any, **kwargs: Any) -> Any: del interp return function(*args, **kwargs) Handler.__name__ = function.__name__ Handler.__qualname__ = function.__qualname__ Handler.__doc__ = function.__doc__ or f"Handler for ``{op_name}``." return Handler return decorator
__all__ = [ "base", "export", "Interpreter", "OpHandler", "handler_fn", "is_abstract", "onnx_export", "InterpreterModule", ]