r"""The sparse polynomial domain: monomials over error symbols.
Where a zonotope must collapse every product into interval noise, a
:class:`PolySp` keeps the monomials themselves:
.. math::
x \;=\; \sum_r d_r \prod_{k \in I_r} \varepsilon_k ,
\qquad \varepsilon_k \in [-1, 1],
stored sparsely — one coefficient tensor row and one ``-1``-padded index row
per monomial. A product of two order-1 values is then an *exact* order-2
value; only monomials beyond the configured ``order``, or with negligible
coefficients (``eps``), are truncated into sound interval noise. Attention
blocks, whose error budget is dominated by ``Q @ K^T`` and ``attn @ V``, are
the motivating case.
The pair-product kernels live in :mod:`boundlab.ops.sparse_poly` (reference
and Triton backends); this module wraps them in the :class:`~boundlab.Expr`
interface.
"""
from __future__ import annotations
from collections import defaultdict
from collections.abc import Sequence
from dataclasses import dataclass, field
from functools import partial
from typing import Literal, Optional, Union, final, override
import torch
from boundlab import Expr, Zeros, utils
from boundlab.error import Error, Reasons
from boundlab.error.alignment import Alignment, Span, SpanTable
from boundlab.ibp import Bias, Noise
from boundlab.ibp.matmul import MatmulBiased
from boundlab.ibp.mul import MulBiased, MulNoised
from boundlab.ibp.softmax import Softmax2ExpReciprocal
from boundlab.interp import Interpreter, OpHandler, base
from boundlab.ops import sparse_poly
from lasso.lad import admm_dict_learning
from boundlab.polysp import legendre
from boundlab.polysp.legendre import Poly2Mul
from boundlab.zono import linearizers
from boundlab.polysp.generator import Generator
from boundlab.utils import Dim, Polynomial, ShapeDtype, TensorFormat, same_shape
from boundlab.zono import Zono
[docs]
@dataclass(frozen=True)
@final
class PolySp(Expr):
"""A polynomial expression component with sparse representation.
One generator stores mixed-order monomials with ``-1``-padded indices;
all monomials index error symbols through one shared span table.
"""
table: SpanTable[Error]
gens: Generator
[docs]
def __init__(
self, table: SpanTable, gens: Generator | Sequence[Generator]
):
assert isinstance(table, SpanTable)
inputs = [gens] if isinstance(gens, Generator) else list(gens)
if not inputs:
raise ValueError("PolySp requires at least one generator.")
expected_error_len = table.validate()
for gen in inputs:
if not isinstance(gen, Generator):
raise TypeError(f"gens must contain Generators, got {type(gen)}.")
if str(gen.error_len) != str(expected_error_len):
raise ValueError(
f"Generator error length {gen.error_len} does not match "
f"sum of expression spans {expected_error_len}."
)
if not all(
same_shape(gen.shape_dtype, inputs[0].shape_dtype)
for gen in inputs[1:]
):
raise ValueError("All generators must have the same expression shape.")
if len(inputs) == 1:
generator = inputs[0]
else:
order = max(gen.order for gen in inputs)
generator = sparse_poly.sum(*(gen.pad_order(order) for gen in inputs))
object.__setattr__(self, "table", table)
object.__setattr__(self, "gens", generator)
@property
@override
def shape_dtype(self) -> ShapeDtype:
return ShapeDtype(
self.gens.shape_dtype.shape,
self.gens.shape_dtype.dtype,
)
@property
def by_order(self) -> dict[int, Generator]:
return {self.gens.order: self.gens}
@property
def nbytes(self):
"""Estimated memory of all generators' data and indices, in bytes."""
return self.gens.nbytes
[docs]
@staticmethod
def make(table: SpanTable, *gens: Generator) -> "PolySp":
"""Build a ``PolySp``, padding and merging its generators."""
return PolySp(table, gens)
[docs]
@classmethod
@override
def convert_from(cls, expr: Expr) -> Optional["PolySp"]:
"""Lift an ``Error``/``Noise`` symbol or a dense ``Zono`` into a sparse
polynomial.
``Bias`` has no polynomial counterpart, so an ``Intervals`` group keeps
its center outside the ``PolySp`` and only its ``Noise`` part converts.
"""
if isinstance(expr, Zeros):
return PolySp.zeros(expr.shape_dtype)
if isinstance(expr, Bias):
return PolySp(SpanTable(), Generator.from_bias(expr))
if isinstance(expr, Noise):
return PolySp.error(Error(expr.reasons, expr.shape_dtype), expr.noise)
if isinstance(expr, Error):
return PolySp.error(expr, 1.0)
if isinstance(expr, Zono):
return PolySp.from_zono(expr)
[docs]
@staticmethod
def error(
err: Error,
amplitude: Union[torch.Tensor, float] = 1.0,
) -> "PolySp":
"""Diagonal polynomial scaling each component of ``err`` by ``amplitude``."""
amplitude = torch.as_tensor(amplitude)
assert amplitude.ndim == 0 or same_shape(amplitude.shape, err.shape_dtype)
size = err.numel()
if amplitude.ndim == 0:
matrix = amplitude * torch.eye(size, dtype=err.shape_dtype.dtype)
else:
matrix = torch.diag(amplitude.reshape(-1))
data = torch.cat(
(
torch.zeros((1, *err.shape_dtype.shape), dtype=matrix.dtype),
matrix.reshape((size, *err.shape_dtype.shape)),
),
dim=0,
)
indices = torch.cat(
(
torch.full((1, 1), -1, dtype=torch.int32),
torch.arange(size, dtype=torch.int32)[:, None],
),
dim=0,
)
return PolySp(SpanTable(err), Generator(data, indices, size))
[docs]
@staticmethod
def zeros(
shape_dtype: ShapeDtype,
table: SpanTable | None = None,
order: int = 1,
) -> "PolySp":
gen = Generator.from_shape(
shape_dtype, 0 if table is None else table.validate(), order
)
table = table if table is not None else SpanTable()
return PolySp(table, gen)
[docs]
@staticmethod
def from_zono(zono: Zono) -> "PolySp":
"""Reinterpret a dense zonotope as an order-1 sparse polynomial."""
error_len = zono.gen.error_len
data = torch.cat(
(
torch.zeros(
(1, *zono.shape_dtype.shape), dtype=zono.shape_dtype.dtype
),
torch.movedim(zono.gen.tensor, -1, 0),
),
dim=0,
)
indices = torch.cat(
(
torch.full((1, 1), -1, dtype=torch.int32),
torch.arange(error_len, dtype=torch.int32)[:, None],
),
dim=0,
)
return PolySp(zono.table, Generator(data, indices, error_len))
[docs]
def full_abs(self) -> "PolySp":
return PolySp(self.table, self.gens.full_abs())
[docs]
@override
def ub(self) -> torch.Tensor:
return self.gens.ub()
[docs]
@override
def lb(self) -> torch.Tensor:
return self.gens.lb()
[docs]
@override
def lbub(self) -> tuple[torch.Tensor, torch.Tensor]:
return self.gens.lbub()
[docs]
@override
def einsum(
self,
subscripts: list[tuple[int, ...]],
*operands: torch.Tensor,
) -> "PolySp":
return PolySp(self.table, self.gens.einsum(subscripts, *operands))
@staticmethod
def _remapped(
gen: Generator,
alignments: Sequence[Alignment],
side: Literal["span1", "span2"],
error_len: Dim,
) -> Generator:
"""Reindex a generator's error symbols into an aligned table."""
pairs = [
(utils.cast(Span, getattr(alignment, side)), alignment.span_out)
for alignment in alignments
if getattr(alignment, side) is not None
]
def remap(indices: torch.Tensor) -> torch.Tensor:
offset = torch.zeros_like(indices)
for old, new in pairs:
inside = (indices >= old.start) & (indices < old.stop)
offset = torch.where(inside, new.start - old.start, offset)
return indices + offset
return gen.map_indices(remap, error_len)
[docs]
def align_fill_zeros(
self,
other: "PolySp",
) -> tuple["PolySp", "PolySp"]:
if not isinstance(other, PolySp):
raise TypeError(f"other must be a PolySp, got {type(other)}.")
if self.table == other.table:
return self, other
table, alignments = self.table.align_with(other.table)
error_len = table.validate()
return (
PolySp(
table,
self._remapped(self.gens, alignments, "span1", error_len),
),
PolySp(
table,
self._remapped(other.gens, alignments, "span2", error_len),
),
)
[docs]
def expanded_to_table(self, table: SpanTable) -> "PolySp":
if not self.table.subset_of(table):
raise ValueError("Cannot expand to a table that omits existing errors.")
if self.table == table:
return self
output_table, alignments = self.table.align_with(table)
error_len = output_table.validate()
return PolySp(
output_table,
self._remapped(self.gens, alignments, "span1", error_len),
)
[docs]
@override
def add(self, other: Expr | torch.Tensor | float) -> "PolySp":
if isinstance(other, (torch.Tensor, float)):
if isinstance(other, float):
other = torch.full(self.shape, other, dtype=self.shape_dtype.dtype)
return self.add(Bias(other))
other = other.to(PolySp)
if not same_shape(self.shape_dtype, other.shape_dtype):
raise ValueError(
f"ShapeDtype mismatch: {self.shape_dtype} and {other.shape_dtype}."
)
left, right = self.align_fill_zeros(other)
order = max(left.gens.order, right.gens.order)
return PolySp(
left.table,
sparse_poly.sum(
left.gens.pad_order(order), right.gens.pad_order(order)
),
)
[docs]
@override
def reshape(self, *shape: Dim) -> "PolySp":
return PolySp(self.table, self.gens.reshape(*shape))
[docs]
@override
def transpose(self, *perm: int) -> "PolySp":
return PolySp(self.table, self.gens.transpose(*perm))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "PolySp":
return PolySp(self.table, self.gens.broadcast_to(*shape))
[docs]
def sparse_matmul(self, other: "PolySp", eps=1e-3, order: int = 2) -> tuple["PolySp", Noise]:
r"""Matrix product with sound truncation.
Every pair of monomials multiplies exactly —
:math:`(d_r \prod_{k \in I_r} \varepsilon_k)(d_s \prod_{k \in I_s} \varepsilon_k)`
is a monomial of degree :math:`|I_r| + |I_s|` — and the result keeps
those whose degree fits ``order`` and whose magnitude clears ``eps``;
everything discarded is enclosed by the returned ``Noise`` (each
dropped monomial contributes at most its coefficient magnitude, since
every :math:`\varepsilon` factor lies in :math:`[-1, 1]`).
"""
s, o = self.align_fill_zeros(other)
product, noise = sparse_poly.matmul(s.gens, o.gens, eps=eps, order=order)
return PolySp(s.table, product), noise
[docs]
def sparse_applypoly(self, poly: Polynomial, eps:float =1e-3, order: int = 2) -> tuple["PolySp", Noise]:
"""Apply a concrete quadratic ``a x^2 + b x + c`` elementwise via the
sparse pair product, with the same ``eps`` / ``order`` truncation."""
if poly.degree == 2:
c, b, a = poly.coeffs
result, noise = sparse_poly.quadratic(self.gens, torch.tensor(a), torch.tensor(b), eps=eps, order=order)
result.data[0] += c
return PolySp(self.table, result), noise
raise ValueError("Only quadratic polynomials are supported.")
def _error_contribution(self, span: Span) -> torch.Tensor:
"""Halfwidth attributable to one error, split evenly across the
symbols of each monomial."""
contribution = torch.zeros(
self.shape_dtype.shape, dtype=self.shape_dtype.dtype
)
indices = self.gens.indices
inside = (indices >= span.start) & (indices < span.stop)
degree = torch.count_nonzero(indices >= 0, dim=1).clamp(min=1)
fraction = inside.sum(1) / degree
weights = fraction.reshape((-1, *([1] * self.ndim)))
contribution += (self.gens.data.abs() * weights).sum(0)
return contribution
[docs]
def reasons_breakdown(self) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
total = torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
reasons = defaultdict(
lambda: torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
)
for err, span in self.table.items():
contribution = self._error_contribution(span)
total += contribution
for name, weight in err.reasons.items():
reasons[name] += contribution * weight
return total, dict(reasons)
[docs]
@override
def torch_print(
self,
group: Literal["reason", "error_ref"] = "reason",
) -> TensorFormat:
if group == "reason":
total, reasons = self.reasons_breakdown()
total = 2 * total.mean()
groups = {
name: 2 * reason.mean()
for name, reason in reasons.items()
}
else:
groups = {
str(error_ref): 2 * self._error_contribution(span).mean()
for error_ref, span in self.table.items()
}
total = sum(groups.values())
fmt = ", ".join(f"{name}={{{name}:.2g}}" for name in groups)
return TensorFormat(
f"PolySp({{total:.8g}}, mem={{mem}}, {fmt})",
total=total,
mem=utils.format_bytes(self.nbytes),
**groups,
)
[docs]
def compress(
self,
nterms: int,
lam: float = 0.5,
max_iter: int = 20,
tol: float = 1e-3,
max_cols: Optional[int] = None,
) -> tuple["PolySp", Noise]:
r"""Re-express the value over at most ``nterms - 1`` fresh symbols,
with a certified residual, via :func:`lasso.lad.admm_dict_learning`
(LAD-lasso dictionary learning from the ``pytorch-lasso`` package).
.. warning::
Not fully tested: soundness is covered by sampling unit tests
and the defaults were tuned on real barrier matrices, but no
complete end-to-end verified run has exercised compression yet.
Every non-constant monomial :math:`\mu_r = \prod_{k \in I_r}
\varepsilon_k` ranges in :math:`[-1, 1]`, so the value is a linear
map of the monomial vector: :math:`x = d_0 + A \mu` with
:math:`A \in \mathbb{R}^{p \times R}` holding one flattened
coefficient column per monomial. LAD-LASSO factorizes
:math:`A \approx D S` with :math:`k' = nterms - 1` dictionary atoms;
each row of :math:`S`, rescaled into the dictionary, defines a fresh
symbol
.. math::
\varepsilon'_j = \frac{S_{j\cdot}\,\mu}{\|S_{j\cdot}\|_1}
\in [-1, 1],
sound when treated as a new independent symbol (the true
:math:`\varepsilon'` set is a subset of the box). The factorization
error is enclosed coefficient-wise — element :math:`i` of the
returned ``Noise`` is :math:`\sum_r |A - DS|_{ir}` — which is
exactly the LAD fit term the solver minimizes, while the L1 penalty
keeps :math:`\|S_{j\cdot}\|_1`, and with it the new symbols'
amplitudes, from inflating.
The fresh symbols carry the value's blended
:class:`~boundlab.error.Reasons`; the residual is tagged
``"compress"``. Correlation with other live values is severed —
the compressed value no longer references the old symbols (see the
cross-value correlation caveat in ``docs/guide/polysp-plan.md``).
Only the ``max_cols`` (default ``16 * nterms``) largest monomials by
coefficient l1 mass are handed to the solver; the tail joins the
residual noise directly. Coefficient mass is heavily concentrated
in practice (on the SST barrier matrices the tail beyond
``16 * 1024`` columns carries ~0.2% of the mass), and this keeps the
solve tractable — LAD-LASSO on all 350k+ monomials of a real
attention barrier would take hours.
The solver runs with ``init="topk"`` (dictionary seeded with the
heaviest monomials, codes warm-started with the matching diagonal,
so iteration 0 is exactly the top-k truncation solution and the
iterations move tail mass from the residual into the atoms), the
ISTA inner lasso solver and the vectorized least-squares +
projection dictionary update (``dict_update="lstsq"``); ``tol`` is
accepted for interface stability but the ADMM loop runs a fixed
``max_iter`` steps. This matters because
residual ``Noise`` propagates through downstream linear layers as
``|W| h`` (no sign cancellation) while generator rows propagate as
``W d``: on the real barrier-0 matrix, cold-started solves leave
~90% of the mass in the residual and inflate downstream width
~7.7x, whereas the warm-started defaults (``lam=0.5``,
``max_iter=20``) reach ~1.02x mean / ~1.4x worst-element with 0.6%
of the mass in ``Noise`` in ~12 s — an order of magnitude below
plain truncation's 3.8%.
Returns:
The compressed ``PolySp`` (constant row plus at most
``nterms - 1`` order-1 rows over one fresh ``Error``) and the
certified residual ``Noise``.
"""
shape = self.gens.data.shape[1:]
dtype = self.gens.data.dtype
if self.gens.nterms <= nterms:
return self, Noise(torch.zeros(shape, dtype=dtype), "compress")
if nterms < 2:
raise ValueError(f"nterms must be at least 2, got {nterms}.")
max_cols = 16 * nterms if max_cols is None else max_cols
gens = self.gens
tail = torch.zeros(shape, dtype=dtype)
if gens.nterms - 1 > max_cols:
l1 = gens.data[1:].reshape(gens.nterms - 1, -1).abs().sum(dim=1)
keep = torch.zeros(gens.nterms, dtype=torch.bool)
keep[0] = True
keep[1:][l1.topk(max_cols).indices] = True
gens, tail = gens.filter_nse(keep)
del tol
coeffs = gens.data[1:]
X = coeffs.reshape(coeffs.shape[0], -1).contiguous() # monomials as samples
k = nterms - 1
D, Z, _ = admm_dict_learning(
X, k, alpha=lam, steps=max_iter, init="topk", progbar=False,
algorithm="ista", maxiter=10, dict_update="lstsq",
return_codes=True,
)
amplitude = Z.abs().sum(dim=0)
rows = (D * amplitude).T.reshape(k, *shape)
residual = tail + (X - Z @ D.T).abs().sum(dim=0).reshape(shape)
total, breakdown = self.reasons_breakdown()
mass = total.sum()
if mass > 0:
reasons = Reasons(
**{name: r.sum() / mass for name, r in breakdown.items()}
)
else:
reasons = Reasons("compress")
fresh = Error(reasons, ShapeDtype((k,), dtype), name="compress")
data = torch.cat((self.gens.data[:1], rows))
indices = torch.cat(
(
torch.full((1, 1), -1, dtype=torch.int32),
torch.arange(k, dtype=torch.int32)[:, None],
)
)
nonzero = rows.reshape(k, -1).abs().amax(dim=1) > 0
gen, _ = Generator(data, indices, k).filter_nse(
torch.cat((torch.ones(1, dtype=torch.bool), nonzero))
)
return PolySp(SpanTable(fresh), gen), Noise(residual, "compress")
[docs]
def keep_ploysp(x, name: str=""):
"""``after_each`` normalizer: lift stray ``Bias`` / ``Noise`` components
into ``PolySp`` so every value stays inside the polynomial world and new
residuals become first-class symbols."""
del name
# Multiple-results primitives (e.g. custom_jvp_call) hand over a list.
if isinstance(x, (list, tuple)):
return type(x)(keep_ploysp(v) for v in x)
if isinstance(x, Expr) and (Bias in x.classset() or Noise in x.classset()):
bias, other = x.split(Bias)
noise, other = other.split(Noise)
return bias.to(PolySp) + noise.to(PolySp) + other
assert isinstance(x, Expr) and x.classset().issubset({PolySp}) \
or isinstance(x, torch.Tensor)
return x
[docs]
@dataclass
class Poly(OpHandler):
"""Apply a concrete polynomial (given as a coefficient list) to a
``PolySp`` via :meth:`PolySp.sparse_applypoly`."""
op: str = "poly"
eps: float = 1e-3
order: int = 2
[docs]
def condition(self, coeffs: list[torch.Tensor], x: Expr, **kwargs):
del kwargs
return isinstance(x, Expr) and x.classset().issubset({Bias, PolySp}) and all(isinstance(c, torch.Tensor) for c in coeffs)
[docs]
def handle(self, interp, coeffs: list[torch.Tensor], x: PolySp, **kwargs) -> Expr:
del kwargs
if Bias in x.classset():
bias, other = x.split(Bias)
noise, other = other.split(Noise)
return bias.to(PolySp) + noise.to(PolySp) + self.handle(interp, coeffs, other)
poly = Polynomial(coeffs)
result, noise = x.sparse_applypoly(poly, eps=self.eps, order=self.order)
return result + noise
[docs]
@dataclass
class Matmul(OpHandler):
"""PolySp × PolySp matrix product via :meth:`PolySp.sparse_matmul`."""
op: str = "matmul"
eps: float = 1e-3
order: int = 2
[docs]
def condition(self, x: Expr, y: Expr, **kwargs):
del kwargs
return isinstance(x, PolySp) and isinstance(y, PolySp)
[docs]
def handle(self, interp, x: PolySp, y: PolySp, **kwargs) -> Expr:
del kwargs
result, noise = x.sparse_matmul(y, eps=self.eps, order=self.order)
return result + noise
@dataclass
class PolySpCompression(OpHandler):
op: str = "marked_identity"
overrides: list[type[OpHandler]] = field(
default_factory=lambda: [base.MarkedIdentity]
)
nterm_limit: int = 2**16
error_len: int = 1024
tol: float = 1e-3
def condition(self, x, **kwargs):
return isinstance(x, PolySp) and 'name' in kwargs and kwargs['name'] == 'layer_barrier'
def handle(self, interp: Interpreter, x: PolySp, **kwargs):
del kwargs
if torch.compiler.is_exporting():
# The solver loop is not traceable, and the term count is a
# data-dependent size under torch.export; exported pipelines
# (benchmark_compilers) run without compression.
return x
if x.gens.nterms > self.nterm_limit:
x, noise = x.compress(nterms=self.error_len, tol=self.tol)
return x + noise
return x
interpret = Interpreter(
base.interpret,
Matmul(eps=1e-4, order=4),
linearizers.MaxWithConst2Relu(),
linearizers.Relu(),
linearizers.Tanh(),
linearizers.Exp(),
linearizers.Reciprocal(),
Softmax2ExpReciprocal(),
PolySpCompression(),
after_each=keep_ploysp
)
"""The sparse polynomial interpreter.
Extends :data:`boundlab.interp.base.interpret` with the exact-to-truncation
:class:`Matmul` pair product (``eps=1e-4``, ``order=4``), the zonotope
linearizers (their residuals are lifted straight back into polynomial
symbols by ``keep_ploysp``), and the softmax decomposition. Swap in
:class:`~boundlab.polysp.softmax.SoftmaxConstrained` for the
simplex-tightened softmax.
"""
# Imported last: softmax pulls PolySp back out of this module.
from boundlab.polysp import softmax # noqa: E402
__all__ = [
"generator",
"legendre",
"softmax",
"Generator",
"Matmul",
"Poly",
"PolySp",
"interpret",
"keep_ploysp",
]