Source code for boundlab.ibp.max


"""Rewrites of ``max(x, const)`` into operations the domains already bound."""

from typing import Any

import torch

from boundlab import Expr, ExprGroup
from boundlab.ibp import Bias
from boundlab.interp import Interpreter, OpHandler

[docs] class MaxWithConstBiased(OpHandler): """Shift the ``Bias`` center out first: ``max(c + i, y) = c + max(i, y - c)`` — exact, and strictly better than relaxing the whole value.""" op = "max"
[docs] def condition(self, x, y, **kwargs): del kwargs if isinstance(x, Expr) and Bias in x.classset() and not isinstance(y, Expr): return True if isinstance(y, Expr) and Bias in y.classset() and not isinstance(x, Expr): return True return False
[docs] def handle( self, interp: Interpreter, x: Any, y: Any, **kwargs, ) -> Expr: del kwargs if not isinstance(x, Expr): x, y = y, x bias, inner = x.split(Bias) return interp.max(inner, y - bias.arr) + bias
[docs] class MaxWithConst2Relu(OpHandler): """``max(x, y) = relu(x - y) + y`` when one side is constant, so ``max`` inherits whatever ReLU relaxation the domain registered.""" 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)
# Splitting the bias off first is strictly better, so it wins whenever both # conditions match (Expr-with-Bias vs. constant). MaxWithConstBiased.overrides = [MaxWithConst2Relu] __all__ = [ "MaxWithConst2Relu", "MaxWithConstBiased", ]