"""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",
]