"""Shared plumbing for differential linearisers.
A differential lineariser looks at the bounds of a triple ``(x, y, d)`` and
returns a :class:`DiffBounds`: an affine enclosure for each branch plus one for
the difference. :func:`apply_diff_bounds` turns that description into a
:class:`~boundlab.diff.expr.DiffExpr3`.
The interesting part is *error sharing*. Each branch's relaxation introduces a
fresh error symbol (``ε_x``, ``ε_y``); the difference may reference those very
same symbols through ``diff_x_error`` / ``diff_y_error`` instead of paying for a
fresh one. Because BoundLab identifies error symbols by identity and aligns
them when expressions meet, that makes the difference track ``x_out − y_out``
*exactly* wherever a lineariser can arrange it — which is what separates a
differential domain from bounding both networks separately.
"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
import torch
from boundlab import Error, Expr
from boundlab.interp import Interpreter, OpHandler
from boundlab.utils import ShapeDtype, is_statically_zero
from boundlab.zono import Zono
from boundlab.zono.linearizers import LinearBounds
from boundlab.diff.expr import DiffExpr2, DiffExpr3
[docs]
@dataclass
class DiffBounds:
"""Affine enclosures for both branches and for their difference.
``x`` and ``y`` are ordinary zonotope relaxations,
``branch ≈ slope·input + bias ± error·ε``.
The difference is enclosed as::
diff ≈ diff_slope·d + diff_x_weight·x + diff_y_weight·y + diff_bias
± diff_x_error·ε_x ± diff_y_error·ε_y ± diff_error·ε_new
where ``ε_x`` and ``ε_y`` are the *same* symbols the ``x`` and ``y``
enclosures introduced. A weight or error left at ``0`` costs nothing.
"""
name: str
x: LinearBounds
y: LinearBounds
diff_slope: torch.Tensor
diff_bias: torch.Tensor
diff_error: torch.Tensor
diff_x_weight: torch.Tensor | float = 0.0
diff_y_weight: torch.Tensor | float = 0.0
diff_x_error: torch.Tensor | float = 0.0
diff_y_error: torch.Tensor | float = 0.0
def _nonzero(value) -> bool:
return not is_statically_zero(value)
[docs]
def tighten_diff(sub: Expr, diff: Expr) -> Expr:
"""Pick, per element, the narrower of two sound enclosures of the difference.
``sub`` is ``x_out − y_out`` — exact whenever the branch relaxations share
their error symbols with the difference — and ``diff`` is the lineariser's
dedicated form. Neither dominates the other: ``sub`` wins when the branches
cancel, ``diff`` wins when the lineariser exploits the strip constraint on
``x − y``. Both enclose the true difference, so any per-element choice is
sound.
"""
sub_lb, sub_ub = sub.lbub()
diff_lb, diff_ub = diff.lbub()
sub_width = sub_ub - sub_lb
diff_width = diff_ub - diff_lb
# A 0/1 blend is only exact when *both* candidates are finite: masking a
# non-finite coefficient evaluates ``0 * inf = NaN`` and poisons every
# element. When one side is non-finite (an honest vacuous reciprocal
# envelope, say), take the finite side wholesale -- both enclose the true
# difference, so any choice is sound.
sub_finite = bool(torch.isfinite(sub_width).all())
diff_finite = bool(torch.isfinite(diff_width).all())
if not diff_finite:
return sub if sub_finite else diff
if not sub_finite:
return diff
sub_narrower = sub_width < diff_width
if bool(sub_narrower.all()):
return sub
if not bool(sub_narrower.any()):
return diff
mask = sub_narrower.to(sub_ub.dtype)
return sub * mask + diff * (1.0 - mask)
[docs]
def apply_diff_bounds(
bounds: DiffBounds,
triple: DiffExpr3,
domain: type[Expr] = Zono,
tighten: bool = True,
) -> DiffExpr3:
"""Realise *bounds* against the input *triple* in the given ``domain``.
``domain`` is the expression class used for fresh error symbols — ``Zono``
for :mod:`boundlab.diff.zono3`, ``PolySp`` for
:mod:`boundlab.diff.polysp3`.
"""
x, y, d = triple.x, triple.y, triple.diff
shape_dtype = ShapeDtype(x.shape_dtype.shape, x.shape_dtype.dtype)
def symbol(suffix: str) -> Error:
return Error(bounds.name, shape_dtype, suffix)
def scaled(error: Error, amplitude) -> Expr:
return domain.error(error, amplitude) # type: ignore[attr-defined]
x_out = x * bounds.x.slope + bounds.x.bias
y_out = y * bounds.y.slope + bounds.y.bias
diff_out = d * bounds.diff_slope + bounds.diff_bias
# Each branch's error symbol is created once and referenced again by the
# difference, so whatever the branches share cancels there instead of
# accumulating. That reuse is the whole point of the domain.
if _nonzero(bounds.x.error):
eps_x = symbol("x")
x_out = x_out + scaled(eps_x, bounds.x.error)
if _nonzero(bounds.diff_x_error):
diff_out = diff_out + scaled(eps_x, bounds.diff_x_error)
if _nonzero(bounds.y.error):
eps_y = symbol("y")
y_out = y_out + scaled(eps_y, bounds.y.error)
if _nonzero(bounds.diff_y_error):
diff_out = diff_out + scaled(eps_y, bounds.diff_y_error)
if _nonzero(bounds.diff_x_weight):
diff_out = diff_out + x * bounds.diff_x_weight
if _nonzero(bounds.diff_y_weight):
diff_out = diff_out + y * bounds.diff_y_weight
if _nonzero(bounds.diff_error):
diff_out = diff_out + scaled(symbol("d"), bounds.diff_error)
if tighten:
diff_out = tighten_diff(x_out - y_out, diff_out)
return DiffExpr3(x_out, y_out, diff_out)
DiffLinearizerFn = Callable[
[tuple[torch.Tensor, torch.Tensor],
tuple[torch.Tensor, torch.Tensor],
tuple[torch.Tensor, torch.Tensor]],
DiffBounds,
]
"""``(x_lbub, y_lbub, diff_lbub) -> DiffBounds``."""
[docs]
def diff_linearizer_fn(op: str) -> Callable[[DiffLinearizerFn], type[OpHandler]]:
"""Turn a differential lineariser into an :class:`~boundlab.interp.OpHandler`.
The handler owns its operator outright: differential inputs take the
lineariser, and plain expressions are delegated to ``fallback`` — the
domain's standard handler, held *inside* this one rather than registered
beside it, so the interpreter never sees two ready handlers for one op.
"""
op_name = op
def decorator(linearizer: DiffLinearizerFn) -> type[OpHandler]:
@dataclass
class DiffLinearizerHandler(OpHandler):
op: str = op_name
domain: type[Expr] = Zono
tighten: bool = True
fallback: OpHandler | None = None
def condition(self, x: Any, **kwargs: Any) -> bool:
if isinstance(x, (DiffExpr2, DiffExpr3)):
return True
return self.fallback is not None and self.fallback.condition(
x, **kwargs
)
def handle(self, interp: Interpreter, x: Any, **kwargs: Any) -> Expr:
if not isinstance(x, (DiffExpr2, DiffExpr3)):
assert self.fallback is not None
return self.fallback.handle(interp, x, **kwargs)
del interp, kwargs
triple = x.to(DiffExpr3)
bounds = linearizer(
triple.x.lbub(), triple.y.lbub(), triple.diff.lbub()
)
return apply_diff_bounds(
bounds, triple, self.domain, self.tighten
)
def __call__(
self,
x_lbub: tuple[torch.Tensor, torch.Tensor],
y_lbub: tuple[torch.Tensor, torch.Tensor],
d_lbub: tuple[torch.Tensor, torch.Tensor],
) -> DiffBounds:
return linearizer(x_lbub, y_lbub, d_lbub)
DiffLinearizerHandler.__name__ = linearizer.__name__
DiffLinearizerHandler.__qualname__ = linearizer.__qualname__
DiffLinearizerHandler.__doc__ = linearizer.__doc__
return DiffLinearizerHandler
return decorator
__all__ = [
"DiffBounds",
"DiffLinearizerFn",
"apply_diff_bounds",
"diff_linearizer_fn",
"tighten_diff",
]