Source code for boundlab.polysp

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