"""Export PyTorch callables to the ONNX IR consumed by BoundLab."""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Any, Callable, Sequence
import warnings
import onnx_ir as ir
import torch
from torch import nn
__all__ = [
"InterpreterModule",
"onnx_export",
]
class _CallableModule(nn.Module):
def __init__(self, function: Callable[..., Any]):
super().__init__()
self.function = function
def forward(self, *args, **kwargs):
return self.function(*args, **kwargs)
def _example(value: torch.Size | Sequence[int] | torch.Tensor) -> torch.Tensor:
if isinstance(value, torch.Tensor):
return value
return torch.zeros(tuple(value), dtype=torch.get_default_dtype())
[docs]
def onnx_export(
function: Callable[..., torch.Tensor] | nn.Module,
args: tuple[torch.Size | Sequence[int] | torch.Tensor, ...],
path: str | Path | None = None,
kwargs: dict[str, torch.Size | Sequence[int] | torch.Tensor] | None = None,
input_names: Sequence[str] | None = None,
output_names: Sequence[str] | None = None,
optimize: bool = True,
) -> ir.Model:
"""Export ``function`` and return its in-memory ONNX IR model.
Tensor arguments are used verbatim. A shape may be supplied instead and
is represented by a float32 zero tensor during export.
"""
module = function if isinstance(function, nn.Module) else _CallableModule(function)
module = module.eval()
example_args = tuple(_example(value) for value in args)
example_kwargs = {
name: _example(value) for name, value in (kwargs or {}).items()
}
exported = torch.export.export(module, example_args, example_kwargs)
destination = str(path) if path is not None else None
registration_logger = logging.getLogger(
"torch.onnx._internal.exporter._registration"
)
previous_level = registration_logger.level
registration_logger.setLevel(logging.ERROR)
try:
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
message=r"`isinstance\(treespec, LeafSpec\)` is deprecated",
category=FutureWarning,
)
program = torch.onnx.export(
exported,
args=(),
f=destination,
export_params=True,
optimize=optimize,
input_names=input_names,
output_names=output_names,
verbose=False,
)
finally:
registration_logger.setLevel(previous_level)
if program is None:
# ``torch.onnx.export`` returns None when saving. Load the saved model
# so callers always receive an interpretable ONNX graph.
assert destination is not None
return ir.load(destination)
return program.model
[docs]
class InterpreterModule(nn.Module):
"""A bound ONNX graph packaged as an ``nn.Module`` returning ``(lb, ub)``.
``input_convert`` lifts the concrete forward inputs into abstract
expressions (e.g. wrapping them in an error ball); the interpreter then
evaluates the graph and the module concretizes. Because the whole
pipeline is expressed in traceable tensor ops, it can itself be exported
or AOT-compiled (:meth:`compiled`).
"""
[docs]
def __init__(
self,
interpreter: Interpreter,
graph,
input_convert: Callable[[torch.Tensor], Expr],
*,
verbose: bool = False,
) -> None:
super().__init__()
# Bind the graph once: this materializes its ONNX initializers, which
# is neither cheap nor traceable to redo on every forward.
self._evaluate = interpreter(graph, verbose=verbose)
self._input_convert = input_convert
[docs]
def forward(self, *args: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Return the lower and upper bounds for the given center point."""
return self._evaluate(self._input_convert(*args)).lbub()
[docs]
def compiled(self, path: str, *args: torch.Tensor) -> AOTICompiledModel:
"""Export + AOTInductor-compile this module and load the package."""
exported = torch.export.export(self, args)
torch._inductor.aoti_compile_and_package(exported, path)
return torch._inductor.aoti_load_package(path)