Source code for boundlab.diff.zono3.bounds

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