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