Source code for boundlab.diff.zono3.bilinear

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