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