Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions docs/API.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
60 changes: 60 additions & 0 deletions docs/tanh_prime_bounds.md
Original file line number Diff line number Diff line change
@@ -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<h(x)<=4*exp(-800)<2**-1074` and return
`[0,2**-1074]` at that endpoint. This prevents accidental return of a
zero upper bound under underflow. At infinite endpoints the limit is zero.

The contexts are independent of the caller's Decimal context and use the
existing activation helpers. There is no additional dependency. The public
result encloses the exact real range with outward rounding.

The PyTorch interval adapter retains its float32 outward padding, clipped to
`[0,1]`. Exact zero and one point results need no padding. Deep hybrid factors
round outward after casting the scalar bounds to their tensor dtype, also
clipping to `[0,1]`. These intervals concern the mathematical derivative;
they do not bound errors from every possible floating-point/autograd formula.
Affine and quadratic PZ approximations are separate representations.

## Verification

Tests use an independent high-precision `1/cosh(x)^2` reference and include
intervals crossing zero, point intervals, signed zero, subnormal inputs and
outputs, saturated tails, unbounded intervals, NaN and shape rejection,
even symmetry, and random intervals. Jacobian and deep-factor tests check
the actual integration points, including float32 tensor conversion and
underflow. Tests supplement the rounding argument; sampled values alone
are not a proof of an enclosure.
3 changes: 2 additions & 1 deletion src/intervalnets/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
"""Interval arithmetic utilities for neural network evaluation."""

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,
Expand Down Expand Up @@ -83,6 +83,7 @@
__all__ = [
"Interval",
"tanh_double_prime_bounds",
"tanh_prime_bounds",
"PolynomialZonotope",
"PZOneJet",
"PZTwoJet",
Expand Down
53 changes: 53 additions & 0 deletions src/intervalnets/activations.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,59 @@ def _critical_bounds() -> 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.
Expand Down
29 changes: 19 additions & 10 deletions src/intervalnets/deep_hybrid.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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,
Expand Down
29 changes: 11 additions & 18 deletions src/intervalnets/pytorch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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)


Expand Down
35 changes: 35 additions & 0 deletions tests/test_deep_hybrid.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand Down
20 changes: 20 additions & 0 deletions tests/test_pytorch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading
Loading