r"""The zonotope domain: affine forms over shared error symbols.
A :class:`Zono` represents
.. math::
x \;=\; \sum_k G_{\cdot k}\, \varepsilon_k ,
\qquad \varepsilon_k \in [-1, 1],
with the coefficients :math:`G` stored in a
:class:`~boundlab.zono.generator.Generator` (trailing error axis) and the
symbol layout in a :class:`~boundlab.error.alignment.SpanTable`. (The center
lives outside, as the group's :class:`~boundlab.ibp.Bias` component.)
Linear maps act on :math:`G` exactly, so nothing widens until a non-linearity
is met; values built over the same symbols stay correlated and cancel under
subtraction. Non-linearities use the sound affine enclosures in
:mod:`~boundlab.zono.linearizers`, whose residual becomes a fresh symbol via
``keep_zono``; products of two zonotopes use the estimates in
:mod:`~boundlab.zono.matmul`.
"""
from __future__ import annotations
from collections import defaultdict
from collections.abc import Sequence
from dataclasses import dataclass
from functools import partial
from typing import Literal, Optional, Self, Union, cast, final, override
import torch
from boundlab import Zeros, utils
from boundlab import (
Error,
Expr,
)
from boundlab import ibp
from boundlab.ibp import Bias, Noise
from boundlab.ibp.matmul import MatmulBiased
from boundlab.interp import Interpreter, OpHandler, base
from boundlab.utils import Dim, ShapeDtype, TensorFormat, same_shape
from boundlab.error.alignment import SpanTable, Span
from boundlab.zono.generator import Generator
[docs]
@dataclass(frozen=True)
@final
class Zono(Expr):
r"""A centerless zonotope :math:`\sum_k G_{\cdot k} \varepsilon_k`.
``gen`` holds the coefficients with the flattened error axis last;
``table`` says which slice of that axis belongs to which
:class:`~boundlab.Error`. Binary operations align the tables first
(:meth:`align_fill_zeros`), which is what keeps shared symbols shared.
"""
table: SpanTable[Error]
gen: Generator
[docs]
def __init__(self, table: SpanTable, gen: Generator):
assert isinstance(table, SpanTable)
expected_error_len = table.validate()
if str(gen.error_len) != str(expected_error_len):
raise ValueError(
f"Generator error length {gen.error_len} does not match "
f"sum of expression spans {expected_error_len}."
)
object.__setattr__(self, "table", table)
object.__setattr__(self, "gen", gen)
@property
@override
def shape_dtype(self) -> ShapeDtype:
return self.gen.shape_dtype
[docs]
@classmethod
@override
def convert_from(cls, expr: Expr) -> Optional["Zono"]:
"""Build a zonotope from a single ``Error`` symbol or ``Noise`` interval.
``Bias`` has no zonotope counterpart, so an ``Intervals`` group keeps its
center outside the ``Zono`` and only its ``Noise`` part converts here.
"""
if isinstance(expr, Noise):
err = Error(expr.reasons, expr.shape_dtype)
return Zono.error(err, expr.noise)
if isinstance(expr, Error):
return Zono.error(expr, 1.0)
if isinstance(expr, Zeros):
return Zono.zeros(expr.shape_dtype)
[docs]
@staticmethod
def single(
err: Expr,
generator: Union[Generator, torch.Tensor],
):
'''Zonotope over one error symbol with an explicit coefficient tensor.'''
assert isinstance(generator, Generator) or generator.ndim + 1 == err.ndim
if isinstance(generator, torch.Tensor):
generator = Generator(generator)
return Zono(
SpanTable(err),
generator,
)
[docs]
@staticmethod
def error(
err: Error,
amplitude: Union[torch.Tensor, float] = 1.0,
) -> "Zono":
"""Diagonal zonotope scaling each component of ``err`` by ``amplitude``."""
amplitude = torch.as_tensor(amplitude)
assert amplitude.ndim == 0 or same_shape(amplitude.shape, err.shape_dtype)
size = err.numel()
if amplitude.ndim == 0:
tensor = amplitude * torch.eye(size, dtype=err.shape_dtype.dtype)
else:
tensor = torch.diag(amplitude.reshape(-1))
return Zono(
SpanTable(err),
Generator(tensor.reshape((*err.shape_dtype.shape, size))),
)
[docs]
@staticmethod
def zeros(
shape_dtype: ShapeDtype,
table: SpanTable | None = None,
) -> "Zono":
gen = Generator.from_shape(shape_dtype, 0 if table is None else table.validate())
table = table if table is not None else SpanTable()
return Zono(table, gen)
[docs]
def full_abs(self) -> "Zono":
return Zono(
self.table,
Generator(self.gen.tensor.abs()),
)
[docs]
@override
def ub(self) -> torch.Tensor:
return self.gen.ub()
[docs]
@override
def lb(self) -> torch.Tensor:
return self.gen.lb()
[docs]
@override
def lbub(self) -> tuple[torch.Tensor, torch.Tensor]:
return self.gen.lbub()
[docs]
@override
def chw(self) -> tuple[torch.Tensor, torch.Tensor]:
return self.gen.chw()
[docs]
@override
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> "Zono":
return Zono(
self.table,
self.gen.einsum(subscripts, *operands),
)
def __getitem__(self, item: Error) -> Generator:
span = self.table[item]
return self.gen.error_span(span)
[docs]
def align_fill_zeros(
self,
other: "Zono",
) -> tuple["Zono", "Zono"]:
"""Rebuild both zonotopes over the merged symbol table.
Spans one side does not carry are zero-filled, so afterwards the two
generators are column-aligned and any coefficient-wise operation is
meaningful.
"""
if self.table == other.table:
return self, other
table, alignments = self.table.align_with(other.table)
def aligned_pairs(x: Union[Span, None], y: Union[Span, None]) -> tuple[Generator, Generator]:
if x is None and y is not None:
return Generator.from_shape(self.shape_dtype, y.size), other.gen.error_span(y)
if y is None and x is not None:
return self.gen.error_span(x), Generator.from_shape(other.shape_dtype, x.size)
if x is not None and y is not None:
return self.gen.error_span(x), other.gen.error_span(y)
else:
raise ValueError("Both spans cannot be None.")
x_li, y_li = utils.unzip([aligned_pairs(alignment.span1, alignment.span2) for alignment in alignments])
return (
Zono(table, Generator.concat(*x_li)),
Zono(table, Generator.concat(*y_li)),
)
[docs]
def expanded_to_table(self, table: SpanTable) -> "Zono":
"""Embed into a larger symbol table, zero-filling the new spans."""
assert self.table.subset_of(table), "Cannot expand to a table that is not a subset."
if self.table == table:
return self
_, alignments = self.table.align_with(table)
return Zono(
table,
Generator.concat(*[
self.gen.error_span(alignment.span1)
if alignment.span1 is not None else
Generator.from_shape(self.shape_dtype, utils.cast(Span, alignment.span2).size)
for alignment in alignments
])
)
[docs]
@override
def add(self, other: "Expr") -> "Zono":
other = other.to(Zono)
if not same_shape(self.shape_dtype, other.shape_dtype):
raise ValueError(f"ShapeDtype mismatch: {self.shape_dtype} and {other.shape_dtype}.")
if self.table == other.table:
return Zono(
self.table,
self.gen.add(other.gen),
)
table, alignments = self.table.align_with(other.table)
def add_aligments(x: Union[Span, None], y: Union[Span, None]) -> Generator:
if x is None and y is not None:
return other.gen.error_span(y)
if y is None and x is not None:
return self.gen.error_span(x)
if x is not None and y is not None:
return self.gen.error_span(x).add(other.gen.error_span(y))
else:
raise ValueError("Both spans cannot be None.")
return Zono(
table,
Generator.concat(
*(
add_aligments(alignment.span1, alignment.span2)
for alignment in alignments
)
)
)
[docs]
@override
def reshape(self, *shape: Dim) -> "Zono":
return Zono(self.table, self.gen.reshape(*shape))
[docs]
@override
def transpose(self, *perm: int) -> "Zono":
return Zono(self.table, self.gen.transpose(*perm))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "Zono":
return Zono(self.table, self.gen.broadcast_to(*shape))
[docs]
def reasons_breakdown(self) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
"""Total halfwidth and its split over the symbols' ``Reasons`` labels —
each symbol contributes its coefficient-slice's concretized width."""
total = torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
reasons = defaultdict(
lambda: torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
)
for err, span in self.table.items():
contribution = self.gen.error_span(span).ub()
total += contribution
expr_reasons = err.reasons
for name, weight in expr_reasons.items():
reasons[name] += contribution * weight
return total, reasons
[docs]
def torch_print(
self,
group: Literal["reason", "error_ref"] = "reason",
) -> TensorFormat:
if group == "reason":
total, reasons = self.reasons_breakdown()
total = 2 * total.mean()
groups = {
name: 2 * reason.mean()
for name, reason in reasons.items()
}
else:
groups = {
str(error_ref): 2 * self.gen.error_span(span).ub().mean()
for error_ref, span in self.table.items()
}
total = sum(groups.values())
fmt = ", ".join(f"{name}={{{name}:.2g}}" for name in groups)
return TensorFormat(f"Zono({{total:.8g}}, {fmt})", total=total, **groups)
from boundlab.zono import linearizers
from boundlab.zono import matmul
from boundlab.zono.matmul import Matmul
from boundlab.ibp.softmax import Softmax2ExpReciprocal
[docs]
def keep_zono(x, name: str=""):
"""``after_each`` normalizer: promote fresh ``Noise`` to new ``Zono`` symbols.
Linearizers emit their residual as ``Noise``; converting it here gives the
residual its own error symbol, so later layers see it as a first-class
coordinate that can cancel instead of an anonymous interval that only
accumulates.
"""
del name
# Multiple-results primitives (e.g. custom_jvp_call) hand over a list.
if isinstance(x, (list, tuple)):
return type(x)(keep_zono(v) for v in x)
if isinstance(x, Expr) and Noise in x.classset():
noise, other = x.split(Noise)
return noise.to(Zono) + other
assert isinstance(x, Expr) and x.classset().issubset({Bias, Zono}) or isinstance(x, torch.Tensor)
return x
interpret = Interpreter(
base.interpret,
MatmulBiased(),
Matmul(),
linearizers.MaxWithConst2Relu(),
linearizers.Relu(),
linearizers.Exp(),
linearizers.Tanh(),
linearizers.Reciprocal(),
Softmax2ExpReciprocal(),
after_each=keep_zono
)
"""The zonotope interpreter.
Extends :data:`boundlab.interp.base.interpret` with the biased matmul split,
the zonotope × zonotope :class:`~boundlab.zono.matmul.Matmul` estimate, the
affine linearizers for ``relu``/``exp``/``tanh``/``reciprocal``, and the
softmax decomposition; ``keep_zono`` promotes every fresh ``Noise`` residual
to a new error symbol after each node.
"""
__all__ = [
"Generator",
"Matmul",
"Zono",
"interpret",
"keep_zono",
"linearizers",
"matmul",
]