diff --git a/docs/API.md b/docs/API.md index 2c85655..ecb6e39 100644 --- a/docs/API.md +++ b/docs/API.md @@ -42,6 +42,31 @@ optimal/tightest interval boxes for every operation. - Scalar division by an interval containing `0` raises `ZeroDivisionError`. - Vector interval division is intentionally not implemented and raises `NotImplementedError`. +## Scalar second derivative of tanh + +`tanh_double_prime_bounds(value: Interval) -> Interval` is available directly +from `intervalnets`, without PyTorch. It returns outward-rounded binary64 bounds +for the mathematical function `tanh''` on a scalar interval: + +```python +from intervalnets import Interval, tanh_double_prime_bounds + +bounds = tanh_double_prime_bounds(Interval(0.5, 1.0)) +# Approximately [-0.769800358919501, -0.639700008449225]. +``` + +The calculation evaluates both endpoints and includes any interior extrema at +`+/-atanh(1/sqrt(3))`. It encloses the exact real range, with a small outward +rounding allowance, rather than repeatedly multiplying dependent intervals. +Zero point intervals stay exactly zero; unbounded intervals use the limits at +infinity. NaN endpoints and nonscalar inputs are rejected. Decimal arithmetic +from the standard library bounds elementary-function and arithmetic errors; +there is no additional dependency. See [the rounding argument](tanh_second_derivative_bounds.md). + +The internal PyTorch adapter retains additional float32 outward padding and +preserves the known sign. This scalar API does not add a network Hessian API +to branches that do not already have one. + ## PyTorch integration (`IntervalTensor` + patching) ### `IntervalTensor` diff --git a/docs/tanh_second_derivative_bounds.md b/docs/tanh_second_derivative_bounds.md new file mode 100644 index 0000000..2221194 --- /dev/null +++ b/docs/tanh_second_derivative_bounds.md @@ -0,0 +1,64 @@ +# Interval bounds for the second derivative of tanh + +For `g(x) = tanh''(x)`, extrema on `[a,b]` occur at the two endpoints +and at any included stationary points `+/-c`, where + +``` +c = atanh(1/sqrt(3)) = log(2 + sqrt(3))/2 +g(-c) = M, g(c) = -M, M = 4/(3*sqrt(3)). +``` + +Indeed, `g'(x) = -2(1-tanh(x)^2)(1-3*tanh(x)^2)`, whose only finite +zeros are `+/-c`. Taking the hull of these candidate values is the exact +real range. It avoids dependency inflation in `-2*T*(1-T*T)`. + +## Machine enclosure argument + +`activations.tanh_double_prime_bounds` encloses this range as follows: + +1. Construct a bracket for `sqrt(3)` with Decimal's correctly rounded square + root and its two neighboring decimal numbers. Directed arithmetic and a + similarly bracketed logarithm give a bracket for `c` and an upper bound + for `M`. Convert these bounds outward to binary64. If an input interval + overlaps a critical-point bracket, include the corresponding global + extremum conservatively. +2. At each finite nonzero endpoint, use `q = exp(-2*abs(x))` and + `abs(g(x)) = 8*q*(1-q)/(1+q)^3`. Bracket the exponent by directed + multiplication. Decimal's exponential is correctly rounded to nearest, + even when a directed context is supplied, so explicitly take the lower + neighbor of the lower exponential and upper neighbor of the upper one. + Intersect this enclosure with `[0,1]`. +3. Every remaining operation uses directed Decimal arithmetic. All factors + are nonnegative: the lower numerator uses `q_lo*(1-q_hi)`, the upper + uses `q_hi*(1-q_lo)`. Divide the lower numerator by the upper denominator + and conversely for the upper bound. Apply the exact sign last. +4. Convert each Decimal bound to binary64, compare its exact Decimal image + against the source bound, and move one binary64 step outward when needed. + This also handles underflow to zero. Near zero, increase decimal precision + to resolve `1-q` without losing relative tightness. +5. For finite `abs(x) >= 400`, use `abs(g(x)) <= 8*exp(-800) < 2**-1074` + and enclose by zero and the smallest positive binary64 subnormal, with + the appropriate sign. The values at zero and the limits at infinity + are exactly zero. Intersect the final range with the rounded global bounds. + +The independent contexts do not modify the caller's decimal context. The +argument relies on the standard library's correctly rounded Decimal +`sqrt`, `ln`, and `exp`, not on an assumed error bound for system `libm`. +See [Python's Decimal documentation](https://docs.python.org/3/library/decimal.html). + +The public scalar bounds enclose the mathematical derivative. They are not +a guarantee about every floating-point/autograd expression for that derivative. +The PyTorch adapter additionally applies the backend's existing float32 +padding, retaining the sign restriction. Its returned bounds can therefore +be wider than the scalar API. Existing affine PZ enclosures are unchanged. + +## Regression coverage + +Tests compare against high-precision positive-exponential reference values, +including both extrema, adjacent floats at the critical points, point +intervals, signed zero, subnormal inputs and outputs, saturated tails, +unbounded intervals, invalid inputs, odd symmetry, and random intervals. +They also verify strict improvement over the former interval product on +representative intervals and independence from the caller's decimal context. +Numerical tests supplement the enclosure argument above; sampled values +alone are not an enclosure proof. diff --git a/src/intervalnets/__init__.py b/src/intervalnets/__init__.py index 62b98ee..9ba2669 100644 --- a/src/intervalnets/__init__.py +++ b/src/intervalnets/__init__.py @@ -1,6 +1,7 @@ """Interval arithmetic utilities for neural network evaluation.""" from .interval import Interval +from .activations import tanh_double_prime_bounds from .polynomial_zonotope import ( PZOneJet, PZTwoJet, @@ -81,6 +82,7 @@ __all__ = [ "Interval", + "tanh_double_prime_bounds", "PolynomialZonotope", "PZOneJet", "PZTwoJet", diff --git a/src/intervalnets/activations.py b/src/intervalnets/activations.py new file mode 100644 index 0000000..5a4eb60 --- /dev/null +++ b/src/intervalnets/activations.py @@ -0,0 +1,102 @@ +"""Scalar activation bounds with explicit control of elementary-function errors.""" + +from __future__ import annotations + +from decimal import Context, Decimal, ROUND_CEILING, ROUND_FLOOR +from math import inf, isinf, isnan, nextafter + +from .interval import Interval + + +def _contexts(precision: int = 60) -> tuple[Context, Context]: + # Do not depend on, or modify, the caller's decimal context. + lower = Context(prec=precision, rounding=ROUND_FLOOR, Emin=-999999, Emax=999999) + upper = Context(prec=precision, rounding=ROUND_CEILING, Emin=-999999, Emax=999999) + return lower, upper + + +def _float_lower(value: Decimal) -> float: + result = float(value) + return nextafter(result, -inf) if Decimal.from_float(result) > value else result + + +def _float_upper(value: Decimal) -> float: + result = float(value) + return nextafter(result, inf) if Decimal.from_float(result) < value else result + + +def _critical_bounds() -> tuple[float, float, float]: + down, up = _contexts() + # Decimal sqrt/ln/exp are correctly rounded, but not necessarily in the + # requested directed mode. Their adjacent decimal numbers bracket them. + root = down.sqrt(Decimal(3)) + root_lo, root_hi = down.next_minus(root), up.next_plus(root) + # atanh(1/sqrt(3)) = log(2 + sqrt(3))/2. + c_lo = down.divide(down.next_minus(down.ln(down.add(2, root_lo))), 2) + c_hi = up.divide(up.next_plus(up.ln(up.add(2, root_hi))), 2) + maximum_hi = up.divide(4, down.multiply(3, root_lo)) + return _float_lower(c_lo), _float_upper(c_hi), _float_upper(maximum_hi) + + +_CRITICAL_LO, _CRITICAL_HI, _MAXIMUM_HI = _critical_bounds() +_MIN_SUBNORMAL = nextafter(0.0, inf) + + +def _tanh_double_prime_point_bounds(x: float) -> tuple[float, float]: + if x == 0.0 or isinf(x): + # The values at infinite endpoints are the limits at infinity. + return 0.0, 0.0 + if abs(x) >= 400.0: + # |tanh''(x)| <= 8*exp(-2*|x|) <= 8*exp(-800) < 2**-1074. + return (-_MIN_SUBNORMAL, 0.0) if x > 0 else (0.0, _MIN_SUBNORMAL) + + magnitude = Decimal.from_float(abs(x)) + # Near zero, 1-exp(-2|x|) loses decimal places. Extra precision keeps + # the enclosure narrow even at binary64 subnormal inputs. + down, up = _contexts(60 + max(0, -magnitude.adjusted())) + exponent_lo = down.multiply(-2, magnitude) + exponent_hi = up.multiply(-2, magnitude) + q_lo = max(Decimal(0), down.next_minus(down.exp(exponent_lo))) + q_hi = min(Decimal(1), up.next_plus(up.exp(exponent_hi))) + + # |tanh''(x)| = 8*q*(1-q)/(1+q)**3, q=exp(-2|x|). + # All factors are nonnegative. Bound each arithmetic operation outward; + # ordinary tanh(x) may have already rounded to +/-1 in the tails. + numerator_lo = down.multiply(8, down.multiply(q_lo, down.subtract(1, q_hi))) + numerator_hi = up.multiply(8, up.multiply(q_hi, up.subtract(1, q_lo))) + base_lo, base_hi = down.add(1, q_lo), up.add(1, q_hi) + denominator_lo = down.multiply(down.multiply(base_lo, base_lo), base_lo) + denominator_hi = up.multiply(up.multiply(base_hi, base_hi), base_hi) + lower = _float_lower(down.divide(numerator_lo, denominator_hi)) + upper = _float_upper(up.divide(numerator_hi, denominator_lo)) + return (-upper, -lower) if x > 0 else (lower, upper) + + +def tanh_double_prime_bounds(value: Interval) -> Interval: + """Enclose the real-valued ``tanh''`` over a scalar interval. + + Extrema occur at the endpoints or at +/-atanh(1/sqrt(3)), with + values -/+4/(3*sqrt(3)). This gives the exact real range before + numerical rounding; the returned binary64 endpoints round outward. + Critical-point membership is conservative if the input overlaps the + rounded bracket of a critical point. Infinite endpoints use limits. + + Decimal elementary functions and directed arithmetic provide the + rounding bounds without a libm accuracy assumption or a new dependency. + This encloses the mathematical function, not errors in an arbitrary + floating-point implementation of its derivative. + """ + if not isinstance(value, Interval) or value.shape != (): + raise TypeError("tanh_double_prime_bounds requires a scalar Interval.") + a, b = float(value.lower), float(value.upper) + if isnan(a) or isnan(b): + raise ValueError("tanh_double_prime_bounds does not accept NaN endpoints.") + + left = _tanh_double_prime_point_bounds(a) + right = _tanh_double_prime_point_bounds(b) + lower, upper = min(left[0], right[0]), max(left[1], right[1]) + if a <= _CRITICAL_HI and b >= _CRITICAL_LO: + lower = -_MAXIMUM_HI + if a <= -_CRITICAL_LO and b >= -_CRITICAL_HI: + upper = _MAXIMUM_HI + return Interval.from_bounds(max(lower, -_MAXIMUM_HI), min(upper, _MAXIMUM_HI)) diff --git a/src/intervalnets/pytorch.py b/src/intervalnets/pytorch.py index 6765024..80ee42b 100644 --- a/src/intervalnets/pytorch.py +++ b/src/intervalnets/pytorch.py @@ -8,6 +8,7 @@ from typing import Any from .interval import Interval +from .activations import tanh_double_prime_bounds from .polynomial_zonotope import PZOneJet, PZTwoJet, PolynomialZonotope from .pz_tanh import ( affine_tanh_double_prime_enclosure, @@ -2178,11 +2179,17 @@ def _interval_second_derivative_bounds_sigmoid(value: Interval) -> Interval: def _interval_second_derivative_bounds_tanh(value: Interval) -> Interval: - tanh_bounds = _apply_monotone_bounds(IntervalTensor((value.lower,), (value.upper,)), tanh) - tanh_interval = Interval(tanh_bounds.lower[0], tanh_bounds.upper[0]) - one = Interval.point(1.0) - two = Interval.point(2.0) - return -(two * tanh_interval * (one - (tanh_interval * tanh_interval))) + bounds = tanh_double_prime_bounds(value) + if bounds.lower == bounds.upper == 0.0: + return bounds + # Retain the PyTorch interval backend's extra float32 outward padding. + lower = _pad_outward(bounds.lower, -inf, include_float32=True) + upper = _pad_outward(bounds.upper, inf, include_float32=True) + if value.lower >= 0.0: + upper = min(upper, 0.0) + if value.upper <= 0.0: + lower = max(lower, 0.0) + return Interval.from_bounds(lower, upper) def _matrix_multiply(left: list[list[Interval]], right: list[list[Interval]]) -> list[list[Interval]]: diff --git a/tests/test_activation_bounds.py b/tests/test_activation_bounds.py new file mode 100644 index 0000000..1331c67 --- /dev/null +++ b/tests/test_activation_bounds.py @@ -0,0 +1,162 @@ +"""High-precision references and regression cases for real tanh'' ranges.""" +from decimal import Decimal, Inexact, localcontext +import math +import random + +import pytest + +from intervalnets import Interval, tanh_double_prime_bounds + + +def _reference(x): + if math.isinf(x): + return Decimal(0) + z = Decimal.from_float(abs(x)) + with localcontext() as ctx: + ctx.prec = 100 + (max(0, -z.adjusted()) if z else 0) + # Independent positive-exponential formula; no tanh saturation. + e = (2*z).exp() + result = -8*e*(e-1)/(e+1)**3 + return result if x >= 0 else -result + + +def _critical_reference(): + with localcontext() as ctx: + ctx.prec = 100 + root = Decimal(3).sqrt() + return (2+root).ln()/2, 4/(3*root) + + +def _range_reference(a, b): + c, maximum = _critical_reference() + with localcontext() as ctx: + ctx.prec = 100 + points = [_reference(a), _reference(b)] + if Decimal.from_float(a) <= c <= Decimal.from_float(b): + points.append(-maximum) + if Decimal.from_float(a) <= -c <= Decimal.from_float(b): + points.append(maximum) + return min(points), max(points) + + +CASES = [ + (-0.1, 0.1), (-1.0, 1.0), (-2.0, 2.0), (0.0, 0.5), + (0.5, 1.0), (1.0, 2.0), (2.0, 3.0), (-3.0, -2.0), + (5.0, 6.0), (20.0, 21.0), (-21.0, -20.0), + (350.0, 351.0), (373.0, 374.0), (400.0, 401.0), + (-math.inf, math.inf), (0.0, math.inf), (-math.inf, 0.0), +] + + +@pytest.mark.parametrize("a,b", CASES) +def test_range_encloses_high_precision_extrema_and_is_tight(a, b): + bounds = tanh_double_prime_bounds(Interval(a, b)) + lo, hi = _range_reference(a, b) + assert Decimal.from_float(bounds.lower) <= lo + assert Decimal.from_float(bounds.upper) >= hi + # Up to two extra binary64 steps permit conservative critical brackets + # and the explicit far-tail underflow enclosure. + assert bounds.lower >= math.nextafter(math.nextafter(float(lo), -math.inf), -math.inf) + assert bounds.upper <= math.nextafter(math.nextafter(float(hi), math.inf), math.inf) + + +@pytest.mark.parametrize("x", [ + 0.0, -0.0, math.nextafter(0.0, math.inf), -math.nextafter(0.0, math.inf), + 1e-300, -1e-300, 1e-20, -1e-20, 0.25, -0.25, 1.0, + 20.0, -20.0, 373.0, -373.0, 400.0, -400.0, +]) +def test_point_intervals_enclose_the_real_value(x): + bounds = tanh_double_prime_bounds(Interval.point(x)) + reference = _reference(x) + assert Decimal.from_float(bounds.lower) <= reference <= Decimal.from_float(bounds.upper) + if x == 0.0: + assert bounds.as_tuple() == (0.0, 0.0) + elif x > 0: + assert bounds.upper <= 0 + else: + assert bounds.lower >= 0 + + +def test_adjacent_floats_around_both_critical_points(): + critical, _ = _critical_reference() + c = float(critical) + for sign in [-1.0, 1.0]: + center = sign*c + points = [ + math.nextafter(center, -math.inf), center, + math.nextafter(center, math.inf), + ] + for i, a in enumerate(points): + for b in points[i:]: + bounds = tanh_double_prime_bounds(Interval(a, b)) + lo, hi = _range_reference(a, b) + assert Decimal.from_float(bounds.lower) <= lo + assert Decimal.from_float(bounds.upper) >= hi + + +def test_random_intervals_contain_samples_and_respect_odd_symmetry(): + rng = random.Random(20260928) + for _ in range(150): + a = rng.uniform(-30, 30) + b = a + 10**rng.uniform(-10, 1) + bounds = tanh_double_prime_bounds(Interval(a, b)) + reflected = tanh_double_prime_bounds(Interval(-b, -a)) + assert reflected.lower == -bounds.upper + assert reflected.upper == -bounds.lower + for i in range(9): + x = a + (b-a)*i/8 + # Guard the sampled float against a rounded endpoint overshoot. + x = min(b, max(a, x)) + assert Decimal.from_float(bounds.lower) <= _reference(x) <= Decimal.from_float(bounds.upper) + + +def test_does_not_depend_on_or_change_decimal_context(): + expected = tanh_double_prime_bounds(Interval(0.5, 1.0)).as_tuple() + with localcontext() as ctx: + ctx.prec = 3 + ctx.traps[Inexact] = True + before = ctx.copy() + assert tanh_double_prime_bounds(Interval(0.5, 1.0)).as_tuple() == expected + assert ctx.prec == before.prec + assert ctx.traps == before.traps + assert ctx.flags == before.flags + + +@pytest.mark.parametrize("a,b", [(math.nan, 1.0), (0.0, math.nan)]) +def test_rejects_nan(a, b): + with pytest.raises(ValueError, match="NaN"): + tanh_double_prime_bounds(Interval(a, b)) + + +@pytest.mark.parametrize("value", [(-1.0, 1.0), Interval([-1.0], [1.0])]) +def test_requires_scalar_interval(value): + with pytest.raises(TypeError, match="scalar Interval"): + tanh_double_prime_bounds(value) + + +def test_large_finite_inputs_do_not_overflow(): + bounds = tanh_double_prime_bounds(Interval(1e300, 1e308)) + assert bounds.as_tuple() == (-math.nextafter(0.0, math.inf), 0.0) + + +@pytest.mark.parametrize("a,b", [(-1.0, 1.0), (0.5, 1.0), (2.0, 3.0)]) +def test_strictly_improves_old_interval_product(a, b): + t = Interval(math.tanh(a), math.tanh(b)) + old = -(Interval.point(2)*t*(Interval.point(1)-t*t)) + new = tanh_double_prime_bounds(Interval(a, b)) + assert old.lower <= new.lower <= new.upper <= old.upper + assert new.upper-new.lower < old.upper-old.lower + + +def test_pytorch_adapter_encloses_scalar_bounds_and_preserves_sign(): + pytest.importorskip("torch") + from intervalnets.pytorch import _interval_second_derivative_bounds_tanh + for a, b in CASES[:12] + [(0.0, 0.0)]: + value = Interval(a, b) + real = tanh_double_prime_bounds(value) + padded = _interval_second_derivative_bounds_tanh(value) + assert padded.lower <= real.lower <= real.upper <= padded.upper + if a >= 0: + assert padded.upper <= 0 + if b <= 0: + assert padded.lower >= 0 diff --git a/tests/test_pytorch.py b/tests/test_pytorch.py index 44c4211..cdb78cf 100644 --- a/tests/test_pytorch.py +++ b/tests/test_pytorch.py @@ -790,6 +790,30 @@ def test_eval_hessian_tanh_network_encloses_corner_second_derivatives() -> None: assert hessian.lower[0][1][1] <= expected_11 <= hessian.upper[0][1][1] +@pytest.mark.parametrize( + "lower,upper,maximum_width", + [(0.5, 1.0, 0.130102), (-1.0, -0.5, 0.130102), (-1.0, 1.0, 1.539602)], +) +def test_tanh_hessian_includes_interior_extrema_with_tight_bounds( + lower, upper, maximum_width +) -> None: + domain = IntervalTensor([lower], [upper]) + hessian = _eval_hessian_bounds(nn.Tanh(), domain) + lo, hi = hessian.lower[0][0][0], hessian.upper[0][0][0] + critical = math.atanh(1.0 / math.sqrt(3.0)) + maximum = 4.0 / (3.0 * math.sqrt(3.0)) + if lower <= critical <= upper: + assert lo <= -maximum <= hi + if lower <= -critical <= upper: + assert lo <= maximum <= hi + assert hi - lo < maximum_width + + +def test_tanh_hessian_at_zero_is_exact_zero() -> None: + hessian = _eval_hessian_bounds(nn.Tanh(), IntervalTensor([0.0], [0.0])) + assert hessian.lower[0][0][0] == hessian.upper[0][0][0] == 0.0 + + def test_sobolev_norm_constant_network_matches_closed_form() -> None: enable_interval_eval() model = nn.Sequential(nn.Linear(1, 1))