Source code for boundlab.diff.expr

"""Paired and tripled expressions for differential verification.

Two structurally identical networks ``f₁`` and ``f₂`` are propagated at the
same time so that their *difference* ``f₁(x) − f₂(x)`` can be bounded far more
tightly than by bounding each network on its own and subtracting.

Two carriers implement that:

``DiffExpr2``
    A pair ``(x, y)``.  Every linear operator applies to both components
    independently; no difference is tracked.  Produced by
    :func:`boundlab.diff.ops.diff_pair` for paired weights, and by the
    interpreter for values that have not yet met a non-linearity.

``DiffExpr3``
    A triple ``(x, y, diff)`` where ``diff`` over-approximates
    ``x − y``.  Constant biases cancel in ``diff``; a shared sub-expression
    (a value both networks compute identically) also cancels exactly.

Both are :class:`boundlab.Expr` subclasses, so every base ONNX operator
(reshape, transpose, reduce, cast, …) works on them unchanged: the four
:class:`~boundlab.Expr` primitives simply map over the components.

Examples
--------
>>> import torch
>>> from boundlab import Error
>>> from boundlab.utils import ShapeDtype
>>> from boundlab.zono import Zono
>>> from boundlab.diff.expr import DiffExpr3
>>> err = Error("input", ShapeDtype((2,), torch.float32))
>>> x = Zono.error(err) * 0.1 + torch.tensor([1.0, 2.0])
>>> y = Zono.error(err) * 0.1 + torch.tensor([1.0, 2.5])
>>> triple = DiffExpr3(x, y, x - y)
>>> [round(v, 4) for v in triple.diff.ub().tolist()]
[0.0, -0.5]
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Optional, final, override

import torch

from boundlab import Expr, Zeros
from boundlab.ibp import Bias
from boundlab.utils import Dim, ShapeDtype, TensorFormat, same_shape


def _as_expr(value, shape_dtype: ShapeDtype | None = None) -> Expr:
    """Lift a tensor/scalar into a :class:`~boundlab.ibp.Bias`."""
    if isinstance(value, Expr):
        return value
    tensor = torch.as_tensor(value)
    if shape_dtype is not None:
        tensor = torch.broadcast_to(tensor, tuple(shape_dtype.shape))
    return Bias(tensor)


[docs] @dataclass(frozen=True) @final class DiffExpr2(Expr): """A pair of expressions ``(x, y)``, one per network. All operators apply component-wise; nothing links the two branches, so a ``DiffExpr2`` is exactly as tight as running both networks separately. It is promoted to :class:`DiffExpr3` by the first differential non-linearity (or by :meth:`to3`). """ x: Expr y: Expr
[docs] def __init__(self, x, y): x, y = _as_expr(x), _as_expr(y) if not same_shape(x.shape_dtype, y.shape_dtype): raise ValueError( f"DiffExpr2 components must share a shape: " f"{x.shape_dtype} and {y.shape_dtype}." ) object.__setattr__(self, "x", x) object.__setattr__(self, "y", y)
# ------------------------------------------------------------------ # Expr protocol # ------------------------------------------------------------------ @property @override def shape_dtype(self) -> ShapeDtype: return self.x.shape_dtype def _map(self, fn) -> "DiffExpr2": return DiffExpr2(fn(self.x), fn(self.y))
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor ) -> "DiffExpr2": return self._map(lambda e: e.einsum(subscripts, *operands))
[docs] @override def reshape(self, *shape: Dim) -> "DiffExpr2": return self._map(lambda e: e.reshape(*shape))
[docs] @override def transpose(self, *perm: int) -> "DiffExpr2": return self._map(lambda e: e.transpose(*perm))
[docs] @override def broadcast_to(self, *shape: Dim) -> "DiffExpr2": return self._map(lambda e: e.broadcast_to(*shape))
[docs] @override def add(self, other: Expr) -> "DiffExpr2": assert isinstance(other, DiffExpr2) return DiffExpr2(self.x + other.x, self.y + other.y)
[docs] @override def ub(self) -> torch.Tensor: raise NotImplementedError()
[docs] @override def lb(self) -> torch.Tensor: raise NotImplementedError()
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: raise NotImplementedError()
# ------------------------------------------------------------------ # Differential algebra # ------------------------------------------------------------------
[docs] def to3(self) -> "DiffExpr3": """Promote to a triple, materialising ``diff = x − y``.""" return DiffExpr3(self.x, self.y, self.x - self.y)
[docs] def get_const(self) -> Optional[tuple[torch.Tensor, torch.Tensor]]: """Return ``(x, y)`` as tensors when both branches are constants.""" if isinstance(self.x, Bias) and isinstance(self.y, Bias): return self.x.arr, self.y.arr if isinstance(self.x, Zeros) and isinstance(self.y, Zeros): zeros = torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype) return zeros, zeros return None
[docs] def is_constant(self) -> bool: return self.get_const() is not None
[docs] @override def torch_print(self, group="reason") -> TensorFormat: return TensorFormat("DiffExpr2") + self.x.torch_print(group)
def __repr__(self) -> str: return f"DiffExpr2(x={self.x!r}, y={self.y!r})"
[docs] @dataclass(frozen=True) @final class DiffExpr3(Expr): """A triple ``(x, y, diff)`` where ``diff`` over-approximates ``x − y``. ``x`` and ``y`` bound each network's activations; ``diff`` bounds their difference and is kept *correlated* with them: linearisation errors that are shared between the branches cancel instead of accumulating. Linear operators apply to all three components. Sums with plain expressions or pairs land in an :class:`~boundlab.ExprGroup` (the base ``Expr.__add__``); :func:`lift_diff` folds such a group back into a triple, with shared addends — an affine bias, say — cancelling in ``diff``. """ x: Expr y: Expr diff: Expr
[docs] def __init__(self, x, y, diff): x, y, diff = _as_expr(x), _as_expr(y), _as_expr(diff) if not same_shape(x.shape_dtype, y.shape_dtype): raise ValueError( f"DiffExpr3 x/y must share a shape: " f"{x.shape_dtype} and {y.shape_dtype}." ) if not same_shape(x.shape_dtype, diff.shape_dtype): raise ValueError( f"DiffExpr3 diff must share the branch shape: " f"{diff.shape_dtype} and {x.shape_dtype}." ) object.__setattr__(self, "x", x) object.__setattr__(self, "y", y) object.__setattr__(self, "diff", diff)
[docs] @classmethod @override def convert_from(cls, expr: Expr) -> Optional["DiffExpr3"]: """Lift any expression into a triple. A :class:`DiffExpr2` gets ``diff = x − y``; any other expression is treated as *shared* between the networks, so its diff is exactly zero. """ if isinstance(expr, DiffExpr2): return expr.to3() zero = torch.zeros( tuple(expr.shape_dtype.shape), dtype=expr.shape_dtype.dtype ) return DiffExpr3(expr, expr, Bias(zero))
# ------------------------------------------------------------------ # Expr protocol # ------------------------------------------------------------------ @property @override def shape_dtype(self) -> ShapeDtype: return self.x.shape_dtype def _map(self, fn) -> "DiffExpr3": """Apply *fn* to all three components (bias-free linear ops).""" return DiffExpr3(fn(self.x), fn(self.y), fn(self.diff))
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor ) -> "DiffExpr3": return self._map(lambda e: e.einsum(subscripts, *operands))
[docs] @override def reshape(self, *shape: Dim) -> "DiffExpr3": return self._map(lambda e: e.reshape(*shape))
[docs] @override def transpose(self, *perm: int) -> "DiffExpr3": return self._map(lambda e: e.transpose(*perm))
[docs] @override def broadcast_to(self, *shape: Dim) -> "DiffExpr3": return self._map(lambda e: e.broadcast_to(*shape))
[docs] @override def add(self, other: Expr) -> "DiffExpr3": assert isinstance(other, DiffExpr3) return DiffExpr3( self.x + other.x, self.y + other.y, self.diff + other.diff )
[docs] @override def ub(self) -> torch.Tensor: raise NotImplementedError()
[docs] @override def lb(self) -> torch.Tensor: raise NotImplementedError()
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: raise NotImplementedError()
[docs] @override def torch_print(self, group="reason") -> TensorFormat: return TensorFormat("DiffExpr3") + self.diff.torch_print(group)
def __repr__(self) -> str: return f"DiffExpr3(diff={self.diff!r})" def __str__(self) -> str: parts = "\n".join( f" {name}: {str(value).replace(chr(10), chr(10) + ' ')}," for name, value in (("x", self.x), ("y", self.y), ("diff", self.diff)) ) return f"DiffExpr3 {{\n{parts}\n}}"
[docs] def lift_diff(value: Expr) -> Expr: """Fold the shared part of a mixed sum into its differential component. Adding a plain expression or constant to a differential one goes through the base :meth:`Expr.__add__`, which places unlike addends side by side in an :class:`~boundlab.ExprGroup`. Anything grouped next to a differential component was contributed identically to both networks, so this lift adds it to both branches — where it cancels in the difference — and folds an absorbed pair in with ``diff += pair.x − pair.y``. Values without a differential component pass through untouched. """ classset = value.classset() if DiffExpr3 in classset: triple, rest = value.split(DiffExpr3) assert isinstance(triple, DiffExpr3) x, y, diff = triple.x, triple.y, triple.diff if DiffExpr2 in classset: pair, rest = rest.split(DiffExpr2) assert isinstance(pair, DiffExpr2) x, y = x + pair.x, y + pair.y diff = diff + (pair.x - pair.y) if not isinstance(rest, Zeros): x, y = x + rest, y + rest # shared: cancels in the difference return DiffExpr3(x, y, diff) if DiffExpr2 in classset: pair, rest = value.split(DiffExpr2) assert isinstance(pair, DiffExpr2) if isinstance(rest, Zeros): return pair return DiffExpr2(pair.x + rest, pair.y + rest) return value
type DiffLike = DiffExpr2 | DiffExpr3 __all__ = ["DiffExpr2", "DiffExpr3", "DiffLike", "lift_diff"]