Source code for boundlab.polysp.legendre

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] def eval_legendre_form(li: list[torch.Tensor], c: torch.Tensor, hw: torch.Tensor, x: torch.Tensor) -> torch.Tensor: """Evaluate ``sum_i li[i] * hw**i * P_i((x - c) / hw)``, safe at ``hw == 0``.""" u = (x - c) / torch.where(hw > 0, hw, 1.0) ps = [ torch.ones_like(u), u, (3 * u**2 - 1) / 2, (5 * u**3 - 3 * u) / 2, (35 * u**4 - 30 * u**2 + 3) / 8, (63 * u**5 - 70 * u**3 + 15 * u) / 8, (231 * u**6 - 315 * u**4 + 105 * u**2 - 5) / 16, ] assert len(li) <= len(ps), "eval_legendre_form supports orders up to 6" return utils.sum(l * hw**i * p for i, (l, p) in enumerate(zip(li, ps)))
[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", ]