Source code for boundlab.ops.unary_fn_opt

r'''Certified elementwise function ranges via cubic Hermite splines.

:func:`spline3_ibp` bounds an arbitrary differentiable scalar function over a
box by interpolating it piecewise with cubics (values and derivatives at the
subinterval endpoints) and covering the interpolation error with the
two-point Hermite remainder

.. math::

   f(x) - p(x) = \frac{f^{(4)}(\xi)}{4!}\,(x - x_0)^2 (x - x_1)^2,
   \qquad |f - p| \le \frac{M_4\, h^4}{384},

where :math:`M_4` bounds the fourth derivative.  The cubic's exact extrema
(endpoints plus the real roots of its derivative) are closed-form, so the
whole bound is differentiable and vectorized — the residual oracle behind the
Legendre relaxations.  The ``quartic_bound_*`` helpers supply :math:`M_4` for
the standard activations.
'''

from collections.abc import Callable
import torch

QuarticBoundFn = Callable[[torch.Tensor, torch.Tensor], torch.Tensor]
'''An elementwise fourth-derivative bound: ``(lb, ub) -> M4 >= max |f''''|``.'''

[docs] def spline3_ibp( fn: Callable[[torch.Tensor], torch.Tensor], lb: torch.Tensor, ub: torch.Tensor, quartic_bound: torch.Tensor | QuarticBoundFn, nintvl: int = 8 ) -> tuple[torch.Tensor, torch.Tensor]: """ Bound ``fn`` over the elementwise interval ``[lb, ub]`` via cubic Hermite spline interpolation. Each interval is split into ``nintvl`` subintervals of width ``h``. On each one, ``fn`` is interpolated by the cubic matching the values and derivatives (via ``torch.func.jvp``) at both endpoints; the exact extrema of that cubic (endpoints and real roots of its derivative) are candidates for the extrema of ``fn``. The interpolation error is covered by the two-point Hermite remainder fn(x) - p(x) = fn''''(xi) / 4! * (x - x0)^2 * (x - x1)^2, whose magnitude is at most ``M4 * h**4 / 384``. ``quartic_bound`` supplies ``M4 >= max |fn''''|``: either an array bounding it over the whole ``[lb, ub]``, or a callable evaluated per subinterval at its endpoints, so each subinterval carries its own (tighter) remainder. ``fn`` (and a callable ``quartic_bound``) are evaluated under ``torch.vmap`` over the sample axis, so they always receive arrays of exactly ``lb``'s shape and may therefore implement a different scalar function per element (e.g. selected by a mask or per-element parameters of that shape). Returns sound elementwise ``(lower, upper)`` bounds on ``fn`` over ``[lb, ub]`` (exact for polynomials of degree <= 3 with ``M4 = 0``). """ assert lb.shape == ub.shape, "Lower and upper bounds must have the same shape." # Fix the sample dtype to the input's: under default dtype configuration a default # float64 linspace would silently promote the whole computation. ts = torch.linspace(0.0, 1.0, nintvl + 1, dtype=lb.dtype) ts = ts.reshape((nintvl + 1,) + (1,) * lb.ndim) xs = lb + (ub - lb) * ts # (nintvl + 1, *shape) ys, dys = torch.vmap( lambda x: torch.func.jvp(fn, (x,), (torch.ones_like(x),)) )(xs) if callable(quartic_bound): quartic_bound = torch.vmap(quartic_bound)(xs[:-1], xs[1:]) # (nintvl, *shape) assert quartic_bound.shape == (nintvl, *lb.shape), \ "Quartic bound must be elementwise over the subinterval endpoints." else: assert quartic_bound.shape == lb.shape, "Quartic bound must have the same shape as the bounds." h = (ub - lb) / nintvl y0, y1 = ys[:-1], ys[1:] # (nintvl, *shape) m0, m1 = dys[:-1] * h, dys[1:] * h # Hermite cubic on u in [0, 1]: p(u) = a*u^3 + b*u^2 + c*u + d. a = 2 * (y0 - y1) + m0 + m1 b = 3 * (y1 - y0) - 2 * m0 - m1 c = m0 d = y0 def p(u: torch.Tensor) -> torch.Tensor: return ((a * u + b) * u + c) * u + d # p'(u) = A*u^2 + B*u + C; its real roots in [0, 1] are the only interior # extremum candidates. Missing a real root here loses an extremum of p # and returns bounds that are too tight, so the roots must stay accurate # when A (or B) is a rounding-level artifact of an exactly-quadratic p: # use the cancellation-free form q = -(B + sign(B) sqrt(disc)) / 2 with # roots q/A and C/q. Degenerate cases only add spurious candidates, # which are harmless: every candidate is p evaluated inside [0, 1]. A, B, C = 3 * a, 2 * b, c disc = B * B - 4 * A * C # Double-where keeps the gradient finite at disc <= 0 (sqrt' blows up at # zero and would poison optimizers differentiating through these bounds); # the value is identical to sqrt(max(disc, 0)). disc_pos = disc > 0 sqrt_disc = torch.where(disc_pos, torch.sqrt(torch.where(disc_pos, disc, 1.0)), 0.0) sign_b = torch.where(B >= 0, 1.0, -1.0) q = -(B + sign_b * sqrt_disc) / 2 a_ok = A != 0 q_ok = q != 0 u1 = torch.where(a_ok, q / torch.where(a_ok, A, 1.0), 0.0) u2 = torch.where(q_ok, C / torch.where(q_ok, q, 1.0), 0.0) u1 = u1.clamp(0.0, 1.0) u2 = u2.clamp(0.0, 1.0) candidates = torch.stack([y0, y1, p(u1), p(u2)]) # (4, nintvl, *shape) pmin = candidates.amin(0) # (nintvl, *shape) pmax = candidates.amax(0) # Each subinterval carries its own remainder; the bound over [lb, ub] is # the envelope of the per-subinterval bounds. err = quartic_bound.abs() * h**4 / 384 return (pmin - err).amin(0), (pmax + err).amax(0)
[docs] def quartic_bound_exp(lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor: """Elementwise bound on ``|exp''''|`` over [lb, ub].""" return torch.exp(torch.maximum(lb, ub))
[docs] def quartic_bound_tanh(lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor: """Elementwise bound on ``|tanh''''|`` over [lb, ub].""" abs_lb = torch.where( (lb <= 0.0) & (ub >= 0.0), 0.0, torch.minimum(lb.abs(), ub.abs()), ) return (7 * torch.exp(-abs_lb.abs())).clamp(max=4.2)
[docs] def quartic_bound_reciprocal_pos(lb: torch.Tensor, ub: torch.Tensor) -> torch.Tensor: """Elementwise bound on ``|reciprocal''''|`` over [lb, ub]; +inf where the interval is not strictly positive (1/x is unbounded there).""" pos = lb > 0 return torch.where(pos, 25 / torch.where(pos, lb, 1.0) ** 5, torch.inf)
__all__ = [ "QuarticBoundFn", "quartic_bound_exp", "quartic_bound_reciprocal_pos", "quartic_bound_tanh", "spline3_ibp", ]