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