Source code for boundlab.ibp

r"""The interval domain: centers (:class:`Bias`) plus halfwidths (:class:`Noise`).

A value is enclosed as

.. math::

   x \in [c - w,\; c + w]
   \quad\Longleftrightarrow\quad
   x = \mathrm{Bias}(c) + \mathrm{Noise}(w),

the cheapest sound abstraction: every transfer function is a couple of tensor
ops.  ``Noise`` components are *independent* — correlations are dropped, so
``x - x`` widens to :math:`\pm 2w` instead of cancelling.  The richer domains
therefore use ``Bias``/``Noise`` as the concrete part they carry alongside
their symbolic components, and fall back to intervals only where symbolic
structure has run out.
"""

from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from functools import partial
from typing import Callable, Literal, Self, final, override
import torch

from boundlab import Expr, ExprGroup, Zeros
from boundlab.error import Reasons
from boundlab.utils import Dim, ShapeDtype, TensorFormat, einsum_formatter, same_shape


[docs] @dataclass(frozen=True) @final class Bias(Expr): """A deterministic additive component: ``lb == ub == arr``. Every linear primitive is exact on it (a linear map of a point is a point), which is why constants never widen anything. """ arr: torch.Tensor
[docs] @classmethod def convert_from(cls, expr: Expr) -> "Bias | None": if isinstance(expr, Zeros): return Bias(torch.zeros(expr.shape_dtype.shape, dtype=expr.shape_dtype.dtype))
@property @override def shape_dtype(self) -> ShapeDtype: return ShapeDtype(self.arr.shape, self.arr.dtype)
[docs] @override def ub(self) -> torch.Tensor: return self.arr
[docs] @override def lb(self) -> torch.Tensor: return self.arr
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: return self.arr, self.arr
[docs] @override def chw(self) -> tuple[torch.Tensor, torch.Tensor]: return self.arr, torch.zeros_like(self.arr)
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "Bias": return Bias( torch.einsum( einsum_formatter(subscripts), self.arr, *operands, ) )
[docs] @override def reshape(self, *shape: Dim) -> "Bias": return Bias(self.arr.reshape(tuple(shape)))
[docs] @override def transpose(self, *perm: int) -> "Bias": return Bias(self.arr.permute(*perm))
[docs] @override def broadcast_to(self, *shape: Dim) -> "Bias": return Bias(torch.broadcast_to(self.arr, tuple(shape)))
[docs] @override def add(self, other: Expr) -> "Bias": if not isinstance(other, Bias): return NotImplemented if not same_shape(self.shape_dtype, other.shape_dtype): raise ValueError(f"Shape mismatch: {self.shape_dtype} and {other.shape_dtype}.") return Bias(self.arr + other.arr)
[docs] @override def torch_print( self, group: Literal["reason", "error_ref"] = "reason", ) -> TensorFormat: del group return TensorFormat("{bias:.5g}", bias=self.arr.mean())
[docs] @dataclass(frozen=True, init=False) @final class Noise(Expr): r"""An independent symmetric interval :math:`\pm w` with provenance. ``einsum`` maps it through the absolute operands — :math:`|A|\,w \ge \sup_{|e| \le w} |A e|` elementwise — which is sound but treats every element as adversarially independent. ``add`` sums the halfwidths and blends the :class:`~boundlab.error.Reasons` tags by mass, so width attribution survives arithmetic. """ noise: torch.Tensor reasons: Reasons = Reasons()
[docs] @classmethod def convert_from(cls, expr: Expr) -> "Noise | None": if isinstance(expr, Zeros): return Noise(torch.zeros(expr.shape_dtype.shape, dtype=expr.shape_dtype.dtype))
[docs] def __new__(cls, noise: torch.Tensor, reasons: Reasons | str = Reasons()) -> "Noise": noise = torch.as_tensor(noise) # Concrete values are checked eagerly (skipped under jit/vmap tracing # and with python -O); a negative or NaN halfwidth is an unsound # interval, never a representable one. if not torch.compiler.is_compiling(): assert bool(torch.all(noise >= -1e-3)), ( "Noise halfwidths must be non-negative (and not NaN): " f"{noise[noise < 0.0].mean()}" ) if isinstance(reasons, str): reasons = Reasons(reasons) instance = super().__new__(cls) object.__setattr__(instance, "noise", noise) object.__setattr__(instance, "reasons", reasons) return instance
@property def shape_dtype(self) -> ShapeDtype: return ShapeDtype(self.noise.shape, self.noise.dtype)
[docs] def ub(self) -> torch.Tensor: return self.noise
[docs] def lb(self) -> torch.Tensor: return -self.noise
[docs] def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: return -self.noise, self.noise
[docs] def chw(self) -> tuple[torch.Tensor, torch.Tensor]: return torch.zeros_like(self.noise), self.noise
[docs] def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "Noise": return Noise( torch.einsum( einsum_formatter(subscripts), self.noise, *(operand.abs() for operand in operands), ), self.reasons, )
[docs] def reshape(self, *shape: Dim) -> "Noise": return Noise(self.noise.reshape(tuple(shape)), self.reasons)
[docs] def transpose(self, *perm: int) -> "Noise": return Noise(self.noise.permute(*perm), self.reasons)
[docs] def broadcast_to(self, *shape: Dim) -> "Noise": return Noise(torch.broadcast_to(self.noise, tuple(shape)), self.reasons)
[docs] def add(self, other: Expr) -> "Noise": if not isinstance(other, Noise): return NotImplemented if not same_shape(self.shape_dtype, other.shape_dtype): raise ValueError(f"Shape mismatch: {self.shape_dtype} and {other.shape_dtype}.") spart = self.noise.sum() opart = other.noise.sum() # Clamp the denominator so two zero-noise operands blend to zero # weights instead of 0/0 = NaN. denom = (spart + opart).clamp(min=torch.finfo(self.noise.dtype).tiny) s, o = spart / denom, opart / denom return Noise( self.noise + other.noise, (self.reasons * s + other.reasons * o), )
[docs] def torch_print( self, group: Literal["reason", "error_ref"] = "reason", ) -> TensorFormat: del group return TensorFormat("±{noise:.8g}", noise=self.noise.mean())
Intervals = ExprGroup[Bias | Noise] from boundlab.interp import Interpreter, base from boundlab.ibp import elementwise from boundlab.ibp.elementwise import Exp, Reciprocal, Relu, Tanh from boundlab.ibp.matmul import MatmulBiased, MatmulNoise from boundlab.ibp.max import MaxWithConst2Relu, MaxWithConstBiased from boundlab.ibp.softmax import Softmax2ExpReciprocal interpret = Interpreter( *base.interpret.values(), MatmulBiased(), MatmulNoise(), MaxWithConst2Relu(), Relu(), Exp(), Tanh(), Reciprocal(), Softmax2ExpReciprocal(), ) """The interval interpreter. Extends :data:`boundlab.interp.base.interpret` with the component-split ``mul``/``matmul`` handlers, endpoint-mapped monotone activations, the ``max``-to-``relu`` rewrites, and the softmax decomposition. Cheapest and least precise: every result is a plain ``Bias + Noise`` box. """ __all__ = [ "matmul", "max", "mul", "softmax", "Bias", "Intervals", "MatmulBiased", "MatmulNoise", "MaxWithConst2Relu", "MaxWithConstBiased", "Noise", "Softmax2ExpReciprocal", "elementwise", "interpret", ]