Source code for boundlab.zono

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