Source code for boundlab.interp.export

"""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)