r"""BoundLab's core expression algebra.
An *abstract value* stands for a set of concrete tensors. BoundLab represents
one as a **sum of typed components**
.. math::
x \;=\; \underbrace{c}_{\text{Bias}}
\;+\; \underbrace{\pm w}_{\text{Noise}}
\;+\; \underbrace{G\,\varepsilon}_{\text{Zono}}
\;+\; \cdots ,
\qquad \varepsilon_k \in [-1, 1],
where each addend is one :class:`Expr` subclass capturing one kind of
uncertainty. Mixed-class sums live in an :class:`ExprGroup` (at most one
component per class), :class:`Zeros` is the empty sum, and
:class:`Intersects` conjoins several enclosures of the *same* value so the
tightest bound of each wins.
Every component class implements a handful of primitives (``einsum``, shape
ops, same-class ``add``, ``ub``/``lb``); everything else — operators, matmul,
reductions, conversion — is derived from those in :class:`Expr`. Soundness
is compositional: if every component's ``ub``/``lb`` encloses its set of
values, so does every derived operation.
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import Sequence
from dataclasses import dataclass
import math
from typing import Any, Generic, Literal, Optional, Self, TypeVar, Union, final, overload, override
import torch
import numpy as np
from plum import Callable
from boundlab import ops
import boundlab
from boundlab import utils
from boundlab.utils import (
Dim,
ShapeDtype,
TensorFormat,
TyDict,
all_same,
all_unique,
same_shape,
)
[docs]
class Expr(ABC):
"""Abstract base of every BoundLab expression component.
Subclasses implement the primitives below; the base class derives the
whole tensor-algebra surface (``+``, ``-``, ``*``, ``@``, reductions,
conversion, concretization) from them. All linear structure is funnelled
through :meth:`einsum`, so a component only has to know how a linear map
acts on its own representation to be sound under every derived operation.
"""
@property
@abstractmethod
def shape_dtype(self) -> ShapeDtype:
"""Allocation-free ``(shape, dtype)`` metadata of the value."""
...
[docs]
@abstractmethod
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> Self:
"""Apply the linear map described by an integer-label einsum.
``subscripts`` holds one label tuple per input followed by the output
labels (see :func:`boundlab.utils.einsum_parser`). This is the single
linear primitive: ``__mul__``, ``__matmul__``, ``sum`` and ``mean``
all lower to it, so implementing it soundly makes every derived
linear operation sound.
"""
...
[docs]
@abstractmethod
def reshape(self, *shape: Dim) -> Self: ...
[docs]
@abstractmethod
def transpose(self, *perm: int) -> Self: ...
[docs]
@abstractmethod
def broadcast_to(self, *shape: Dim) -> Self: ...
[docs]
@abstractmethod
def add(self, other: Self) -> Self:
"""Same-class addition; ``__add__`` dispatches here when classes match
and otherwise groups the addends in an :class:`ExprGroup`."""
...
[docs]
@abstractmethod
def ub(self) -> torch.Tensor:
"""Sound elementwise upper bound on every concrete value represented."""
...
[docs]
@abstractmethod
def lb(self) -> torch.Tensor:
"""Sound elementwise lower bound on every concrete value represented."""
...
@property
def shape(self) -> Sequence[Dim]:
return self.shape_dtype.shape
@property
def dtype(self) -> torch.dtype:
return self.shape_dtype.dtype
[docs]
def lbub(self) -> tuple[torch.Tensor, torch.Tensor]:
"""``(lb, ub)`` in one call; components override it when computing
both at once is cheaper than two passes."""
return self.lb(), self.ub()
[docs]
def absub(self) -> torch.Tensor:
r"""Bound on the magnitude: :math:`\max(|lb|, ub) \ge \sup |x|`."""
lb, ub = self.lbub()
return torch.maximum(-lb, ub)
[docs]
def chw(self) -> tuple[torch.Tensor, torch.Tensor]:
"""Center/halfwidth form: ``c = (ub + lb) / 2``, ``w = (ub - lb) / 2``."""
lower, upper = self.lbub()
return (upper + lower) * 0.5, (upper - lower) * 0.5
[docs]
@final
def to_intervals(self, name: str = "") -> "Expr":
"""Collapse to the box hull ``Bias(c) + Noise(w)``.
Sound but lossy: every correlation between error symbols is dropped,
so downstream cancellation (``x - x = 0``) no longer happens.
"""
del name
# Zeros (empty classset) must stay Zeros: converting it to a zero
# Bias + Noise group makes mul handlers recurse on it forever.
if self.classset().issubset({Bias, Noise}):
return self # type: ignore
center, halfwidth = self.chw()
return Bias(center) + Noise(halfwidth)
@property
@final
def ndim(self) -> int:
return len(self.shape_dtype.shape)
[docs]
@final
def numel(self) -> Dim:
return math.prod(self.shape_dtype.shape)
[docs]
@staticmethod
def from_exprlike(expr: ExprLike, shape: Sequence[Dim]) -> "Expr":
'''Lift a tensor/scalar into a broadcast ``Bias``; pass ``Expr`` through.'''
if isinstance(expr, Expr):
return expr
else:
return boundlab.ibp.Bias(torch.broadcast_to(torch.as_tensor(expr), tuple(shape)))
[docs]
@final
def __add__(self, other: ExprLike) -> Expr:
other = Expr.from_exprlike(other, self.shape_dtype.shape)
if other.__class__ is self.__class__:
return self.add(other) # type: ignore
return ExprGroup.from_sum(self, other)
@final
def __radd__(self, other: ConstLike) -> "Expr":
return self.__add__(other)
@final
def __sub__(self, other: ExprLike) -> "Expr":
return self.__add__(-other)
@final
def __rsub__(self, other: ConstLike) -> "Expr":
return (self * -1.0).__add__(other)
[docs]
def __mul__(self, other: ConstLike) -> Self:
other = torch.as_tensor(other)
if not (other.ndim == 0 or same_shape(other.shape, self.shape_dtype)):
raise ValueError(
f"Elementwise multiplication does not broadcast: "
f"{self.shape_dtype} and {other.shape}."
)
subscript = tuple(range(self.ndim))
other_subscript = subscript if other.ndim else ()
return self.einsum([subscript, other_subscript, subscript], other)
@final
def __rmul__(self, other: Any) -> Self:
return self.__mul__(other)
@final
def __neg__(self) -> Self:
return self * -1.0
@property
@final
def T(self) -> Self:
return self.transpose(*tuple(range(self.ndim - 1, -1, -1)))
[docs]
def squeeze(self, axes: Optional[Union[int, Sequence[int]]] = None) -> Self:
if axes is None:
axes = [
index for index, dim in enumerate(self.shape_dtype.shape) if str(dim) == "1"
]
elif isinstance(axes, int):
axes = [axes]
axes = [axis if axis >= 0 else axis + self.ndim for axis in axes]
if any(str(self.shape_dtype.shape[axis]) != "1" for axis in axes):
raise ValueError(
f"Cannot squeeze non-singleton axes {axes} from {self.shape_dtype}."
)
return self.reshape(*
tuple(dim for index, dim in enumerate(self.shape_dtype.shape) if index not in axes)
)
[docs]
def unsqueeze(self, axes: Union[int, Sequence[int]]) -> Self:
if isinstance(axes, int):
axes = [axes]
output_rank = self.ndim + len(axes)
axes = sorted(axis if axis >= 0 else axis + output_rank for axis in axes)
shape = list(self.shape_dtype.shape)
for axis in axes:
shape.insert(axis, 1)
return self.reshape(*tuple(shape))
[docs]
def sum(self, axis: int, keepdims: bool = False) -> Self:
axis = axis if axis >= 0 else axis + self.ndim
inputs = tuple(range(self.ndim))
outputs = inputs[:axis] + inputs[axis + 1 :]
result = self.einsum([inputs, outputs])
return result.unsqueeze(axis) if keepdims else result
[docs]
def mean(self, axis: int, keepdims: bool = False) -> Self:
axis = axis if axis >= 0 else axis + self.ndim
return self.sum(axis, keepdims) * (1.0 / self.shape_dtype.shape[axis])
def __matmul__(self, other: torch.Tensor) -> Self:
ranks = (self.ndim, len(other.shape))
subscripts = {
(1, 1): [(0,), (0,), ()],
(1, 2): [(0,), (0, 1), (1,)],
(2, 1): [(0, 1), (1,), (0,)],
(2, 2): [(0, 1), (1, 2), (0, 2)],
}
if ranks not in subscripts:
raise ValueError("MatMul currently supports rank-1 and rank-2 operands.")
return self.einsum(subscripts[ranks], other)
def __rmatmul__(self, other: torch.Tensor) -> Self:
ranks = (len(other.shape), self.ndim)
subscripts = {
(1, 1): [(0,), (0,), ()],
(1, 2): [(0, 1), (0,), (1,)],
(2, 1): [(1,), (0, 1), (0,)],
(2, 2): [(1, 2), (0, 1), (0, 2)],
}
if ranks not in subscripts:
raise ValueError("MatMul currently supports rank-1 and rank-2 operands.")
return self.einsum(subscripts[ranks], other)
[docs]
def matmul(self, other: torch.Tensor) -> Self:
return self.__matmul__(other)
[docs]
def rmatmul(self, other: torch.Tensor) -> Self:
return self.__rmatmul__(other)
[docs]
def torch_print(
self,
group: Literal["reason", "error_ref"] = "reason",
) -> TensorFormat:
"""Diagnostics as a ``TensorFormat``.
The result is safe to hand to ``Torch diagnostic output`` under tracing:
only plain arrays cross the callback boundary, never expression
metadata (which may hold tracers, e.g. reason weights).
"""
del group
return TensorFormat("")
[docs]
@classmethod
def convert_from(cls, expr: Expr) -> Optional[Self]:
'''Hook: build an instance of ``cls`` representing exactly ``expr``,
or ``None`` when this class cannot. One half of :meth:`to`.'''
return None
[docs]
def convert_to[U: Expr](self, expr_type: type[U]) -> Optional[U]:
'''Hook: convert ``self`` into ``expr_type``, or ``None`` when this
class does not know how. The other half of :meth:`to`.'''
return None
[docs]
@final
def to[U: Expr](self, expr_type: type[U]) -> U:
'''Convert to another component class, exactly.
Identity short-circuits; otherwise the target's ``convert_from`` is
tried, then this class's ``convert_to``. Raises ``TypeError`` when
neither side knows the conversion — conversions never approximate.
'''
if self.__class__ is expr_type:
return self # type: ignore
if expr := expr_type.convert_from(self):
return expr
if expr := self.convert_to(expr_type):
return expr
raise TypeError(
f"Cannot convert {self.__class__.__name__} to {expr_type.__name__}."
)
[docs]
@final
def classset(self) -> set[type[Expr]]:
"""The set of component classes present in this value.
A single component reports ``{type(self)}``, an :class:`ExprGroup`
its member classes, and :class:`Zeros` the empty set. Handlers use
this to decide which part of a value they know how to transform.
"""
if isinstance(self, ExprGroup):
return {cls for cls in self.keys()}
if isinstance(self, Zeros):
return set()
return {self.__class__}
[docs]
def split[U: Expr](self, ty: type[U]) -> tuple[U, Expr]:
"""Split into ``(matching, rest)`` so that ``matching + rest == self``.
The main way handlers peel off the component class they transform
while passing the remainder through untouched.
"""
if isinstance(self, ExprGroup):
return self.split(ty)
if isinstance(self, ty):
return self, Zeros(self.shape_dtype)
else:
return Zeros(self.shape_dtype).to(ty), self
[docs]
@dataclass(frozen=True)
class ExprGroup[T: Expr](Expr, TyDict[T]):
"""A typed sum of components: at most one :class:`Expr` per class.
Linear primitives map over the members; ``ub``/``lb`` add the members'
bounds (sound since the components are summed). Build one with
:meth:`from_sum`, which merges same-class addends via ``add`` and never
nests groups.
"""
[docs]
def __init__(self, *exprs: T):
super().__init__(*exprs)
if len(self) <= 1:
raise ValueError("ExprGroup must contain at least 2 expression.")
assert all(not isinstance(expr, Zeros) for expr in self.values())
assert all_unique(expr.__class__ for expr in self.values())
if not all_same(expr.shape_dtype for expr in self.values()):
raise ValueError(
"Expression shape mismatch in ExprGroup."
)
assert all(not isinstance(expr, ExprGroup) for expr in self.values()), (
"ExprGroup cannot contain other ExprGroups. "
"Use ExprGroup.add() to combine expressions instead."
)
[docs]
@staticmethod
def from_sum(*args: Expr) -> Expr:
"""Normalized sum of arbitrary components.
Same-class addends are merged with ``add``, :class:`Zeros` vanish,
and the result collapses to a bare component (or ``Zeros``) whenever
fewer than two classes remain.
"""
result = {}
for arg in args:
if isinstance(arg, Zeros):
continue
elif isinstance(arg, ExprGroup):
for cls, expr in arg.items():
if cls in result:
result[cls] = result[cls].add(expr)
assert result[cls] is not NotImplemented
else:
result[cls] = expr
else:
cls = arg.__class__
if cls in result:
result[cls] = result[cls].add(arg)
else:
result[cls] = arg
if len(result.values()) == 0:
return Zeros(args[0].shape_dtype) # type: ignore
if len(result.values()) == 1:
return next(iter(result.values()))
return ExprGroup(*result.values())
@property
@override
def shape_dtype(self) -> ShapeDtype:
return next(iter(self.values())).shape_dtype
def _map(self, func: Callable[[T], T]) -> "ExprGroup[T]":
return ExprGroup(*[func(v) for v in self.values()])
[docs]
def tree_flatten(self):
return tuple(self.values()), None
[docs]
@classmethod
def tree_unflatten(cls, auxiliary, children):
del auxiliary
return cls(*children)
[docs]
@override
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> "ExprGroup[T]":
return self._map(lambda expr: expr.einsum(subscripts, *operands))
[docs]
@override
def reshape(self, *shape: Dim) -> "ExprGroup[T]":
return self._map(lambda expr: expr.reshape(*shape))
[docs]
@override
def transpose(self, *perm: int) -> "ExprGroup[T]":
return self._map(lambda expr: expr.transpose(*perm))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "ExprGroup[T]":
return self._map(lambda expr: expr.broadcast_to(*shape))
[docs]
@override
def add(self, other: "Expr") -> ExprGroup:
return ExprGroup.from_sum(self, other) # type: ignore
[docs]
@override
def ub(self) -> torch.Tensor:
return torch.stack([expr.ub() for expr in self.values()]).sum(0)
[docs]
@override
def lb(self) -> torch.Tensor:
return torch.stack([expr.lb() for expr in self.values()]).sum(0)
[docs]
@override
def lbub(self) -> tuple[torch.Tensor, torch.Tensor]:
lb, ub = utils.unzip([expr.lbub() for expr in self.values()])
lbsum = torch.stack(lb).sum(0)
ubsum = torch.stack(ub).sum(0)
return lbsum, ubsum
[docs]
@override
def chw(self) -> tuple[torch.Tensor, torch.Tensor]:
c, hw = utils.unzip([expr.chw() for expr in self.values()])
csum = torch.stack(c).sum(0)
hwsum = torch.stack(hw).sum(0)
return csum, hwsum
@override
def __repr__(self) -> str:
return f"ExprGroup({', '.join(f'{cls.__name__}: {repr(expr)}' for cls, expr in self.items())})"
[docs]
@override
def torch_print(
self,
group: Literal["reason", "error_ref"] = "reason",
) -> TensorFormat:
return TensorFormat.join(
(part for expr in self.values() if (part := expr.torch_print(group)).fmt),
sep=" + ",
)
@overload
def split[U: Expr](self, u: type[U], /) -> tuple[U, Expr]: ...
@overload
def split[U: Expr, V: Expr](self, u: type[U], v: type[V], /) -> tuple[ExprGroup[U | V], Expr]: ...
@overload
def split[U: Expr, V: Expr, W: Expr](self, u: type[U], v: type[V], w: type[W], /) -> tuple[ExprGroup[U | V | W], Expr]: ...
@overload
def split[U: Expr, V: Expr, W: Expr, X: Expr](self, u: type[U], v: type[V], w: type[W], x: type[X], /) -> tuple[ExprGroup[U | V | W | X], Expr]: ...
[docs]
def split(self, *types: type) -> tuple[Expr, Expr]: # type: ignore[override]
assert all(issubclass(cls, Expr) for cls in types), "All types must be subclasses of Expr."
assert all_unique(types), "All types must be unique."
matching = {cls: expr for cls, expr in self.items() if cls in types}
non_matching = {cls: expr for cls, expr in self.items() if cls not in types}
if len(matching) == 0:
main = Zeros(self.shape_dtype)
elif len(matching) == 1:
main = next(iter(matching.values()))
else:
main = ExprGroup(*matching.values())
if len(non_matching) == 0:
rest = Zeros(main.shape_dtype)
elif len(non_matching) == 1:
rest = next(iter(non_matching.values()))
else:
rest = ExprGroup(*non_matching.values())
return main, rest
[docs]
class Zeros(Expr):
"""The empty sum: exactly the all-zeros value.
The identity for ``add`` and the result of splitting off a class that is
not present. Kept distinct from ``Bias(0)`` so structural checks
(``classset() == set()``) can recognize it.
"""
[docs]
def __init__(self, shape_dtype_struct: ShapeDtype):
self.shape_dtype_struct = shape_dtype_struct
@property
@override
def shape_dtype(self) -> ShapeDtype:
return self.shape_dtype_struct
[docs]
@override
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> "Zeros":
return Zeros(self.shape_dtype)
[docs]
@override
def reshape(self, *shape: Dim) -> "Zeros":
return Zeros(ShapeDtype(tuple(shape), self.shape_dtype.dtype))
[docs]
@override
def transpose(self, *perm: int) -> "Zeros":
return Zeros(ShapeDtype(tuple(self.shape_dtype.shape[i] for i in perm), self.shape_dtype.dtype))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "Zeros":
return Zeros(ShapeDtype(tuple(shape), self.shape_dtype.dtype))
[docs]
@override
def add(self, other: "Expr") -> "Expr":
return other
[docs]
@override
def ub(self) -> torch.Tensor:
return torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
[docs]
@override
def lb(self) -> torch.Tensor:
return torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
type ConstLike = Union[torch.Tensor, np.ndarray, float, int]
type ExprLike = Union[Expr, ConstLike]
[docs]
class Intersects(Expr):
r"""Several sound enclosures of the *same* value, intersected.
Each member independently encloses the one true value, so any pointwise
bound may take the best member:
.. math::
lb = \max_i lb_i, \qquad ub = \min_i ub_i .
Linear operations map over the members (each image still encloses the
image of the value). Used e.g. by the softmax handler to run the
zonotope denominator and its interval version side by side and keep
whichever is tighter per element.
"""
[docs]
def __init__(self, *exprs: T):
self.exprs = exprs
if len(exprs) == 0:
raise ValueError("Intersects must contain at least one expression.")
assert all_same(expr.shape_dtype for expr in exprs), (
"Expression shape mismatch in Intersects."
)
def __len__(self) -> int:
return len(self.exprs)
def __getitem__(self, index: int) -> T:
return self.exprs[index]
def __iter__(self):
return iter(self.exprs)
@property
@override
def shape_dtype(self) -> ShapeDtype:
return self.exprs[0].shape_dtype
[docs]
def lb(self) -> torch.Tensor:
return torch.stack([expr.lb() for expr in self.exprs]).max(0).values
[docs]
def ub(self) -> torch.Tensor:
return torch.stack([expr.ub() for expr in self.exprs]).min(0).values
[docs]
@override
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> "Intersects":
return Intersects(*(expr.einsum(subscripts, *operands) for expr in self.exprs))
[docs]
@override
def reshape(self, *shape: Dim) -> "Intersects":
return Intersects(*(expr.reshape(*shape) for expr in self.exprs))
[docs]
@override
def transpose(self, *perm: int) -> "Intersects":
return Intersects(*(expr.transpose(*perm) for expr in self.exprs))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "Intersects":
return Intersects(*(expr.broadcast_to(*shape) for expr in self.exprs))
[docs]
@override
def add(self, other: "Expr") -> "Intersects":
if isinstance(other, Intersects):
assert len(self) == len(other), "Intersects must have the same number of expressions to add."
return Intersects(*(expr + other_expr for expr, other_expr in zip(self, other)))
assert False
from boundlab.error import Error
from boundlab import ibp
from boundlab.ibp import Bias, Noise
__all__ = [
"Bias",
"Error",
"Expr",
"ExprGroup",
"ExprLike",
"Intersects",
"Noise",
"Zeros",
"ops",
]