latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
Source code for tensorplay.onnx
"""ONNX export for TensorPlay.
Export happens in two stages: capture a model into an
:class:`tensorplay.export.ExportedProgram` (via :func:`tensorplay.export`), then
translate the resulting graph into an ONNX ``ModelProto``.
The translation runs the captured graph once on the recorded example inputs so
every intermediate value carries a shape and dtype. Handlers registered in
:mod:`tensorplay.onnx._composite_ops` use that metadata to choose between
lowerings that differ only by rank or shape (``Gemm`` vs ``MatMul``,
``GlobalAveragePool`` vs ``AveragePool``, ``perm`` vectors, ...).
"""
from __future__ import annotations
import warnings
from typing import Any, Mapping, Sequence
from onnx import TensorProto, checker, helper, numpy_helper, shape_inference
from ..export import ExportedProgram, export as tp_export
from ..graph._utils import _iter_nodes
from . import _external_data, _passes, errors, testing, utils, verification
from ._composite_ops import (
GraphBuilder,
OpContext,
Value,
lookup_function_handler,
lookup_method_handler,
)
from ._type_mapping import (
_dtype_to_numpy,
_np_dtype_to_onnx,
_size_to_tuple,
_to_numpy,
)
from ._verify import VerificationError, VerificationResult, verify_model
from .errors import (
OnnxExporterError,
OnnxExporterWarning,
UnsupportedOperatorError,
)
__all__ = [
"DEFAULT_OPSET_VERSION",
"MIN_OPSET_VERSION",
"OnnxExporterError",
"OnnxExporterWarning",
"UnsupportedOperatorError",
"VerificationError",
"VerificationResult",
"errors",
"export",
"is_supported",
"testing",
"utils",
"verification",
]
DEFAULT_OPSET_VERSION = 18
MIN_OPSET_VERSION = 13
#: Factories whose result is fixed by their arguments alone. The ones that
#: leave their contents unset or draw them at random are not among them.
_SETTLED_FACTORIES = frozenset(
{
"arange",
"eye",
"full",
"full_like",
"linspace",
"logspace",
"new_full",
"new_ones",
"new_zeros",
"ones",
"ones_like",
"zeros",
"zeros_like",
}
)
# ---------------------------------------------------------------------------
# Shape / dtype propagation
# ---------------------------------------------------------------------------
def _is_tensor(value: Any) -> bool:
return hasattr(value, "shape") and hasattr(value, "numpy")
def _propagate_metadata(
graph_module: Any, example_inputs: Mapping[str, Any]
) -> dict[str, Any]:
"""Run the captured graph once to learn every node's shape and dtype."""
from ..graph.interpreter import Interpreter
try:
interpreter = Interpreter(graph_module, garbage_collect_values=False)
interpreter.run(**dict(example_inputs))
except Exception as exc: # noqa: BLE001 - metadata is best effort
warnings.warn(
"ONNX export could not evaluate the captured graph on the example "
f"inputs ({type(exc).__name__}: {exc}); shape-dependent operators "
"will fail to lower.",
RuntimeWarning,
stacklevel=3,
)
return {}
return {node.name: value for node, value in interpreter.env.items()}
def _annotate(result: Any, sample: Any) -> Any:
"""Attach the sampled shape/dtype to the value(s) a handler produced."""
if isinstance(result, (list, tuple)):
samples = sample if isinstance(sample, (list, tuple)) else ()
return [
_annotate(item, samples[index] if index < len(samples) else None)
for index, item in enumerate(result)
]
if isinstance(result, Value):
return result
if not _is_tensor(sample):
return Value(result)
return Value(result, _size_to_tuple(sample.shape), _numpy_dtype(sample))
def _numpy_dtype(tensor: Any) -> Any:
# Annotate from the dtype metadata first: a device-resident sample cannot
# be materialized on the host without a transfer, and reading an input's
# type must never force one. Exotic dtypes fall back to host
# materialization (bfloat16 rounds through float32 there).
dtype = getattr(tensor, "dtype", None)
if dtype is not None:
try:
return _dtype_to_numpy(dtype)
except TypeError:
pass
try:
return _to_numpy(tensor).dtype
except Exception: # noqa: BLE001 - exotic dtypes stay unannotated
return None
# ---------------------------------------------------------------------------
# Graph conversion
# ---------------------------------------------------------------------------
class _Converter:
"""Walks the captured graph and emits the equivalent ONNX nodes."""
def __init__(
self,
graph_module: Any,
example_inputs: Mapping[str, Any],
*,
opset_version: int,
input_names: Sequence[str] | None,
output_names: Sequence[str] | None,
dynamic_axes: Mapping[str, Any] | None,
state_values: Mapping[str, Any] | None = None,
num_mutations: int = 0,
) -> None:
self.graph_module = graph_module
self.example_inputs = dict(example_inputs)
self.state_values = dict(state_values or {})
self.num_mutations = int(num_mutations or 0)
self.input_names = list(input_names) if input_names else None
self.output_names = list(output_names) if output_names else None
self.dynamic_axes = dict(dynamic_axes or {})
self.builder = GraphBuilder(opset_version)
self.env: dict[str, Any] = {}
self.samples = _propagate_metadata(
graph_module, {**self.state_values, **self.example_inputs}
)
self.graph_inputs: list[Any] = []
self.eager_outputs: Any = None
# -- helpers ------------------------------------------------------------
def _resolve(self, value: Any) -> Any:
from ..graph.node import Node
if isinstance(value, Node):
try:
return self.env[value.name]
except KeyError: # pragma: no cover - lint guarantees ordering
raise UnsupportedOperatorError(
f"value {value.name!r} is used before it is produced"
) from None
if isinstance(value, tuple):
return tuple(self._resolve(item) for item in value)
if isinstance(value, list):
return [self._resolve(item) for item in value]
if isinstance(value, dict):
return {key: self._resolve(item) for key, item in value.items()}
if isinstance(value, slice):
return slice(
self._resolve(value.start),
self._resolve(value.stop),
self._resolve(value.step),
)
return value
@staticmethod
def _target_id(target: Any) -> tuple[str, str]:
module = getattr(target, "__module__", "") or ""
name = (
getattr(target, "__qualname__", "")
or getattr(target, "__name__", "")
or repr(target)
)
return module, name.rsplit(".", 1)[-1]
# -- node kinds ---------------------------------------------------------
def _placeholder(self, node: Any, index: int) -> Any:
name = (
self.input_names[index]
if self.input_names is not None and index < len(self.input_names)
else node.name
)
self.builder.reserve(name)
sample = self.example_inputs.get(node.name)
if sample is None:
raise UnsupportedOperatorError(
f"input {node.name!r} has no example value; export the program "
"with example inputs for every argument"
)
if _is_tensor(sample):
value = Value(name, _size_to_tuple(sample.shape), _numpy_dtype(sample))
else:
array = _to_numpy(sample)
value = Value(name, tuple(array.shape), array.dtype)
info = self._value_info(name, value)
if info is None:
raise UnsupportedOperatorError(
f"input {node.name!r} has an unsupported example value of type "
f"{type(sample).__name__}"
)
self.graph_inputs.append(info)
return value
def _get_attr(self, node: Any) -> Any:
attribute = self.graph_module._get_attr(str(node.target))
if not _is_tensor(attribute):
return attribute
name = self.builder.unique(str(node.target).replace(".", "_"))
array = _to_numpy(attribute)
self.builder.initializers.append(numpy_helper.from_array(array, name))
return Value(name, tuple(array.shape), array.dtype)
def _state_initializer(self, node: Any) -> Any:
"""Emit a lifted state placeholder as a constant initializer."""
value = self.state_values[node.name]
if not _is_tensor(value):
return Value(str(value))
name = self.builder.unique(node.name)
array = _to_numpy(value)
self.builder.initializers.append(numpy_helper.from_array(array, name))
return Value(name, tuple(array.shape), array.dtype)
def _settled_factory(self, node: Any) -> Any:
"""The tensor a factory call makes, when its arguments settle it.
A capture records ``zeros((2,), dtype=int64)`` as a call so each run
of the program makes its own tensor. An exported model holds values,
not storage, and a factory whose arguments are all plain values makes
the same values every time: those go into the model as an initializer.
A factory that draws at random, leaves its contents unset, or takes
its extent from the graph is not settled and is lowered as a call.
"""
if node.op != "call_function":
return None
name = getattr(node.target, "__name__", "")
if name not in _SETTLED_FACTORIES:
return None
if any(True for _ in _iter_nodes((node.args, node.kwargs))):
return None
sample = self.samples.get(node.name)
return sample if _is_tensor(sample) else None
def _call(self, node: Any) -> Any:
settled = self._settled_factory(node)
if settled is not None:
name = self.builder.unique(node.name)
array = _to_numpy(settled)
self.builder.initializers.append(numpy_helper.from_array(array, name))
return Value(name, tuple(array.shape), array.dtype)
args = [self._resolve(arg) for arg in node.args]
kwargs = {key: self._resolve(value) for key, value in node.kwargs.items()}
if node.op == "call_function":
module, name = self._target_id(node.target)
entry = lookup_function_handler(module, name)
description = f"{module}.{name}" if module else name
else:
name = str(node.target)
entry = lookup_method_handler(name)
description = f"Tensor.{name}"
if entry is None:
raise UnsupportedOperatorError(
f"{description} has no ONNX lowering; register one in "
"tensorplay/onnx/_composite_ops.py or rewrite the model to use "
"a supported operator"
)
handler, params = entry
sample = self.samples.get(node.name)
context = OpContext(
self.builder,
node.name,
params,
args,
kwargs,
out_shape=_size_to_tuple(sample.shape) if _is_tensor(sample) else None,
out_dtype=_numpy_dtype(sample) if _is_tensor(sample) else None,
)
try:
result = handler(context)
except UnsupportedOperatorError:
raise
except Exception as exc: # noqa: BLE001 - point at the offending node
raise UnsupportedOperatorError(
f"failed to lower {description} (node {node.name!r}): {exc}"
) from exc
return _annotate(result, sample)
# -- outputs ------------------------------------------------------------
def _value_info(self, name: str, value: Value) -> Any:
if value.dtype is None:
return None
shape = list(value.shape) if value.shape is not None else None
spec = self.dynamic_axes.get(name)
if shape is not None and spec is not None:
if isinstance(spec, Mapping):
for axis, axis_name in spec.items():
if 0 <= int(axis) < len(shape):
shape[int(axis)] = str(axis_name)
else:
for axis in spec:
if 0 <= int(axis) < len(shape):
shape[int(axis)] = f"{name}_dim_{int(axis)}"
return helper.make_tensor_value_info(
name, _np_dtype_to_onnx(value.dtype), shape
)
def _flatten_outputs(self, value: Any) -> list[Value]:
if isinstance(value, (list, tuple)):
flattened: list[Value] = []
for item in value:
flattened.extend(self._flatten_outputs(item))
return flattened
if isinstance(value, Value):
return [value]
return [Value(str(value))]
# -- driver -------------------------------------------------------------
def convert(self) -> Any:
placeholder_index = 0
outputs: list[Value] = []
for node in self.graph_module.graph.nodes:
if node.op == "placeholder":
if node.name in self.state_values:
# lifted state becomes a constant, not a graph input
self.env[node.name] = self._state_initializer(node)
continue
self.env[node.name] = self._placeholder(node, placeholder_index)
placeholder_index += 1
elif node.op == "get_attr":
self.env[node.name] = self._get_attr(node)
elif node.op in ("call_function", "call_method"):
self.env[node.name] = self._call(node)
elif node.op == "output":
flattened = self._flatten_outputs([self._resolve(node.args[0])])
outputs = flattened[self.num_mutations:]
self.eager_outputs = self.samples.get(node.name)
elif node.op == "call_module":
raise UnsupportedOperatorError(
f"call_module node {node.target!r} reached the ONNX exporter; "
"export inlines submodules, so this graph was captured with a "
"tracer that keeps module boundaries"
)
else: # pragma: no cover - Graph.lint rejects other kinds
raise UnsupportedOperatorError(f"unsupported node kind {node.op!r}")
graph_outputs = []
seen: set[str] = set()
input_names = {info.name for info in self.graph_inputs}
for index, value in enumerate(outputs):
if self.output_names is not None and index < len(self.output_names):
requested = self.builder.reserve(self.output_names[index])
self.builder.op("Identity", [value.name], outputs=[requested])
value = Value(requested, value.shape, value.dtype)
elif value.name in seen or value.name in input_names:
# A value returned twice (or returned unchanged) still needs a
# distinct graph output name.
copied = self.builder.unique(f"{value.name}_out")
self.builder.op("Identity", [value.name], outputs=[copied])
value = Value(copied, value.shape, value.dtype)
seen.add(value.name)
info = self._value_info(value.name, value)
if info is None:
info = helper.make_tensor_value_info(
value.name, TensorProto.UNDEFINED, None
)
graph_outputs.append(info)
return helper.make_graph(
self.builder.nodes,
self.builder.name,
self.graph_inputs,
graph_outputs,
initializer=self.builder.initializers,
value_info=self.builder.value_info,
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def _to_exported_program(model: Any, dynamic_axes: Any) -> tuple[ExportedProgram, Any]:
if isinstance(model, ExportedProgram):
if dynamic_axes is None and model.dynamic_shapes:
dynamic_axes = _dynamic_shapes_to_dynamic_axes(
model.dynamic_shapes, model.graph_signature.user_inputs
)
return model, dynamic_axes
if isinstance(model, (list, tuple)) and model:
callable_, *rest = model
kwargs: dict[str, Any] = {}
if rest and isinstance(rest[-1], dict):
kwargs = dict(rest.pop())
if "dynamic_shapes" in kwargs and dynamic_axes is None:
dynamic_axes = _dynamic_shapes_to_dynamic_axes(
kwargs["dynamic_shapes"], None
)
program = tp_export(callable_, *rest, **kwargs)
return program, dynamic_axes
raise TypeError(
"expected an ExportedProgram or a (model, *args, kwargs) sequence, got "
f"{type(model).__name__}"
)
def _program_state_values(program: Any) -> dict[str, Any]:
"""Resolve lifted state placeholder names to their tensor values."""
from ..export.graph_signature import InputKind
root = program.graph_module.root
values: dict[str, Any] = {}
for spec in program.graph_signature.input_specs:
if spec.kind is InputKind.USER_INPUT or not isinstance(spec.target, str):
continue
value: Any = root
try:
for atom in spec.target.split("."):
value = getattr(value, atom)
except AttributeError:
continue
values[spec.arg.name] = value
return values
[docs]
def export(
exported_program: ExportedProgram | Any,
f: Any = None,
*,
input_names: Sequence[str] | None = None,
output_names: Sequence[str] | None = None,
opset_version: int | None = None,
dynamic_axes: Mapping[str, Mapping[int, str] | Sequence[int]] | None = None,
do_constant_folding: bool = True,
verify: bool = False,
rtol: float = 1e-4,
atol: float = 1e-5,
external_data: bool | None = None,
external_data_location: str | None = None,
check_model: bool = True,
) -> Any:
"""Export a TensorPlay model to ONNX.
Args:
exported_program: an :class:`~tensorplay.export.ExportedProgram`, or a
``(model, *args, kwargs)`` sequence captured on the fly.
f: file path or writable binary file object. When omitted the
``ModelProto`` is returned instead of being written.
input_names: names for the graph inputs, in placeholder order.
output_names: names for the graph outputs.
opset_version: target ONNX opset (default 18, minimum 13).
dynamic_axes: ``{value_name: {axis: axis_name}}`` (or a list of axis
indices) marking dimensions that vary at runtime. Applies to both
inputs and outputs.
do_constant_folding: fold subgraphs whose inputs are all constants.
verify: run the exported model under onnxruntime and compare against
eager execution of ``exported_program``.
rtol/atol: tolerances used by ``verify``.
external_data: store initializers in a side-car file. ``None`` decides
from the model size (models at or above the 2 GiB protobuf limit).
external_data_location: side-car file name for ``external_data``.
check_model: run ``onnx.checker`` over the finished model.
Returns:
The :class:`onnx.ModelProto` when ``f`` is ``None``, else ``None``.
"""
program, dynamic_axes = _to_exported_program(exported_program, dynamic_axes)
opset = DEFAULT_OPSET_VERSION if opset_version is None else int(opset_version)
if opset < MIN_OPSET_VERSION:
raise ValueError(
f"opset_version must be >= {MIN_OPSET_VERSION}, got {opset}"
)
converter = _Converter(
program.graph_module,
program.example_inputs,
opset_version=opset,
input_names=input_names,
output_names=output_names,
dynamic_axes=dynamic_axes,
state_values=_program_state_values(program),
num_mutations=int(
getattr(program.graph_module, "meta", {}).get("num_mutations", 0) or 0
),
)
graph = converter.convert()
model = helper.make_model(
graph,
producer_name="tensorplay",
producer_version=_producer_version(),
opset_imports=[helper.make_opsetid("", opset)],
)
# Declaring a newer IR version than the opset needs makes older runtimes
# reject an otherwise valid model.
model.ir_version = helper.find_min_ir_version_for(
[helper.make_opsetid("", opset)], ignore_unknown=True
)
_passes.optimize(model, do_constant_folding=do_constant_folding)
try:
model = shape_inference.infer_shapes(model, strict_mode=False)
except Exception: # noqa: BLE001 - inference is advisory
pass
if check_model and not _external_data.needs_external_data(model):
checker.check_model(model)
if verify:
expected = converter.eager_outputs
if expected is None:
expected = program(**program.example_inputs)
verify_model(
model,
expected=expected,
input_names=[value.name for value in model.graph.input],
example_inputs=program.example_inputs,
rtol=rtol,
atol=atol,
)
if f is not None:
_external_data.save_model(
model,
f,
external_data=external_data,
location=external_data_location,
)
return None
return model
[docs]
def is_supported(target: Any) -> bool:
"""Whether a captured ``call_function`` target has an ONNX lowering."""
module = getattr(target, "__module__", "") or ""
name = getattr(target, "__qualname__", "") or getattr(target, "__name__", "")
return lookup_function_handler(module, name.rsplit(".", 1)[-1]) is not None
def _producer_version() -> str:
try:
from ..version import __version__
return str(__version__)
except Exception: # noqa: BLE001 - version metadata is optional
return "dev"
def _dynamic_shapes_to_dynamic_axes(
dynamic_shapes: Mapping[str, Mapping[int, Any]],
user_inputs: Sequence[str] | None,
) -> dict | None:
"""Translate ``export(dynamic_shapes=...)`` into ONNX ``dynamic_axes``."""
if not dynamic_shapes:
return None
names = list(user_inputs or dynamic_shapes.keys())
result: dict[str, dict[int, str]] = {}
for index, (argument, dims) in enumerate(dynamic_shapes.items()):
axes: dict[int, str] = {}
for axis, spec in dims.items():
if hasattr(spec, "name"):
axes[int(axis)] = str(spec.name)
if axes:
result[names[index] if index < len(names) else argument] = axes
return result or NoneHelp improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

