Source code for boundlab.ibp.softmax
r"""Softmax via the DeepT shift-invariant decomposition.
Dividing numerator and denominator by :math:`e^{\nu_i}` turns softmax into
.. math::
\sigma_i(\nu) = \frac{e^{\nu_i}}{\sum_j e^{\nu_j}}
= \frac{1}{\sum_j e^{\nu_j - \nu_i}},
a composition of pairwise differences, ``exp``, a reduce-sum, and one
``reciprocal`` — no product of two abstract values, and the differences keep
the exponent range centered so ``exp`` stays well-conditioned.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Optional
from boundlab import Expr, Intersects
from boundlab.ibp.elementwise import Exp
from boundlab.interp import Interpreter, OpHandler
[docs]
@dataclass
class Softmax2ExpReciprocal(OpHandler):
r"""Rewrite softmax as :math:`1 / \sum_j e^{\nu_j - \nu_i}` over the last axis.
The ``exp`` and ``reciprocal`` steps go through the enclosing domain's
handlers (overridable via ``interp_exp`` / ``interp_reciprocal``). The
denominator is also evaluated in plain interval arithmetic and joined
with the symbolic version through :class:`~boundlab.Intersects`, so the
reciprocal sees the tighter of the two enclosures per element — the
interval one wins exactly where symbolic cancellation has degenerated.
"""
op: str = "softmax"
interp_exp: Optional[OpHandler] = field(default=None)
interp_reciprocal: Optional[OpHandler] = field(default=None)
[docs]
def condition(self, x, **params) -> bool:
return isinstance(x, Expr) and params.get("axis", -1) == -1
[docs]
def handle(self, interp: Interpreter, x: Expr, **params):
del params
def exp(interp: Interpreter,x: Expr, **params):
if self.interp_exp is not None:
return self.interp_exp.handle(interp, x, **params)
else:
return interp.exp(x, **params)
def reciprocal(interp: Interpreter, x: Expr, **params):
if self.interp_reciprocal is not None:
return self.interp_reciprocal.handle(interp, x, **params)
else:
return interp.reciprocal(x, **params)
size = x.shape[-1]
pairwise_shape = (*x.shape, size)
x_i = x.reshape(*x.shape, 1).broadcast_to(*pairwise_shape)
x_j = x.reshape(*x.shape[:-1], 1, size).broadcast_to(*pairwise_shape)
denominator = exp(interp, x_j - x_i).sum(axis=-1)
denominator_itvl = Exp()(*(x_j - x_i).lbub()).sum(axis=-1) # type: ignore
result = reciprocal(interp, Intersects(denominator, denominator_itvl))
intersect, result = result.split(Intersects)
return intersect[0] + result
__all__ = ["Softmax2ExpReciprocal"]