diff --git a/docs/API.md b/docs/API.md index ecb6e39..0470277 100644 --- a/docs/API.md +++ b/docs/API.md @@ -42,6 +42,34 @@ 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 first derivative of tanh + +`tanh_prime_bounds(value: Interval) -> Interval` is available directly from +`intervalnets`, without PyTorch. It encloses the mathematical derivative on a +scalar interval using its evenness and monotonic decrease with distance from +zero: + +```python +from intervalnets import Interval, tanh_prime_bounds + +bounds = tanh_prime_bounds(Interval(20.0, 21.0)) +# Approximately [2.299808905717424e-18, 1.699341702116636e-17]. +``` + +Endpoint values use the stable expression `4*q/(1+q)**2`, where +`q = exp(-2*abs(x))`, with outward-rounded Decimal arithmetic. This avoids +the cancellation and saturation error in `1-tanh(x)**2`. The exact maximum +is 1 on intervals containing zero, and the point interval at zero returns +`[1,1]`. All results lie in `[0,1]`; finite saturated/underflow cases keep a +positive upper bound, and infinite endpoints use limits. NaN endpoints and +nonscalar inputs are rejected. See [the rounding argument](tanh_prime_bounds.md). + +The interval Jacobian/Hessian backend uses this helper with its extra float32 +padding, clipped to `[0,1]`. The deep hybrid derivative interval factors also +use it and round outward after conversion to their tensor dtype. The scalar +API encloses the real derivative rather than arbitrary floating-point +autograd errors. + ## Scalar second derivative of tanh `tanh_double_prime_bounds(value: Interval) -> Interval` is available directly diff --git a/docs/tanh_prime_bounds.md b/docs/tanh_prime_bounds.md new file mode 100644 index 0000000..28317e0 --- /dev/null +++ b/docs/tanh_prime_bounds.md @@ -0,0 +1,60 @@ +# Outward-rounded interval bounds for tanh' + +The function `h(x) = tanh'(x) = sech(x)^2` is even and decreases as `abs(x)` +increases. Its only finite maximum is `h(0)=1`. Thus the real range on `[a,b]` +is obtained at the largest and smallest distance to zero: + +``` +far = max(abs(a), abs(b)) +near = 0 if a <= 0 <= b else min(abs(a), abs(b)) +range = [h(far), h(near)]. +``` + +The previous endpoint calculation `1-tanh(x)**2` could round to zero even when +the real derivative was representable and positive. For example, on `[20,21]` +the real range is approximately `[2.2998089057e-18, 1.6993417021e-17]`, whereas +the padded old result was approximately `[-1.4013e-45, 1.4013e-45]`. + +## Rounding argument + +Use `q=exp(-2*abs(x))` and `h(x)=4*q/(1+q)^2`. The rational function is +increasing on `[0,1]`, since its derivative there is `4*(1-q)/(1+q)^3`. + +1. Convert the binary64 endpoint to an exact Decimal. Bound the multiplication + by -2 with directed Decimal arithmetic. The resulting exponent interval + brackets the real exponent, including very small inputs. +2. Decimal exponential is correctly rounded to nearest even. Take its lower + neighbor at the lower exponent and upper neighbor at the upper exponent, + and intersect with `[0,1]`. This gives `q_lo <= q <= q_hi` without a + system-libm accuracy assumption. +3. By monotonicity, it suffices to bound `h(q_lo)` below and `h(q_hi)` above. + For the lower bound, round `4*q_lo` down and `(1+q_lo)^2` up before + dividing downward. For the upper bound, use the opposite rounding + directions at `q_hi`. All factors are positive. +4. Convert the Decimal bounds outward to binary64, comparing the exact + Decimal image of the converted float with the original Decimal bound. + Clip to the exact global range `[0,1]`. At zero use `[1,1]` directly. +5. For finite `abs(x)>=400`, use `0 tuple[float, float, float]: _MIN_SUBNORMAL = nextafter(0.0, inf) +def _tanh_prime_point_bounds(x: float) -> tuple[float, float]: + if x == 0.0: + return 1.0, 1.0 + if isinf(x): + return 0.0, 0.0 + if abs(x) >= 400.0: + # 0 < tanh'(x) <= 4*exp(-800) < 2**-1074 at finite endpoints. + return 0.0, _MIN_SUBNORMAL + + down, up = _contexts() + magnitude = Decimal.from_float(abs(x)) + 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))) + + # h(q)=4*q/(1+q)**2 is increasing for q in [0,1]. Evaluate h(q_lo) + # downward and h(q_hi) upward, rather than subtracting tanh(x)**2 + # from 1. This also avoids cancellation near zero. + base_lo_upper = up.add(1, q_lo) + base_hi_lower = down.add(1, q_hi) + lower = down.divide( + down.multiply(4, q_lo), up.multiply(base_lo_upper, base_lo_upper) + ) + upper = up.divide( + up.multiply(4, q_hi), down.multiply(base_hi_lower, base_hi_lower) + ) + return max(0.0, _float_lower(lower)), min(1.0, _float_upper(upper)) + + +def tanh_prime_bounds(value: Interval) -> Interval: + """Enclose the mathematical ``tanh'`` on a scalar interval. + + ``tanh'`` is even and decreases with distance from zero. Evaluate the + largest distance for the minimum and the smallest for the maximum; + a domain containing zero has maximum exactly 1. Endpoint values use + outward-rounded Decimal arithmetic for 4*q/(1+q)**2 with q=exp(-2|x|). + Bounds stay within [0,1], including subnormal/underflow cases. Infinite + endpoints use limits. No assumption about libm accuracy is needed. + """ + if not isinstance(value, Interval) or value.shape != (): + raise TypeError("tanh_prime_bounds requires a scalar Interval.") + a, b = float(value.lower), float(value.upper) + if isnan(a) or isnan(b): + raise ValueError("tanh_prime_bounds does not accept NaN endpoints.") + + far = max(abs(a), abs(b)) + near = 0.0 if a <= 0.0 <= b else min(abs(a), abs(b)) + far_bounds = _tanh_prime_point_bounds(far) + upper = far_bounds[1] if near == far else _tanh_prime_point_bounds(near)[1] + return Interval.from_bounds(far_bounds[0], upper) + + 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. diff --git a/src/intervalnets/deep_hybrid.py b/src/intervalnets/deep_hybrid.py index ab72d15..dee32cd 100644 --- a/src/intervalnets/deep_hybrid.py +++ b/src/intervalnets/deep_hybrid.py @@ -26,6 +26,7 @@ from typing import Any, Literal from .interval import Interval +from .activations import tanh_prime_bounds from .polynomial_zonotope import PolynomialZonotope from .pz_integration import PZIntegrationCell from .shallow_hybrid import ( @@ -414,21 +415,29 @@ def deep_scalar_hybrid_onejet_reverse( derivative_linear_coefficients = ( derivative_linears + 2.0 * quadratics * z_center ).unsqueeze(1) * z_coefficients - endpoint_lower = 1.0 - torch.tanh(lower).square() - endpoint_upper = 1.0 - torch.tanh(upper).square() - derivative_lower = torch.minimum(endpoint_lower, endpoint_upper) - crosses_zero = (lower <= 0.0) & (upper >= 0.0) - derivative_upper = torch.where( - crosses_zero, - torch.ones_like(lower), - torch.maximum(endpoint_lower, endpoint_upper), + derivative_bounds = [ + tanh_prime_bounds(Interval(float(lo), float(hi))) + for lo, hi in zip( + lower.detach().cpu().tolist(), upper.detach().cpu().tolist() + ) + ] + derivative_lower = torch.tensor( + [bounds.lower for bounds in derivative_bounds], + dtype=lower.dtype, + device=lower.device, + ) + derivative_upper = torch.tensor( + [bounds.upper for bounds in derivative_bounds], + dtype=upper.dtype, + device=upper.device, ) + # Round outward again after conversion to the factor's tensor dtype. derivative_lower = torch.nextafter( derivative_lower, torch.full_like(derivative_lower, -torch.inf) - ) + ).clamp(min=0.0, max=1.0) derivative_upper = torch.nextafter( derivative_upper, torch.full_like(derivative_upper, torch.inf) - ) + ).clamp(min=0.0, max=1.0) raw_factors.append( { "preactivation_center": z_center, diff --git a/src/intervalnets/pytorch.py b/src/intervalnets/pytorch.py index 80ee42b..7045d87 100644 --- a/src/intervalnets/pytorch.py +++ b/src/intervalnets/pytorch.py @@ -8,7 +8,7 @@ from typing import Any from .interval import Interval -from .activations import tanh_double_prime_bounds +from .activations import tanh_double_prime_bounds, tanh_prime_bounds from .polynomial_zonotope import PZOneJet, PZTwoJet, PolynomialZonotope from .pz_tanh import ( affine_tanh_double_prime_enclosure, @@ -280,9 +280,10 @@ def _pz_onejet_trace_record( try: import torch from torch import nn -except ImportError: # pragma: no cover - environment dependent - torch = None - nn = None +except ImportError as exc: # pragma: no cover - environment dependent + # The package's optional-import guard must receive ImportError before + # class definitions attempt to inherit from nn.Module. + raise ImportError("PyTorch is required for intervalnets.pytorch.") from exc class IntervalTensor(Interval): @@ -2148,20 +2149,12 @@ def _interval_derivative_bounds_sigmoid(value: Interval) -> Interval: def _interval_derivative_bounds_tanh(value: Interval) -> Interval: - lower = float(value.lower) - upper = float(value.upper) - tanh_lower = tanh(lower) - tanh_upper = tanh(upper) - - derivative_lower_endpoint = 1.0 - tanh_lower * tanh_lower - derivative_upper_endpoint = 1.0 - tanh_upper * tanh_upper - - maximum = max(derivative_lower_endpoint, derivative_upper_endpoint) - if lower <= 0.0 <= upper: - maximum = 1.0 - minimum = min(derivative_lower_endpoint, derivative_upper_endpoint) - lower_out = _pad_outward(minimum, -inf, include_float32=True) - upper_out = _pad_outward(maximum, inf, include_float32=True) + bounds = tanh_prime_bounds(value) + if bounds.lower == bounds.upper and bounds.lower in (0.0, 1.0): + return bounds + # Retain float32 padding, intersected with the exact global range [0,1]. + lower_out = max(0.0, _pad_outward(bounds.lower, -inf, include_float32=True)) + upper_out = min(1.0, _pad_outward(bounds.upper, inf, include_float32=True)) return Interval.from_bounds(lower_out, upper_out) diff --git a/tests/test_deep_hybrid.py b/tests/test_deep_hybrid.py index 058f374..0cc18f7 100644 --- a/tests/test_deep_hybrid.py +++ b/tests/test_deep_hybrid.py @@ -111,6 +111,41 @@ def test_deep_factored_hybrid_retains_all_factors_and_is_sound() -> None: assert torch.all(gradients <= enclosure_upper) +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +@pytest.mark.parametrize("lower,upper", [(-1.0, 1.0), (20.0, 21.0), (60.0, 61.0)]) +def test_deep_derivative_intervals_enclose_real_values_including_saturation( + dtype, lower, upper +) -> None: + from decimal import Decimal, localcontext + + model = nn.Sequential( + nn.Linear(1, 1), nn.Tanh(), nn.Linear(1, 1), nn.Tanh(), nn.Linear(1, 1) + ).to(dtype=dtype) + with torch.no_grad(): + for layer in (model[0], model[2], model[4]): + layer.weight.fill_(1.0) + layer.bias.zero_() + domain = PolynomialZonotope.from_box( + torch.tensor([lower], dtype=dtype), torch.tensor([upper], dtype=dtype) + ) + result = scalar_hybrid_onejet_reverse(model, domain) + for factor in result.factors: + assert torch.all(factor.derivative_lower >= 0) + assert torch.all(factor.derivative_upper <= 1) + a, b = factor.preactivation_lower.item(), factor.preactivation_upper.item() + lo, hi = factor.derivative_lower.item(), factor.derivative_upper.item() + assert factor.derivative_lower.dtype == factor.center.dtype + if a <= 0 <= b: + assert hi == 1.0 + with localcontext() as ctx: + ctx.prec = 100 + for x in (a, (a+b)/2, b): + z = Decimal.from_float(x) + cosh = (z.exp() + (-z).exp()) / 2 + real = 1 / cosh**2 + assert Decimal.from_float(lo) <= real <= Decimal.from_float(hi) + + def test_deep_factored_integral_encloses_sampled_w12() -> None: model = _deep_model() box = IntervalTensor.from_bounds([-0.35, -0.35], [0.35, 0.35]) diff --git a/tests/test_pytorch.py b/tests/test_pytorch.py index cdb78cf..cb626f9 100644 --- a/tests/test_pytorch.py +++ b/tests/test_pytorch.py @@ -1080,6 +1080,26 @@ def test_sigmoid_jacobian_encloses_autograd_corner_gradients() -> None: assert jacobian.lower[row][col] <= exact <= jacobian.upper[row][col] +def test_tanh_jacobian_at_zero_is_exact_identity() -> None: + jacobian = _eval_jacobian_bounds(nn.Tanh(), IntervalTensor([0.0], [0.0])) + assert jacobian.lower[0][0] == jacobian.upper[0][0] == 1.0 + + +def test_tanh_jacobian_encloses_real_derivative_in_saturated_tail() -> None: + from decimal import Decimal, localcontext + + jacobian = _eval_jacobian_bounds(nn.Tanh(), IntervalTensor([20.0], [21.0])) + lo, hi = jacobian.lower[0][0], jacobian.upper[0][0] + assert 0.0 < lo <= hi < 2e-17 + with localcontext() as ctx: + ctx.prec = 100 + for x in [20, 20.5, 21]: + z = Decimal(str(x)) + cosh = (z.exp() + (-z).exp()) / 2 + real = 1 / cosh**2 + assert Decimal.from_float(lo) <= real <= Decimal.from_float(hi) + + def test_tanh_jacobian_encloses_autograd_corner_gradients() -> None: enable_interval_eval() tanh = nn.Tanh() diff --git a/tests/test_tanh_prime_bounds.py b/tests/test_tanh_prime_bounds.py new file mode 100644 index 0000000..c6ef614 --- /dev/null +++ b/tests/test_tanh_prime_bounds.py @@ -0,0 +1,155 @@ +"""References for tanh' via high-precision cosh, independent of the q formula.""" +from decimal import Decimal, Inexact, localcontext +import math +import random + +import pytest + +from intervalnets import Interval, tanh_prime_bounds + + +def _reference(x): + if math.isinf(x): + return Decimal(0) + z = Decimal.from_float(abs(x)) + with localcontext() as ctx: + # Resolve the O(x**2) departure from 1 even at subnormal inputs. + ctx.prec = 100 + (2*max(0, -z.adjusted()) if z else 0) + cosh = (z.exp() + (-z).exp()) / 2 + return 1 / (cosh*cosh) + + +def _reference_range(a, b): + values = [_reference(a), _reference(b)] + if a <= 0 <= b: + values.append(Decimal(1)) + return min(values), max(values) + + +CASES = [ + (0.0, 0.0), (-1.0, 1.0), (-2.0, 0.5), (-0.5, 2.0), + (2.0, 3.0), (-3.0, -2.0), (5.0, 6.0), (10.0, 11.0), + (20.0, 21.0), (-21.0, -20.0), (100.0, 101.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), + (2.0, math.inf), (-math.inf, -2.0), +] + + +@pytest.mark.parametrize("a,b", CASES) +def test_prime_range_encloses_reference_extrema_and_is_tight(a, b): + result = tanh_prime_bounds(Interval(a, b)) + lower, upper = _reference_range(a, b) + assert 0.0 <= result.lower <= result.upper <= 1.0 + assert Decimal.from_float(result.lower) <= lower + assert Decimal.from_float(result.upper) >= upper + assert result.lower >= math.nextafter(math.nextafter(float(lower), -math.inf), -math.inf) + assert result.upper <= math.nextafter(math.nextafter(float(upper), math.inf), math.inf) + if a <= 0 <= b: + assert result.upper == 1.0 + + +@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, 1e-8, -1e-8, 0.25, + 20.0, -20.0, 373.0, -373.0, 400.0, -400.0, math.inf, -math.inf, +]) +def test_prime_point_values_are_outward_rounded(x): + result = tanh_prime_bounds(Interval.point(x)) + value = _reference(x) + assert Decimal.from_float(result.lower) <= value <= Decimal.from_float(result.upper) + if x == 0: + assert result.as_tuple() == (1.0, 1.0) + elif math.isinf(x): + assert result.as_tuple() == (0.0, 0.0) + else: + assert result.lower < 1.0 + assert result.upper > 0.0 + + +def test_prime_saturated_tail_has_positive_lower_bound(): + # math.tanh(20) is 1, so the old 1-tanh(x)**2 expression returned 0. + result = tanh_prime_bounds(Interval(20.0, 21.0)) + assert result.lower > 2e-18 + assert result.upper < 2e-17 + assert math.tanh(20.0) == 1.0 + + +def test_prime_bounds_are_even_and_contain_random_samples(): + rng = random.Random(20261001) + for _ in range(150): + a = rng.uniform(-60, 60) + b = a + 10**rng.uniform(-10, 1) + result = tanh_prime_bounds(Interval(a, b)) + assert tanh_prime_bounds(Interval(-b, -a)).as_tuple() == result.as_tuple() + for i in range(9): + x = min(b, max(a, a+(b-a)*i/8)) + value = _reference(x) + assert Decimal.from_float(result.lower) <= value <= Decimal.from_float(result.upper) + + +def test_prime_uses_independent_decimal_context(): + expected = tanh_prime_bounds(Interval(20.0, 21.0)).as_tuple() + with localcontext() as ctx: + ctx.prec = 2 + ctx.traps[Inexact] = True + before = ctx.copy() + assert tanh_prime_bounds(Interval(20.0, 21.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_prime_rejects_nan(a, b): + with pytest.raises(ValueError, match="NaN"): + tanh_prime_bounds(Interval(a, b)) + + +@pytest.mark.parametrize("value", [(-1.0, 1.0), Interval([-1.0], [1.0])]) +def test_prime_requires_a_scalar_interval(value): + with pytest.raises(TypeError, match="scalar Interval"): + tanh_prime_bounds(value) + + +def test_prime_large_finite_inputs_keep_nonzero_upper_bound(): + result = tanh_prime_bounds(Interval(1e300, 1e308)) + assert result.as_tuple() == (0.0, math.nextafter(0.0, math.inf)) + + +def test_prime_pytorch_adapter_preserves_range_and_saturated_values(): + pytest.importorskip("torch") + from intervalnets.pytorch import _interval_derivative_bounds_tanh + for a, b in CASES: + value = Interval(a, b) + scalar = tanh_prime_bounds(value) + padded = _interval_derivative_bounds_tanh(value) + assert 0 <= padded.lower <= scalar.lower <= scalar.upper <= padded.upper <= 1 + if a <= 0 <= b: + assert padded.upper == 1.0 + if a == b == 0: + assert padded.as_tuple() == (1.0, 1.0) + assert _interval_derivative_bounds_tanh(Interval(20.0, 21.0)).lower > 0 + + +def test_scalar_activation_api_imports_without_optional_torch(): + import subprocess + import sys + from pathlib import Path + + source_dir = str(Path(__file__).resolve().parents[1] / "src") + script = f""" +import sys +from importlib.abc import MetaPathFinder +sys.path.insert(0, {source_dir!r}) +class NoTorch(MetaPathFinder): + def find_spec(self, fullname, path, target=None): + if fullname == 'torch' or fullname.startswith('torch.'): + raise ModuleNotFoundError('PyTorch disabled for this test', name='torch') +sys.meta_path.insert(0, NoTorch()) +from intervalnets import Interval, tanh_prime_bounds, tanh_double_prime_bounds +assert tanh_prime_bounds(Interval(20,21)).lower > 0 +assert tanh_prime_bounds(Interval.point(0)).as_tuple() == (1.0,1.0) +assert tanh_double_prime_bounds(Interval.point(0)).as_tuple() == (0.0,0.0) +""" + subprocess.run([sys.executable, "-c", script], check=True, capture_output=True, text=True)