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"]