"""Named error symbols and their provenance.
An :class:`Error` is the atom of uncertainty in BoundLab: a tensor-shaped
symbolic variable ranging over :math:`[-1, 1]^{shape}`, identified by a
globally unique ``ID``. Domains build abstract values as functions of these
symbols; two values referencing the *same* symbol stay correlated, which is
what lets ``x - x`` concretize to exactly zero.
:class:`Reasons` records where a symbol's error mass came from ("input",
"relu", "matmul", ...) as a weighted set of labels, so bound widths can be
attributed back to the operations that produced them.
"""
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from functools import partial
import itertools
import math
from typing import Literal, Self, override
import torch
from boundlab import Expr
from boundlab import utils
from boundlab.utils import Dim, ShapeDtype, same_shape
[docs]
@dataclass(frozen=True)
class Reasons(dict[str, torch.Tensor]):
"""A collection of named reasons for an error."""
[docs]
def __init__(self, *args: str, **kwargs: torch.Tensor):
assert len(args) == 0 or len(kwargs) == 0, "Reasons must be initialized with either positional or keyword arguments, not both."
if len(args) > 0:
kwargs = {arg: torch.as_tensor(1/len(args)) for arg in args}
super().__init__(kwargs)
[docs]
def __mul__(self, other: torch.Tensor) -> "Reasons":
return Reasons(**{k: v * other for k, v in self.items()})
def __truediv__(self, other: torch.Tensor) -> "Reasons":
return Reasons(**{k: v * (1 / other) for k, v in self.items()})
[docs]
def __add__(self, other: Self) -> "Reasons":
zero = torch.zeros((), dtype=torch.float32)
return Reasons(**{k: self.get(k, zero) + other.get(k, zero) for k in set(self) | set(other)})
def __str__(self) -> str:
return "&".join(f"{k}" for k in self)
[docs]
def tree_flatten(self):
keys = tuple(sorted(self))
return tuple(self[key] for key in keys), keys
[docs]
@classmethod
def tree_unflatten(cls, keys, values):
return cls(**dict(zip(keys, values)))
[docs]
@staticmethod
def weighted_avg(*args: tuple["Reasons", torch.Tensor]) -> "Reasons":
"""Weighted average of multiple Reasons."""
total_weight = sum(weight for _, weight in args)
if total_weight == 0:
return Reasons()
assert len(args) > 0, "No reasons provided for weighted average."
return utils.sum(reason * weight for reason, weight in args) / total_weight
_COUNTER = itertools.count()
[docs]
@dataclass(frozen=True)
class Error(Expr):
"""A named independent error expression with values in ``[-1, 1]``."""
ID: int
reasons: Reasons
name: str
shape_dtype_struct: ShapeDtype
[docs]
def __init__(
self,
reason: str | Reasons,
shape_dtype_struct: ShapeDtype | torch.Tensor,
name: str = "",
):
if isinstance(shape_dtype_struct, torch.Tensor):
shape_dtype_struct = ShapeDtype.from_value(shape_dtype_struct)
if not isinstance(shape_dtype_struct, ShapeDtype):
raise TypeError("Error shape metadata must be a ShapeDtype or tensor.")
object.__setattr__(self, "ID", next(_COUNTER))
if isinstance(reason, str):
reason = Reasons(reason)
object.__setattr__(self, "reasons", reason)
object.__setattr__(self, "name", str(name))
object.__setattr__(self, "shape_dtype_struct", shape_dtype_struct)
@property
def shape_dtype(self) -> ShapeDtype:
return self.shape_dtype_struct
[docs]
@override
def ub(self) -> torch.Tensor:
return torch.ones(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
[docs]
@override
def lb(self) -> torch.Tensor:
return -self.ub()
[docs]
@override
def einsum(
self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor
) -> Expr:
from boundlab.zono import Zono
return self.to(Zono).einsum(subscripts, *operands)
[docs]
@override
def reshape(self, *shape: Dim) -> Expr:
from boundlab.zono import Zono
return self.to(Zono).reshape(*shape)
[docs]
@override
def transpose(self, *perm: int) -> Expr:
from boundlab.zono import Zono
return self.to(Zono).transpose(*perm)
[docs]
@override
def broadcast_to(self, *shape: Dim) -> Expr:
from boundlab.zono import Zono
return self.to(Zono).broadcast_to(*shape)
[docs]
@override
def add(self, other: Expr) -> Expr:
raise NotImplementedError("Cannot add to an Error directly. Use Error.to(Zono) + other.")
def __str__(self) -> str:
return f"{self.reasons}@{self.ID}{self.name}"
def __repr__(self) -> str:
return str(self)
def __lt__(self, other: Self) -> bool:
return self.ID < other.ID
def __le__(self, other: Self) -> bool:
return self.ID <= other.ID
@override
def __hash__(self) -> int:
return hash(self.ID)
__all__ = [
"alignment",
"Error",
"Reasons",
]