"""Differential products.
Both element-wise and matrix products follow the same identity, which keeps the
difference linear in the *tracked* difference of each factor::
a₁·b₁ − a₂·b₂ = a₁·Δb + Δa·b₂
A₁@B₁ − A₂@B₂ = A₁@ΔB + ΔA@B₂
Each of the four products is delegated back to the interpreter, so whatever
relaxation the domain uses for ``Expr × Expr`` (dense zonotope matmul, sparse
polynomial matmul, …) applies here unchanged. When one factor is a constant
pair — the usual case, two networks' weights joined by
:func:`~boundlab.diff.ops.diff_pair` — every product is exact and the
difference carries no relaxation error at all.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
from boundlab import Expr, Zeros
from boundlab.ibp import Bias
from boundlab.interp import Interpreter, OpHandler
from boundlab.utils import is_statically_zero
from boundlab.diff.expr import DiffExpr2, DiffExpr3
from .bounds import tighten_diff
def _is_zero(value: Any) -> bool:
"""Whether *value* is the constant zero (so a product with it can be skipped)."""
if isinstance(value, Zeros):
return True
if isinstance(value, Bias):
return is_statically_zero(value.arr)
if isinstance(value, torch.Tensor):
return is_statically_zero(value)
return False
def _unwrap(value: Any) -> Any:
"""Hand constants to the interpreter as plain tensors.
``Bias @ Bias`` has no handler in the standard domains (weights normally
arrive as ONNX initializers, i.e. tensors); unwrapping keeps constant
operands on the exact tensor path instead of the relaxed abstract one.
"""
return value.arr if isinstance(value, Bias) else value
def _zero_like(value: Any) -> Expr:
shape_dtype = value.shape_dtype
return Bias(
torch.zeros(tuple(shape_dtype.shape), dtype=shape_dtype.dtype)
)
[docs]
def diff_bilinear(interp_op, a: Any, b: Any, tighten: bool = True) -> DiffExpr3:
"""Apply ``interp_op`` (``mul`` or ``matmul``) to a differential pair."""
a, b = a.to(DiffExpr3), b.to(DiffExpr3)
out_x = interp_op(_unwrap(a.x), _unwrap(b.x))
out_y = interp_op(_unwrap(a.y), _unwrap(b.y))
terms = []
if not _is_zero(b.diff):
terms.append(interp_op(_unwrap(a.x), _unwrap(b.diff)))
if not _is_zero(a.diff):
terms.append(interp_op(_unwrap(a.diff), _unwrap(b.y)))
if terms:
out_diff = terms[0] if len(terms) == 1 else terms[0] + terms[1]
else:
out_diff = _zero_like(out_x)
if tighten:
out_diff = tighten_diff(out_x - out_y, out_diff)
return DiffExpr3(out_x, out_y, out_diff)
[docs]
def diff_pairwise(interp_op, a: Any, b: Any) -> DiffExpr2:
"""Apply ``interp_op`` to two branch pairs, keeping them independent."""
a = a if isinstance(a, DiffExpr2) else DiffExpr2(a, a)
b = b if isinstance(b, DiffExpr2) else DiffExpr2(b, b)
return DiffExpr2(
interp_op(_unwrap(a.x), _unwrap(b.x)),
interp_op(_unwrap(a.y), _unwrap(b.y)),
)
class _DiffBinary(OpHandler):
"""Common dispatch for the differential ``mul`` / ``matmul`` handlers."""
tighten: bool
def condition(self, x: Any, y: Any, **kwargs: Any) -> bool:
del kwargs
return (
isinstance(x, Expr)
and isinstance(y, Expr)
and not isinstance(x, Zeros)
and not isinstance(y, Zeros)
and (
isinstance(x, (DiffExpr2, DiffExpr3))
or isinstance(y, (DiffExpr2, DiffExpr3))
)
)
def _op(self, interp: Interpreter):
raise NotImplementedError
def handle(self, interp: Interpreter, x: Any, y: Any, **kwargs: Any) -> Expr:
del kwargs
operation = self._op(interp)
# Two pairs stay a pair: nothing is lost, and the difference is only
# materialised once a non-linearity needs it.
if not isinstance(x, DiffExpr3) and not isinstance(y, DiffExpr3):
return diff_pairwise(operation, x, y)
return diff_bilinear(operation, x, y, self.tighten)
[docs]
@dataclass
class DiffMul(_DiffBinary):
"""Element-wise product of differential expressions."""
op: str = "mul"
tighten: bool = True
def _op(self, interp: Interpreter):
return interp.mul
[docs]
@dataclass
class DiffMatmul(_DiffBinary):
"""Matrix product of differential expressions."""
op: str = "matmul"
tighten: bool = True
def _op(self, interp: Interpreter):
return interp.matmul
__all__ = ["DiffMatmul", "DiffMul", "diff_bilinear", "diff_pairwise"]