Source code for boundlab.interp.base

"""Base ONNX operators shared by every BoundLab abstract domain."""

from __future__ import annotations

from typing import Any

import torch

from boundlab import Expr, Zeros
from boundlab.interp import Interpreter, OpHandler, handler_fn


@handler_fn("add")
def Add(x, y, **params):
    """Addition via ``Expr.__add__`` — exact for every component class."""
    del params
    return x + y


@handler_fn("sub")
def Sub(x, y, **params):
    """Subtraction ``x + (-1) * y`` — exact; shared symbols cancel."""
    del params
    return x - y


[docs] class MulSimple(OpHandler): """Product with at most one abstract operand — exact via ``einsum``.""" op = "mul"
[docs] def condition(self, x, y, **kwargs): del kwargs return ( not isinstance(x, Expr) or not isinstance(y, Expr) or isinstance(x, Zeros) or isinstance(y, Zeros) )
[docs] def handle(self, interp: Interpreter, x: Any, y: Any, **kwargs) -> Any: del interp, kwargs if isinstance(x, Zeros): return x if isinstance(y, Zeros): return y return x * y
[docs] class DivSimple(OpHandler): """Division with a concrete divisor (``x * (1/y)``) or dividend (``reciprocal(y) * x``); the fully abstract case belongs to the domain.""" op = "div"
[docs] def condition(self, x, y, **kwargs): del kwargs return not (isinstance(x, Expr) and isinstance(y, Expr))
[docs] def handle(self, interp: Interpreter, x: Any, y: Any, **kwargs) -> Any: del kwargs if isinstance(x, Expr): return x * (1 / y) if isinstance(y, Expr): return interp.reciprocal(y) * x return x / y
@handler_fn("neg") def Neg(value, **params): """Negation — exact.""" del params return -value @handler_fn("reshape") def Reshape(value, *, new_sizes, **params): """ONNX ``Reshape`` with a constant target shape.""" del params return value.reshape(*new_sizes) @handler_fn("transpose") def Transpose(value, *, permutation, **params): """ONNX ``Transpose``; expressions use ``Expr.transpose``, tensors ``permute``.""" del params if isinstance(value, Expr): return value.transpose(*permutation) return value.permute(*permutation) @handler_fn("broadcast_to") def BroadcastTo(value, *, shape, **params): """ONNX ``Expand`` to a constant shape.""" del params if isinstance(value, Expr): return value.broadcast_to(*shape) return torch.broadcast_to(value, shape) @handler_fn("reduce_sum") def ReduceSum(value, *, axes=None, keepdims=False, **params): """ONNX ``ReduceSum`` axis by axis via ``Expr.sum`` — exact (linear).""" del params if axes is None: axes = tuple(range(value.ndim)) for axis in sorted(axes, reverse=True): value = value.sum(axis, keepdims) return value @handler_fn("reduce_mean") def ReduceMean(value, *, axes=None, keepdims=False, **params): """ONNX ``ReduceMean`` axis by axis via ``Expr.mean`` — exact (linear).""" del params if axes is None: axes = tuple(range(value.ndim)) for axis in sorted(axes, reverse=True): value = value.mean(axis, keepdims) return value @handler_fn("unsqueeze") def Unsqueeze(value, *, axes, **params): """ONNX ``Unsqueeze``: insert size-1 axes.""" del params for axis in sorted(axes): value = value.unsqueeze(axis) return value @handler_fn("squeeze") def Squeeze(value, *, axes=None, **params): """ONNX ``Squeeze``: drop size-1 axes.""" del params return value.squeeze(axes) @handler_fn("identity") def Identity(value, **params): """ONNX ``Identity`` — pass through.""" del params return value _ONNX_DTYPES = { 1: torch.float32, 6: torch.int32, 7: torch.int64, 9: torch.bool, 10: torch.float16, 11: torch.float64, 16: torch.bfloat16, } @handler_fn("cast") def Cast(value, *, to, **params): """ONNX ``Cast``; abstract expressions only allow the no-op cast.""" del params if isinstance(value, Expr): if value.dtype == _ONNX_DTYPES[to]: return value raise TypeError("Casting an abstract expression to a new dtype is unsupported.") return value.to(_ONNX_DTYPES[to])
[docs] class MatmulSimple(OpHandler): """Matrix product with one concrete operand — exact, a linear map on the abstract one.""" op = "matmul"
[docs] def condition(self, x, y, **kwargs): del kwargs return not (isinstance(x, Expr) and isinstance(y, Expr))
[docs] def handle(self, interp: Interpreter, x: Any, y: Any, **kwargs) -> Any: del interp, kwargs return x @ y
[docs] class Gemm(OpHandler): """ONNX ``Gemm``: ``alpha * op(x) @ op(weight) + beta * bias``, lowered to the ``transpose`` / ``matmul`` / ``mul`` / ``add`` handlers.""" op = "gemm"
[docs] def handle( self, interp: Interpreter, x, weight, bias=None, *, alpha=1.0, beta=1.0, transA=0, transB=0, **kwargs, ): del kwargs if transA: x = interp.transpose(x, permutation=(1, 0)) if transB: weight = interp.transpose(weight, permutation=(1, 0)) result = interp.matmul(x, weight) if alpha != 1.0: result = interp.mul(result, alpha) if bias is not None: result = interp.add(result, interp.mul(bias, beta)) return result
class MarkedIdentity(OpHandler): """Emit a custom ``boundlab::MarkedIdentity`` ONNX node marking ``x`` as a residual connection, so interpreters can treat the skip path specially instead of seeing an anonymous ``Add``.""" op = "marked_identity" def handle(self, interp: Interpreter, x, **kwargs): del interp, kwargs return x interpret = Interpreter( Add(), Sub(), MulSimple(), DivSimple(), MatmulSimple(), Gemm(), Neg(), Reshape(), Transpose(), BroadcastTo(), ReduceSum(), ReduceMean(), Unsqueeze(), Squeeze(), Identity(), Cast(), MarkedIdentity(), ) """The domain-independent base interpreter. Bundles the operators every domain shares — implemented purely against the :class:`~boundlab.Expr` primitives, so they are exact for any component class: ``add``/``sub``/``neg``, ``mul``/``div``/``matmul`` with at least one concrete operand, ``gemm``, the shape operators, reductions, ``identity`` and ``cast``. Domains extend it: ``Interpreter(base.interpret, <nonlinear handlers>, ...)``. """ __all__ = [ "Add", "BroadcastTo", "Cast", "DivSimple", "Gemm", "Identity", "MatmulSimple", "MulSimple", "Neg", "Reshape", "ReduceMean", "ReduceSum", "Squeeze", "Sub", "Transpose", "Unsqueeze", "interpret", ]