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
25 changes: 25 additions & 0 deletions docs/API.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`
Expand Down
64 changes: 64 additions & 0 deletions docs/tanh_second_derivative_bounds.md
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.
2 changes: 2 additions & 0 deletions src/intervalnets/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Interval arithmetic utilities for neural network evaluation."""

from .interval import Interval
from .activations import tanh_double_prime_bounds

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Keep the scalar export importable without PyTorch

When the package is installed with only its base dependencies (PyTorch is optional in pyproject.toml), from intervalnets import tanh_double_prime_bounds still crashes: the later .pytorch import catches the missing-PyTorch ImportError, sets nn = None, and then evaluates class IntervalAdd(nn.Module), raising AttributeError. 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 👍 / 👎.

from .polynomial_zonotope import (
PZOneJet,
PZTwoJet,
Expand Down Expand Up @@ -81,6 +82,7 @@

__all__ = [
"Interval",
"tanh_double_prime_bounds",
"PolynomialZonotope",
"PZOneJet",
"PZTwoJet",
Expand Down
102 changes: 102 additions & 0 deletions src/intervalnets/activations.py
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))
17 changes: 12 additions & 5 deletions src/intervalnets/pytorch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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]]:
Expand Down
Loading
Loading