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