Repository navigation
Tighten tanh Hessian intervals using stationary extrema #191
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
MoritzMaibaum
merged 1 commit into
better-int
from
codex/tanh-second-derivative-better-int
Sep 28, 2026
Merged
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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)) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
When the package is installed with only its base dependencies (PyTorch is optional in
pyproject.toml),from intervalnets import tanh_double_prime_boundsstill crashes: the later.pytorchimport catches the missing-PyTorchImportError, setsnn = None, and then evaluatesclass IntervalAdd(nn.Module), raisingAttributeError. Consequently, the newly exported scalar API cannot be used in the explicitly documented no-PyTorch environment; the package initializer must avoid loading that module or the optional module must remain import-safe when PyTorch is absent.Useful? React with 👍 / 👎.