Source code for boundlab

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