r"""Triple-zonotope abstract interpretation for differential verification.
Two structurally identical networks are propagated together as a triple
:class:`~boundlab.diff.expr.DiffExpr3` ``(x, y, d)``, where ``x`` and ``y``
bound each network's activations and ``d`` bounds their *difference*. Because
``d`` is tracked as a first-class zonotope over the same error symbols as the
branches, everything the two networks share cancels instead of accumulating —
which is what makes ``f₁(x) − f₂(x)`` provable at perturbation radii where
bounding each network separately gives nothing.
Affine operations need no special handling: the four
:class:`~boundlab.Expr` primitives map over the components, biases cancel in
``d``, and the base ONNX operators work unchanged. Non-linearities use the
differential linearisers in this package — a nine-case ReLU split following
VeryDiff (Teuber et al., 2024), and hexagon-Chebyshev envelopes for
``exp`` / ``tanh`` / ``reciprocal`` that bound the divided difference over the
*feasible* region rather than over a merged interval.
Examples
--------
>>> import torch
>>> from torch import nn
>>> from boundlab import Error
>>> from boundlab.utils import ShapeDtype
>>> from boundlab.zono import Zono
>>> from boundlab.diff.expr import DiffExpr3
>>> from boundlab.diff.zono3 import interpret
>>> err = Error("input", ShapeDtype((4,), torch.float32))
>>> x = Zono.error(err) * 0.05 + torch.zeros(4)
>>> triple = DiffExpr3(x, x, torch.zeros(4))
>>> model = nn.Sequential(nn.Linear(4, 5), nn.ReLU(), nn.Linear(5, 3))
>>> out = interpret(model)(triple)
>>> tuple(out.diff.ub().shape)
(3,)
"""
from __future__ import annotations
from typing import Any
from boundlab import Expr
from boundlab.ibp import Bias
from boundlab.ibp.matmul import MatmulBiased
from boundlab.ibp.softmax import Softmax2ExpReciprocal
from boundlab.interp import Interpreter, base
from boundlab.zono import Zono, keep_zono, linearizers, matmul
from boundlab.diff import ops
from boundlab.diff.expr import DiffExpr2, DiffExpr3, lift_diff
from .bounds import DiffBounds, apply_diff_bounds, diff_linearizer_fn, tighten_diff
from .bilinear import DiffMatmul, DiffMul, diff_bilinear, diff_pairwise
from .exp import DiffExp
from .heaviside import DiffHeavisidePruning, DiffTopKPruning
from .reciprocal import DiffReciprocal
from .relu import DiffRelu
from .softmax import DiffSoftmax, DiffSoftmaxPruning
from .tanh import DiffTanh
[docs]
def keep_diff(inner):
"""Lift a domain's ``after_each`` normaliser over differential components.
A node whose result mixes a differential expression with a shared one
(a bias add, a residual constant) produces an :class:`~boundlab.ExprGroup`
— the base ``Expr.__add__`` has no differential rules — so each result is
first folded back into a pair or triple by
:func:`~boundlab.diff.expr.lift_diff`. Constant
(:class:`~boundlab.ibp.Bias`) components are left alone: paired weights
stay concrete, which keeps every product involving them exact.
"""
def normalize(value: Any, name: str = "") -> Any:
if isinstance(value, (list, tuple)):
return type(value)(normalize(item, name) for item in value)
if isinstance(value, Bias):
return value
if isinstance(value, Expr):
value = lift_diff(value)
if isinstance(value, DiffExpr3):
return DiffExpr3(
normalize(value.x, name),
normalize(value.y, name),
normalize(value.diff, name),
)
if isinstance(value, DiffExpr2):
return DiffExpr2(normalize(value.x, name), normalize(value.y, name))
return inner(value, name)
return normalize
[docs]
def diff_interpreter(
*extra,
domain: type[Expr] = Zono,
tighten: bool = True,
after_each=None,
) -> Interpreter:
"""Assemble a differential interpreter over ``domain``.
``domain`` is the expression class fresh error symbols are built in —
:class:`~boundlab.zono.Zono` here, :class:`~boundlab.polysp.PolySp` for
:mod:`boundlab.diff.polysp3`. ``tighten`` keeps the per-element choice
between the lineariser's difference form and ``x_out − y_out``; turning it
off is cheaper but loses precision.
Each differential handler *owns* its operator: the standard handler for
plain expressions rides along as its ``fallback`` instead of being
registered beside it, so no operator ever has two ready handlers.
"""
return Interpreter(
base.interpret,
ops.DiffPair(),
DiffMul(tighten=tighten),
DiffMatmul(tighten=tighten),
DiffRelu(domain=domain, tighten=tighten, fallback=linearizers.Relu()),
DiffExp(domain=domain, tighten=tighten, fallback=linearizers.Exp()),
DiffTanh(domain=domain, tighten=tighten, fallback=linearizers.Tanh()),
DiffReciprocal(
domain=domain, tighten=tighten, fallback=linearizers.Reciprocal()
),
DiffSoftmax(fallback=Softmax2ExpReciprocal()),
DiffSoftmaxPruning(),
DiffHeavisidePruning(),
DiffTopKPruning(),
*extra,
after_each=after_each,
)
interpret = diff_interpreter(
MatmulBiased(),
matmul.Matmul(),
linearizers.MaxWithConst2Relu(),
domain=Zono,
after_each=keep_diff(keep_zono),
)
"""Differential interpreter over dense zonotopes.
Feed it a :class:`~boundlab.diff.expr.DiffExpr3` ``(x, y, d)``, a
:class:`~boundlab.diff.expr.DiffExpr2` pair, or a plain
:class:`~boundlab.Expr` — the last case falls straight through to standard
:mod:`boundlab.zono` interpretation.
"""
__all__ = [
"DiffBounds",
"DiffExp",
"DiffHeavisidePruning",
"DiffMatmul",
"DiffMul",
"DiffReciprocal",
"DiffRelu",
"DiffSoftmax",
"DiffSoftmaxPruning",
"DiffTanh",
"DiffTopKPruning",
"apply_diff_bounds",
"diff_bilinear",
"diff_interpreter",
"diff_linearizer_fn",
"diff_pairwise",
"interpret",
"keep_diff",
"tighten_diff",
]