Source code for boundlab.ibp.elementwise

"""Interval transfer functions for elementwise activations.

Each activation here is monotone, so its exact range over ``[lb, ub]`` is
``[f(lb), f(ub)]``; the handler returns that range as a fresh
``Bias + Noise`` pair.  Exact as an interval — but a *new* interval: any
correlation with the input is lost, which is precisely what the zonotope
linearizers (:mod:`boundlab.zono.linearizers`) avoid.
"""

from __future__ import annotations

from collections.abc import Callable
from dataclasses import dataclass

import torch

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


[docs] def interval_fn(op: str): """Wrap an endpoint-mapping function ``(lb, ub) -> (lb', ub')`` as the interval handler for ``op``; fires only on pure ``Bias + Noise`` input.""" _op = op def decorator( fn: Callable[[torch.Tensor, torch.Tensor], tuple[torch.Tensor, torch.Tensor]], ) -> type[OpHandler]: @dataclass class IntervalFnHandler(OpHandler): op: str = _op def condition(self, x, **kwargs): del kwargs return isinstance(x, Expr) and x.classset().issubset({Bias, Noise}) def handle(self, interp, x: Expr, **kwargs) -> Expr: del interp, kwargs return self(*x.lbub()) def __call__(self, lb: torch.Tensor, ub: torch.Tensor) -> Expr: lb, ub = fn(lb, ub) return Bias((ub + lb) * 0.5) + Noise( (ub - lb) * 0.5, self.op ) IntervalFnHandler.__name__ = fn.__name__ IntervalFnHandler.__qualname__ = fn.__qualname__ IntervalFnHandler.__doc__ = ( fn.__doc__ or f"Interval transfer for ``{_op}``: map both endpoints through the concrete function." ) return IntervalFnHandler return decorator
[docs] @interval_fn("relu") def Relu(lb: torch.Tensor, ub: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Exact interval ReLU: ``[max(lb, 0), max(ub, 0)]`` (monotone).""" return lb.clamp(min=0.0), ub.clamp(min=0.0)
[docs] @interval_fn("exp") def Exp(lb: torch.Tensor, ub: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Exact interval exp: ``[exp(lb), exp(ub)]`` (monotone).""" return torch.exp(lb), torch.exp(ub)
[docs] @interval_fn("tanh") def Tanh(lb: torch.Tensor, ub: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Exact interval tanh: ``[tanh(lb), tanh(ub)]`` (monotone).""" return torch.tanh(lb), torch.tanh(ub)
[docs] @interval_fn("reciprocal") def Reciprocal(lb: torch.Tensor, ub: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Exact interval reciprocal ``[1/ub, 1/lb]`` (monotone decreasing); only sound on intervals that exclude zero.""" return 1.0 / ub, 1.0 / lb
__all__ = [ "Exp", "Reciprocal", "Relu", "Tanh", "interval_fn", ]