"""Differential softmax, and softmax under score-based pruning.
Both handlers use the DeepT rewrite
σᵢ(ν) = 1 / Σⱼ exp(νⱼ − νᵢ)
so softmax reduces to ``exp → reduce-sum → reciprocal``, each of which already
has a differential lineariser. No product is needed, and the shift by ``νᵢ``
keeps the exponent range bounded.
Two soundness guards sit around that pipeline:
*Denominator floor.* The reciprocal's domain is strictly positive, and a
propagated lower bound that dips to ``≤ 0`` makes it emit vacuous envelopes.
But the denominator has a *provable* floor independent of the propagated one:
``Dᵢ = Σⱼ exp(νⱼ − νᵢ) ≥ max(1, exp(maxⱼ ν_lbⱼ − ν_ubᵢ))``, since the ``j = i``
term is ``exp(0) = 1`` and every term is positive. Raising a propagated lower
bound to a proven floor is sound.
*Range intersection.* Both branches output probabilities in ``[0, 1]``, so
their difference lies in ``[−1, 1]``. Intersecting the propagated envelope with
those known ranges is sound and keeps any residual reciprocal blow-up from
cascading into later layers.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
from boundlab import Expr
from boundlab.ibp import Bias, Noise
from boundlab.interp import Interpreter, OpHandler
from boundlab.diff.expr import DiffExpr2, DiffExpr3
from .heaviside import apply_mask, const_scores, heaviside_mask
def _box(lower: torch.Tensor, upper: torch.Tensor, reason: str) -> Expr:
center = (lower + upper) / 2
return Bias(center) + Noise((upper - lower) / 2, reason)
def clamp_expr(value: Expr, lo: float, hi: float, reason: str) -> Expr:
"""Intersect ``value``'s envelope with the known interval ``[lo, hi]``.
Elements already inside are kept verbatim, preserving the correlations that
downstream cancellation relies on; elements that escape are replaced by the
box hull of the intersection. Sound either way: the true value lies in both
the propagated interval and ``[lo, hi]``.
"""
lower, upper = value.lbub()
finite = torch.isfinite(lower) & torch.isfinite(upper)
inside = finite & (lower >= lo - 1e-6) & (upper <= hi + 1e-6)
if bool(inside.all()):
return value
new_lower = torch.clamp(lower, min=lo)
new_upper = torch.clamp(upper, max=hi)
# Non-finite envelopes, and empty intersections (fp slop, or an unsound
# upstream bound), fall back to the full mathematical range: it is the only
# information left to trust.
degenerate = (~finite) | (new_lower > new_upper)
new_lower = torch.where(degenerate, torch.full_like(new_lower, lo), new_lower)
new_upper = torch.where(degenerate, torch.full_like(new_upper, hi), new_upper)
box = _box(new_lower, new_upper, reason)
if not bool(finite.all()):
# Blending against non-finite coefficients would evaluate ``0 * inf``;
# drop the correlations and keep the pure box.
return box
mask = inside.to(new_lower.dtype)
return value * mask + box * (1.0 - mask)
def floor_expr(value: Expr, floor: torch.Tensor, reason: str) -> Expr:
"""Raise ``value``'s lower bound to the proven per-element ``floor``."""
lower, upper = value.lbub()
finite = torch.isfinite(lower) & torch.isfinite(upper)
needed = finite & (lower < floor)
if not bool(needed.any()):
return value
new_lower = torch.maximum(lower, floor)
new_upper = torch.maximum(upper, new_lower) # trust the floor on contradiction
zeros = torch.zeros_like(upper)
box = _box(
torch.where(needed, new_lower, zeros),
torch.where(needed, new_upper, zeros),
reason,
)
# ``needed`` implies ``finite``, so every masked-out coefficient is finite.
return value * (~needed).to(upper.dtype) + box
def _pairwise_shift(value: Expr) -> tuple[Expr, int]:
"""Return ``ν_j − ν_i`` with ``i`` on axis ``-2`` and ``j`` on axis ``-1``."""
size = value.shape[-1]
pairwise = (*value.shape, size)
value_i = value.reshape(*value.shape, 1).broadcast_to(*pairwise)
value_j = value.reshape(*value.shape[:-1], 1, size).broadcast_to(*pairwise)
return value_j - value_i, size
def _denominator_floor(value: Expr, mask: torch.Tensor | None = None) -> torch.Tensor:
"""Proven lower bound on ``Σⱼ [maskⱼ] exp(νⱼ − νᵢ)``.
Every term is positive, so keeping only the largest one is already a valid
floor. Without a mask the ``j = i`` term contributes ``exp(0) = 1``. The
exponent is capped at 80 to keep the floor finite — lowering an exponent
only weakens the floor, so it stays valid.
"""
lower, upper = value.lbub()
kept = lower if mask is None else torch.where(
mask.bool(), lower, torch.full_like(lower, float("-inf"))
)
floor = torch.exp(
(kept.amax(dim=-1, keepdim=True) - upper).clamp(max=80.0)
)
return floor if mask is not None else floor.clamp(min=1.0)
def _apply_floor(denominator: DiffExpr3, data: DiffExpr3, mask=None) -> DiffExpr3:
return DiffExpr3(
floor_expr(
denominator.x, _denominator_floor(data.x), "softmax_denominator"
),
floor_expr(
denominator.y, _denominator_floor(data.y, mask), "softmax_denominator"
),
denominator.diff,
)
def _clamp_result(result: Any, reason: str) -> Any:
if not isinstance(result, DiffExpr3):
return result
return DiffExpr3(
clamp_expr(result.x, 0.0, 1.0, reason),
clamp_expr(result.y, 0.0, 1.0, reason),
clamp_expr(result.diff, -1.0, 1.0, reason),
)
[docs]
@dataclass
class DiffSoftmax(OpHandler):
"""Differential softmax over the last axis.
Plain expressions are delegated to ``fallback`` (the domain's standard
softmax handler), so this handler owns the operator outright.
"""
op: str = "softmax"
fallback: OpHandler | None = None
[docs]
def condition(self, x: Any, **params: Any) -> bool:
if isinstance(x, (DiffExpr2, DiffExpr3)):
return params.get("axis", -1) == -1
return self.fallback is not None and self.fallback.condition(x, **params)
[docs]
def handle(self, interp: Interpreter, x: Any, **params: Any) -> Expr:
if not isinstance(x, (DiffExpr2, DiffExpr3)):
assert self.fallback is not None
return self.fallback.handle(interp, x, **params)
del params
data = x.to(DiffExpr3)
shifted, _ = _pairwise_shift(data)
denominator = interp.exp(shifted).sum(axis=-1)
denominator = _apply_floor(denominator.to(DiffExpr3), data)
return _clamp_result(interp.reciprocal(denominator), "diff_softmax")
[docs]
class DiffSoftmaxPruning(OpHandler):
"""Differential handler for ``boundlab::SoftmaxPruning`` over the last axis.
Network 1 keeps the full softmax; network 2 sees the mask applied to both
the numerator and every denominator term, i.e.
``h(sᵢ)·exp(dᵢ) / Σⱼ h(sⱼ)·exp(dⱼ)``.
"""
op = "softmax_pruning"
[docs]
def handle(
self,
interp: Interpreter,
scores: Any,
data: Any,
*,
dim: int = -1,
**params: Any,
) -> Expr:
del params
if isinstance(data, torch.Tensor):
data = Bias(data) # a graph constant: shared between the networks
triple = data.to(DiffExpr3)
ndim = len(triple.shape_dtype.shape)
if dim not in (-1, ndim - 1):
raise NotImplementedError(
f"Differential softmax pruning supports the last axis only, got {dim}."
)
mask = heaviside_mask(const_scores(scores))
shifted, size = _pairwise_shift(triple)
# Scores select *keys* (axis j), broadcast over the query axis i.
key_mask = torch.broadcast_to(
mask.reshape(*mask.shape[:-1], 1, size), tuple(shifted.shape_dtype.shape)
)
terms = apply_mask(key_mask, interp.exp(shifted))
denominator = _apply_floor(
terms.sum(axis=-1).to(DiffExpr3), triple, mask
)
result = apply_mask(mask, interp.reciprocal(denominator))
return _clamp_result(result, "diff_softmax_pruning")
__all__ = [
"DiffSoftmax",
"DiffSoftmaxPruning",
"clamp_expr",
"floor_expr",
]