r'''Softmax over sparse polynomials, tightened by the simplex constraint.
Softmax outputs satisfy :math:`\sum_i \sigma_i = 1` exactly, but the
propagated abstract value only satisfies it approximately. The residual
polynomial :math:`E = \big(\sum_i \sigma_i\big) - 1` is therefore *known to
be zero*, so any multiple of it may be added to the output without changing
the represented function. :func:`softmax_opt` exploits this: it picks, per
output element, the linear combination of the constraint equations that
minimizes the resulting interval width (least squares on the error rows,
rescaled by the exact L1-optimal step), and adds it. A sound rescaling of
the error symbols (:func:`restrict_errors`, justified by the bounds from
:func:`bounding_error_terms`) first shrinks symbols the constraint pins down.
'''
from __future__ import annotations
from dataclasses import dataclass, field
import torch
from boundlab.ibp.softmax import Softmax2ExpReciprocal
from boundlab.interp import Interpreter, OpHandler
from boundlab.polysp import PolySp
from boundlab.polysp.generator import Generator
[docs]
@dataclass
class SoftmaxConstrained(OpHandler):
'''Run the domain's softmax decomposition, then apply the sum-to-one
constraint optimization (:func:`softmax_opt`) to the result.'''
op: str = "softmax"
inner: OpHandler = field(default_factory=Softmax2ExpReciprocal)
[docs]
def condition(self, x, **params) -> bool:
return isinstance(x, PolySp) and params.get("axis", -1) == -1
[docs]
def handle(self, interp: Interpreter, x: PolySp, **params):
softmax = self.inner.handle(interp, x, **params)
softmax = interp.after_each(softmax, "softmax")
return add_sum_constraint(softmax)
[docs]
def softmax_opt(polysp: PolySp, axis: int = -1) -> PolySp:
r'''Add the optimal multiple of the zero-valued constraint
:math:`\sum_i \sigma_i - 1 = 0` to each element.
Sound because the added expression is identically zero on every concrete
input; effective because its error rows are anti-correlated with the
output's and cancel them.
'''
assert axis == -1, "Softmax optimization only supports axis=-1"
equations = polysp.sum(-1).add(-1.0).reshape(-1)
c, hw = bounding_error_terms(equations.gens)
scale = (c.abs() + hw).clamp(max=1.0)
polysp = PolySp(
polysp.table,
restrict_errors(polysp.gens, scale)
)
equations = PolySp(
equations.table,
restrict_errors(equations.gens, scale)
)
# Skip the constant row: only the error rows contribute to the width, and
# letting lstsq chase the centre picks huge multipliers that inflate them.
lhs = equations.gens.data.flatten(1)[1:]
rhs = polysp.gens.data.flatten(1)[1:]
if lhs.shape[0] == 0:
return polysp
# lstsq minimizes the L2 norm of the rows while the interval width is
# their L1 norm, so rescale its correction by the exact L1-optimal step.
solution = torch.linalg.lstsq(lhs, -rhs).solution
step = l1_optimal_step(lhs @ solution, rhs)
coeffs = (solution * step).T
result = polysp.add(equations.einsum([(0,), (1,0), (1,)], coeffs).reshape(*polysp.shape))
return result
add_sum_constraint = softmax_opt
[docs]
def l1_optimal_step(direction: torch.Tensor, residual: torch.Tensor) -> torch.Tensor:
"""Per-column step ``t`` minimizing ``|residual + t * direction|.sum(0)``.
The minimizer of this piecewise-linear objective is the weighted median of
the ratios ``-residual / direction`` under the weights ``|direction|``.
Zero is always a candidate, so stepping never increases the L1 norm.
"""
weight = direction.abs()
active = weight > 0
ratio = torch.where(
active,
-residual / torch.where(active, direction, torch.ones_like(direction)),
torch.inf,
)
ratio, order = ratio.sort(0)
cumulative = weight.gather(0, order).cumsum(0)
total = cumulative[-1]
median = (cumulative < total / 2).sum(0, keepdim=True)
step = ratio.gather(0, median).squeeze(0)
return torch.where(total > 0, step, torch.zeros_like(step))
[docs]
def restrict_errors(poly: Generator, hw: torch.Tensor) -> Generator:
"""Return updated generator with old_error = hw * new_error range.
Args:
poly: Generator to restrict.
hw: torch.Tensor of shape (poly.error_len,) representing the coefficient of the new error.
"""
hw = hw.to(poly.data.dtype)
one = torch.ones((), dtype=poly.data.dtype)
factors = torch.where(poly.indices >= 0, hw[poly.indices.clamp(min=0).long()], one)
scale = factors.prod(1).reshape((-1,) + (1,) * (poly.data.ndim - 1))
return poly.with_data(poly.data * scale)
[docs]
def bounding_error_terms(equations: Generator) -> tuple[torch.Tensor, torch.Tensor]:
"""Given a set of equations, return the estimated bound [c - hw, c + hw] for each error term.
This is done simply by 1) concretizing highorder terms, 2) Compute intersection of bounds for each equation.
Returns:
c: torch.Tensor of shape (equations.error_len,) representing the constant term.
hw: torch.Tensor of shape (equations.error_len,) representing the coefficient of the new
"""
data = equations.data.reshape(equations.data.shape[0], -1)
degree = equations.nse_order()
constant = data[degree == 0].sum(0)
high = data[degree >= 2].abs().sum(0)
linear = degree == 1
coeff = torch.zeros((equations.error_len, data.shape[1]), dtype=data.dtype)
coeff.index_add_(0, equations.indices[linear].amax(1).long(), data[linear])
# Solve `constant + coeff[j] * e_j + rest = 0` for each e_j, where `rest`
# spans the remaining linear terms plus the concretized high-order terms,
# then intersect over the equations and the trivial range [-1, 1].
active = coeff != 0
safe = torch.where(active, coeff, torch.ones_like(coeff))
slack = coeff.abs().sum(0) - coeff.abs() + high
center = -constant / safe
radius = slack / safe.abs()
ub = torch.where(active, center + radius, torch.inf).amin(1).clamp(-1.0, 1.0)
lb = torch.where(active, center - radius, -torch.inf).amax(1).clamp(-1.0, 1.0)
return (ub + lb) / 2, (ub - lb).clamp(min=0.0) / 2
__all__ = [
"SoftmaxConstrained",
"add_sum_constraint",
"bounding_error_terms",
"l1_optimal_step",
"restrict_errors",
"softmax_opt",
]