Source code for boundlab.polysp.softmax

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