Source code for boundlab.zono.linearizers

r"""Sound affine enclosures of scalar non-linearities.

A *linearizer* maps elementwise input bounds :math:`[l, u]` to a
:class:`LinearBounds` :math:`(\lambda, \mu, \beta)` certifying

.. math::

   \lambda x + \mu - \beta \;\le\; f(x) \;\le\; \lambda x + \mu + \beta
   \qquad \text{for all } x \in [l, u].

Applying it keeps the input's symbolic structure — ``x * slope + bias`` is a
linear map — and pays only the residual :math:`\pm\beta` as fresh
``Noise``.  Each linearizer below chooses :math:`\lambda` to (approximately)
minimize :math:`\beta`, which for a fixed slope is half the range of
:math:`f(x) - \lambda x` over :math:`[l, u]` (the Chebyshev center of the
residual).
"""

from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Callable, overload

import torch

from boundlab import Expr
from boundlab.ibp import Bias, Intervals, Noise
from boundlab.interp import OpHandler, handler_fn

[docs] class MaxWithConst2Relu(OpHandler): """``max(x, c) = relu(x - c) + c`` when one operand is constant, so ``max`` inherits the domain's ReLU relaxation.""" op = "max"
[docs] def condition(self, x, y, **kwargs): del kwargs return not (isinstance(x, Expr) and isinstance(y, Expr))
[docs] def handle(self, interp, x, y, **kwargs): del kwargs if isinstance(x, Expr) and not isinstance(y, Expr): return interp.relu(x - y) + y elif isinstance(y, Expr) and not isinstance(x, Expr): return interp.relu(y - x) + x elif isinstance(x, torch.Tensor) and isinstance(y, torch.Tensor): return torch.maximum(x, y)
[docs] @dataclass class LinearBounds: r"""A sound enclosure ``f(x) ∈ slope·x + bias ± error`` over the queried box. Calling it applies the enclosure to an expression: ``x * slope + bias + Noise(error, name)`` — the affine part exact, the residual as fresh interval noise tagged with ``name``. """ name: str slope: torch.Tensor bias: torch.Tensor error: torch.Tensor
[docs] def __call__(self, x: Expr) -> Expr: from boundlab.zono import Zono if not isinstance(self.slope, torch.Tensor) and not isinstance(self.bias, torch.Tensor) and not isinstance(self.error, torch.Tensor): if not torch.compiler.is_compiling(): assert torch.isfinite(self.slope).all(), ( f"Non-finite slope in {self.name}: {self.slope}" ) assert torch.isfinite(self.bias).all(), ( f"Non-finite bias in {self.name}: {self.bias}" ) assert torch.isfinite(self.error).all(), ( f"Non-finite error in {self.name}: {self.error}" ) return ( x * self.slope + self.bias + Noise(self.error, self.name) )
[docs] def linearizer_fn(op: str): """Wrap a bounds function ``(lb, ub) -> LinearBounds`` as the zonotope handler for ``op``; also callable directly on bounds (and optionally an expression) for reuse by other domains.""" _op = op def decorator(linearizer: Callable[[torch.Tensor, torch.Tensor], LinearBounds]): @dataclass class LinearizerHandler(OpHandler): op: str = _op apply_bound: Callable[[LinearBounds, Expr], Expr] = lambda bounds, x: bounds(x) def condition(self, x: Expr, **kwargs): del kwargs return not x.classset().issubset({Noise, Bias}) def handle(self, interp, x: Expr, **kwargs) -> Expr: del kwargs return self.apply_bound(linearizer(*x.lbub()), x) @overload def __call__(self, l: torch.Tensor, u: torch.Tensor, x: None = None) -> LinearBounds: ... @overload def __call__(self, l: torch.Tensor, u: torch.Tensor, x: Expr) -> Expr: ... def __call__(self, l: torch.Tensor, u: torch.Tensor, x: Expr | None = None): if x is not None: return self.apply_bound(linearizer(l, u), x) else: return linearizer(l, u) LinearizerHandler.__name__ = linearizer.__name__ LinearizerHandler.__qualname__ = linearizer.__qualname__ LinearizerHandler.__doc__ = ( linearizer.__doc__ or f"Zonotope linearizer for ``{_op}``: a sound ``slope * x + bias ± error`` enclosure." ) return LinearizerHandler return decorator
[docs] @linearizer_fn("relu") def Relu(lower: torch.Tensor, upper: torch.Tensor) -> LinearBounds: r"""Triangle relaxation of ReLU. The chord slope :math:`\lambda = (\mathrm{relu}(u) - \mathrm{relu}(l)) / (u - l)` makes the residual :math:`\mathrm{relu}(x) - \lambda x` range over :math:`[0,\; \mathrm{relu}(u) - \lambda u]` on a crossing interval (extremes at the kink and at the endpoints), so centering gives :math:`\mu = \beta = (\mathrm{relu}(u) - \lambda u)/2`. On stable intervals :math:`\lambda \in \{0, 1\}` and :math:`\beta = 0` — exact. This slope choice minimizes :math:`\beta` among all sound affine enclosures (the DeepZ/triangle optimum). """ relu_lower = lower.clamp(min=0.0) relu_upper = upper.clamp(min=0.0) slope = (relu_upper - relu_lower) / (upper - lower + 1e-30) error = (relu_upper - slope * upper) * 0.5 return LinearBounds( name="relu", slope=slope, bias=(relu_upper + relu_lower - slope * (upper + lower)) * 0.5, error=error, )
[docs] @linearizer_fn("exp") def Exp(lb: torch.Tensor, ub: torch.Tensor) -> LinearBounds: r"""Chord-slope enclosure of :math:`e^x`. With the chord slope :math:`\lambda = (e^u - e^l)/(u - l)` (mean value: :math:`\lambda = e^\xi` for some interior :math:`\xi`), convexity puts the residual's maximum at the endpoints — where the chord touches, value :math:`U = e^u - \lambda u` — and its minimum at the tangency point :math:`x^* = \log\lambda`, value :math:`L = \lambda(1 - \log\lambda)`. Centering yields :math:`\mu = (U + L)/2`, :math:`\beta = (U - L)/2`. Near-degenerate intervals switch to the midpoint derivative, and any non-finite element falls back to the sound interval enclosure :math:`e^u/2 \pm e^u/2 \supseteq [0, e^u]`. """ shape = lb.shape lb = lb.reshape(-1) ub = ub.reshape(-1) expm1_lb = torch.expm1(lb) expm1_ub = torch.expm1(ub) slope = (expm1_ub - expm1_lb) / (ub - lb) slope = torch.where(ub - lb >= 1e-5, slope, torch.exp((ub + lb) / 2)) slope = slope.clamp(min=torch.finfo(slope.dtype).tiny) slope_point = torch.log(slope) U = 1 + torch.maximum(expm1_ub - slope * ub, expm1_lb - slope * lb) L = slope * (1 - slope_point) beta = (U - L) / 2 mu = (U + L) / 2 bad = ~(torch.isfinite(mu) & torch.isfinite(beta) & torch.isfinite(slope)) slope = torch.where(bad, torch.zeros_like(slope), slope) mu = torch.where(bad, (1 + expm1_ub) / 2, mu) beta = torch.where(bad, (1 + expm1_ub) / 2, beta) slope = slope.reshape(*shape) mu = mu.reshape(*shape) beta = beta.reshape(*shape) return LinearBounds( name="exp", slope=slope, bias=mu, error=beta, )
[docs] @linearizer_fn("tanh") def Tanh(lower: torch.Tensor, upper: torch.Tensor) -> LinearBounds: r"""Minimal-derivative enclosure of :math:`\tanh` (DeepZ-style). The slope :math:`\lambda = \min(1 - \tanh^2 l,\; 1 - \tanh^2 u)` never exceeds the true derivative anywhere in :math:`[l, u]` (:math:`\tanh'` is unimodal), so the residual :math:`\tanh(x) - \lambda x` is monotone and attains its extremes at the endpoints; :math:`\mu` and :math:`\beta` center that endpoint range. """ tanh_lower = torch.tanh(lower) tanh_upper = torch.tanh(upper) slope = torch.minimum(1.0 - tanh_lower**2, 1.0 - tanh_upper**2) bias = (tanh_upper + tanh_lower - slope * (upper + lower)) * 0.5 error = ((tanh_upper - tanh_lower - slope * (upper - lower)) * 0.5).abs() return LinearBounds( name="tanh", slope=slope, bias=bias, error=error )
[docs] @linearizer_fn("reciprocal") def Reciprocal(lb: torch.Tensor, ub: torch.Tensor) -> LinearBounds: r"""Tangent-line enclosure of :math:`1/x` on positive intervals. A tangent at :math:`t` has slope :math:`-1/t^2` and, by convexity, lower bounds :math:`1/x` after its own offset; the geometric mean :math:`t = \sqrt{lu}` equalizes the gap at both endpoints (the classical minimal-\ :math:`\beta` choice), guarded by :math:`t \ge u/2 + 0.01` against collapse. The upper offset is the endpoint maximum of :math:`1/x - \lambda x`; centering the two offsets gives :math:`\mu, \beta`. Degenerate (point) intervals return the exact value. """ degen = (ub - lb).abs() < 1e-12 t_crit = torch.sqrt(ub * lb) t_crit2 = 0.5 * ub + 0.01 t_opt = torch.maximum(t_crit, t_crit2) slope = -1.0 / (t_opt**2) val_at_t = 1.0 / t_opt c_lower = val_at_t - slope * t_opt c_upper = torch.maximum( 1.0 / lb - slope * lb, 1.0 / ub - slope * ub, ) mu = 0.5 * (c_upper + c_lower) beta = 0.5 * (c_upper - c_lower) slope = torch.where(degen, torch.zeros_like(slope), slope) mu = torch.where(degen, 1.0 / lb, mu) beta = torch.where(degen, torch.zeros_like(beta), beta.abs()) return LinearBounds( name="reciprocal", slope=slope, bias=mu, error=beta, )
__all__ = [ "Exp", "LinearBounds", "MaxWithConst2Relu", "Reciprocal", "Relu", "Tanh", "linearizer_fn", ]