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