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