r"""Polynomial (Legendre-projection) relaxations of scalar activations.
Instead of an affine enclosure, approximate :math:`f` on :math:`[c - a, c + a]`
by a degree-\ :math:`n` polynomial and bound the residual. The polynomial is
the truncated Legendre series
.. math::
f(c + a u) \approx \sum_{i \le n} l_i\, a^i\, P_i(u),
\qquad
l_i = \frac{2i + 1}{2 a^i} \int_{-1}^{1} P_i(u)\, f(c + a u)\, du ,
the :math:`L^2`-optimal degree-\ :math:`n` approximation on the interval
(Legendre polynomials are orthogonal on :math:`[-1, 1]`). The residual
``f - p`` is bounded soundly by :func:`~boundlab.ops.unary_fn_opt.spline3_ibp`
and becomes interval noise; an optional Adam loop then perturbs the
coefficients per element and keeps whichever iterate certified the smallest
residual. The polynomial itself is applied through the domain's ``poly``
handler — exactly, for a sparse polynomial input.
"""
from abc import abstractmethod
from dataclasses import dataclass, field
import math
from typing import Any, cast
import torch
from boundlab import Expr, utils
from boundlab.ibp import Noise
from boundlab.interp import OpHandler
from boundlab.ops.legendre import legendre_polynomial
from boundlab.ops.unary_fn_opt import (
quartic_bound_exp,
quartic_bound_reciprocal_pos,
spline3_ibp,
)
[docs]
class Poly2Mul(OpHandler):
"""Evaluate a concrete-coefficient polynomial by Horner's rule, each step
one ``mul`` and one ``add`` through the enclosing domain."""
op = "poly"
[docs]
def condition(self, coeffs, x, **kwargs: Any) -> bool:
del kwargs
return isinstance(coeffs, list)
[docs]
def handle(self, interp, coeffs, x, **kwargs: Any) -> "Expr":
assert len(coeffs) > 0
acc = coeffs[-1]
for coeff in reversed(coeffs[:-1]):
acc = interp.mul(x, acc) + coeff
return acc
[docs]
class Poly2Square(OpHandler):
"""Quadratic special case routed through the domain's ``square`` handler:
``a x^2 + b x + c`` with a single non-linear step."""
op = "poly"
[docs]
def condition(self, coeffs, x, **kwargs: Any) -> bool:
del kwargs
return isinstance(coeffs, list) and len(coeffs) == 3
[docs]
def handle(self, interp, coeffs, x, **kwargs: Any) -> "Expr":
assert len(coeffs) == 3
return interp.square(x) * coeffs[2] + x * coeffs[1] + coeffs[0]
[docs]
@dataclass(frozen=True)
class AdamConfig:
"""Hyper-parameters for the optional per-element coefficient refinement
(Adam with annealed gradient noise)."""
learning_rate: float = 3e-3
beta1: float = 0.9
beta2: float = 0.999
epsilon: float = 1e-8
noise_eta: float = 1e-3
noise_gamma: float = 0.55
seed: int = 0
[docs]
class LegendreHandler(OpHandler):
"""Base for activations relaxed by Legendre projection.
Subclasses provide the scalar function (:meth:`fn`), its closed-form
Legendre coefficients (:meth:`legendre_coeffs`), and a fourth-derivative
bound (:meth:`quartic_bound`) for the sound residual evaluation. The
handler then emits ``poly(coeffs, x) + Noise(residual)``.
"""
order: int
nintvl: int
optimizer: AdamConfig | None
opt_iters: int
[docs]
@abstractmethod
def fn(self, x: torch.Tensor) -> torch.Tensor:
"""The scalar function being approximated (elementwise, differentiable)."""
...
[docs]
@abstractmethod
def quartic_bound(self, lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor:
"""Elementwise bound on ``|fn''''|`` over ``[lb, ub]`` (for ``spline3_ibp``)."""
...
[docs]
@abstractmethod
def legendre_coeffs(self, c: torch.Tensor, hw: torch.Tensor, **kwargs: Any) -> list[torch.Tensor]:
r"""
Compute the Legendre coefficients for center ``c`` and half-width ``hw``.
$$l_i = \frac{2i+1}{2 a^i} \int_{-1}^{1} P_i(x) f(a x + b) dx$$
with ``a = hw``, ``b = c``, so that
``f(x) \approx \sum_i l_i a^i P_i((x - b) / a)``. Must be finite for
``a == 0`` (where ``l_i`` tends to the Taylor limit).
"""
...
[docs]
def approx_lbub(self, polyc: utils.Polynomial, c: torch.Tensor, hw: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Elementwise range of the residual ``fn - p`` over ``[c - hw, c + hw]``.
The interpolant ``p`` has degree <= 3, so the residual's fourth
derivative is ``fn''''`` and ``quartic_bound`` applies unchanged.
"""
def residual(x: torch.Tensor) -> torch.Tensor:
return self.fn(x) - polyc(x - c)
polyc_d4 = polyc.derivative(4)
def quartic_bound(lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor:
hwx = (ub - lb) / 2
cx = (lb + ub) / 2
l, u = polyc_d4.ibp(cx - c, hwx)
return self.quartic_bound(lb, ub) + torch.maximum(l.abs(), u.abs())
return spline3_ibp(residual, c - hw, c + hw, quartic_bound, nintvl=self.nintvl)
[docs]
def optimized_polynomial(self, c: torch.Tensor, hw: torch.Tensor, **kwargs: Any) -> tuple[list[torch.Tensor], Noise]:
"""Legendre coefficients plus a sound symmetric bound on the residual.
The residual's center is folded into ``li[0]`` so the returned noise is
the tightest symmetric enclosure of ``fn - p``.
Non-finite values saturate soundly instead of leaking NaN into the
bounds: bad coefficients (overflow, out-of-domain inputs) are zeroed
before the residual is bounded against them, and an overflowing
residual bound becomes an infinite noise with a finite center.
With an ``optimizer``, every element independently keeps the
coefficients of the tightest residual width seen anywhere along the
trajectory — the closed form is iterate 0, so no element ever ends
looser than it — and the noise is recomputed from that selection.
"""
li = self.legendre_coeffs(c, hw, **kwargs)
polyc = legendre_polynomial(li, hw)
polyc = utils.Polynomial([
torch.where(torch.isfinite(coefficient), coefficient, 0.0)
for coefficient in polyc.coeffs
])
if self.order >= 1 and self.optimizer:
delta = [torch.zeros_like(c) for _ in polyc.coeffs[1:]]
def objective(ds: list[torch.Tensor]) -> tuple[torch.Tensor, torch.Tensor]:
p = polyc + utils.Polynomial([0.0] + ds)
lo, hi = self.approx_lbub(p, c, hw)
w = hi - lo
return torch.where(torch.isfinite(w), w, 0.0).sum(), w
def scalar_objective(ds: list[torch.Tensor]) -> torch.Tensor:
return objective(ds)[0]
# Elementwise best-iterate selection: each input keeps the delta
# of the tightest width seen so far. Non-finite widths (diverged
# or NaN deltas) compare as inf and are never selected, and the
# summed objective's per-element compromises (or a noisy
# optimizer's jitter) cannot leak into the result.
best, best_w = delta, torch.full_like(c, torch.inf)
def select_best(ds: list[torch.Tensor], w: torch.Tensor) -> None:
nonlocal best, best_w
w = torch.where(torch.isfinite(w), w, torch.inf)
improved = w < best_w
best = [torch.where(improved, d, b) for d, b in zip(ds, best)]
best_w = torch.minimum(best_w, w)
config = self.optimizer
delta = [value.detach().requires_grad_(True) for value in delta]
first_moment = [torch.zeros_like(value) for value in delta]
second_moment = [torch.zeros_like(value) for value in delta]
generator = torch.Generator(
device=torch.get_default_device()
).manual_seed(config.seed)
for step in range(1, self.opt_iters + 1):
value, w = objective(delta)
select_best([item.detach() for item in delta], w.detach())
gradients = torch.autograd.grad(value, delta)
noise_scale = math.sqrt(config.noise_eta / step**config.noise_gamma)
next_delta = []
for index, (item, gradient) in enumerate(zip(delta, gradients)):
gradient = gradient + noise_scale * torch.randn(
gradient.shape,
dtype=gradient.dtype,
generator=generator,
)
first_moment[index] = (
config.beta1 * first_moment[index]
+ (1 - config.beta1) * gradient
)
second_moment[index] = (
config.beta2 * second_moment[index]
+ (1 - config.beta2) * gradient.square()
)
first_hat = first_moment[index] / (1 - config.beta1**step)
second_hat = second_moment[index] / (1 - config.beta2**step)
updated = item - config.learning_rate * first_hat / (
torch.sqrt(second_hat) + config.epsilon
)
next_delta.append(updated.detach().requires_grad_(True))
delta = next_delta
# The final iterate was never scored inside the loop; saturate its
# non-finite entries and let it compete.
final = [torch.where(torch.isfinite(d), d, 0.0).detach() for d in delta]
select_best(final, objective(final)[1])
polyc = polyc + utils.Polynomial([0.0] + best)
lo, hi = self.approx_lbub(polyc, c, hw)
mid = (lo + hi) / 2
half = (hi - lo) / 2
polyc.coeffs[0] = polyc.coeffs[0] + torch.where(torch.isfinite(mid), mid, 0.0)
polyc0 = [torch.as_tensor(coefficient) for coefficient in polyc.coeffs]
return polyc0, Noise(torch.where(torch.isfinite(half), half, torch.inf), self.op)
[docs]
def condition(self, *args: Any, **kwargs: Any) -> bool:
return len(args) == 1 and isinstance(args[0], Expr)
[docs]
def handle(self, interp, x: "Expr", **kwargs: Any) -> "Expr":
c, hw = x.chw()
polyc, noise = self.optimized_polynomial(c, hw, **kwargs)
return interp.poly(polyc, x - c) + noise
[docs]
@dataclass
class Relu(LegendreHandler):
r'''Legendre relaxation of ReLU (orders 1–6, closed-form coefficients).
On a crossing interval the projection integrals split at the kink; the
resulting coefficients are polynomial in :math:`c/a`, evaluated in closed
form per order. ``quartic_bound`` is zero — ReLU is piecewise linear, so
the spline residual evaluation is exact up to the kink handling — and the
optional Adam refinement tightens the per-element coefficients further.
'''
op: str = "relu"
order: int = 2
nintvl: int = None # type: ignore default to 1 for order <= 3, 16 for order >= 4
# Swept on a (c, hw) grid: diagonal optimizers decouple the summed
# per-element widths (lbfgs's shared line search cannot), so adam wins;
# lr 3e-3 at 500 iterations was the tightest configuration. Annealed
# gradient noise explores the nonconvex order-4 bound landscape (a single
# noisy trajectory beat 16 random restarts there) and is harmless for the
# convex orders <= 3; the fixed seed keeps results deterministic.
optimizer: AdamConfig | None = field(default_factory=AdamConfig)
opt_iters: int = 500
def __post_init__(self):
assert 1 <= self.order <= 6
if self.nintvl is None:
self.nintvl = 1 if self.order <= 3 else 16
[docs]
def fn(self, x: torch.Tensor) -> torch.Tensor:
return x.clamp(min=0.0)
[docs]
def quartic_bound(self, lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor:
return torch.zeros_like(lb)
[docs]
def legendre_coeffs(self, c: torch.Tensor, hw: torch.Tensor, **kwargs: Any) -> list[torch.Tensor]:
"""
Implemented based on::
l(i) = (2n+1)/(2a^n) integrate(relu(a x + b) legendrep(n,x), {x,-1,1})
l(0) = -1/4 (cb - 1) (a + 2 b + a cb)
l(1) = -1/4 (cb2m1(1) * 3 b + 2 (cb - 1) (1 + cb + cb^2))
l(2) = -5/16 cb2m1(2) (a + 4 b cb + 3 a cb^2)
l(3) = -7/16 cb2m1(3) (-b + 5 b cb^2 + 4 a cb^3)
l(4) = -3/32 cb2m1(4) (-a - 18 b cb - 10 a cb^2 + 42 b cb^3 + 35 a cb^4)
l(5) = -11/32 cb2m1(5) (b - 14 b cb^2 - 10 a cb^3 + 21 b cb^4 + 18 a cb^5)
l(6) = -13/256 cb2m1(6) (a + 40 b cb + 21 a cb^2 - 240 b cb^3 - 189 a cb^4 + 264 b cb^5 + 231 a cb^6)
where cb = clip(-b/a, -1, 1)
cb2m1(n) = (cb^2 - 1) / a^n
``cb`` and ``cb2m1`` are handled safely for ``a == 0``.
"""
del kwargs
a, b = hw, c
crossing = b.abs() < a
cb = torch.where(crossing, (-b / a).clamp(-1.0, 1.0), torch.sign(-b))
def cb2m1(n: int) -> torch.Tensor:
return torch.where(crossing, (cb * cb - 1.0) / a**n, 0.0)
li = [
-(cb - 1) * (a + 2 * b + a * cb) / 4,
-(cb2m1(1) * 3 * b + 2 * (cb - 1) * (1 + cb + cb * cb)) / 4,
-5 / 16 * cb2m1(2) * (a + 4 * b * cb + 3 * a * cb**2),
-7 / 16 * cb2m1(3) * (-b + 5 * b * cb**2 + 4 * a * cb**3),
-3 / 32 * cb2m1(4) * (-a - 18 * b * cb - 10 * a * cb**2 + 42 * b * cb**3 + 35 * a * cb**4),
-11 / 32 * cb2m1(5) * (b - 14 * b * cb**2 - 10 * a * cb**3 + 21 * b * cb**4 + 18 * a * cb**5),
-13 / 256 * cb2m1(6) * (a + 40 * b * cb + 21 * a * cb**2 - 240 * b * cb**3 - 189 * a * cb**4 + 264 * b * cb**5 + 231 * a * cb**6),
]
return li[: self.order + 1]
[docs]
def approx_lbub(self, polyc: utils.Polynomial, c: torch.Tensor, hw: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
lbx, ubx = c - hw, c + hw
l4, u4 = polyc.derivative(4).ibp(torch.zeros_like(c), hw)
quartic_bound = torch.broadcast_to(
torch.maximum(l4.abs(), u4.abs()), lbx.shape
)
mid = torch.zeros_like(lbx).clamp(min=lbx, max=ubx)
lo1, hi1 = spline3_ibp(lambda x: torch.as_tensor(-polyc(x - c)),
lbx, mid, quartic_bound, nintvl=self.nintvl)
lo2, hi2 = spline3_ibp(lambda x: x - polyc(x - c),
mid, ubx, quartic_bound, nintvl=self.nintvl)
neg, pos = lbx <= 0, ubx >= 0 # which segments actually exist
inf = torch.inf
lo = torch.minimum(torch.where(neg, lo1, inf), torch.where(pos, lo2, inf))
hi = torch.maximum(torch.where(neg, hi1, -inf), torch.where(pos, hi2, -inf))
return lo, hi
[docs]
class Tanh(LegendreHandler):
"""Placeholder for a Legendre tanh relaxation (not yet implemented)."""
__all__ = [
"AdamConfig",
"LegendreHandler",
"Poly2Mul",
"Poly2Square",
"Relu",
"Tanh",
"eval_legendre_form",
]