Source code for boundlab.polysp.generator


r"""Sparse monomial storage behind a :class:`~boundlab.polysp.PolySp`."""

import math
from collections.abc import Callable
from dataclasses import dataclass
from functools import partial
from typing import Optional, override
import torch

from boundlab import Expr
from boundlab.ibp import Bias
from boundlab.ops import sparse_poly
from boundlab.ops.sparse_poly.structure import partition
from boundlab.utils import Dim, ShapeDtype, einsum_formatter


[docs] @dataclass(frozen=True) class Generator(Expr): r"""A sparse polynomial in ``error_len`` symbols. Row ``r`` of ``data`` (shape ``(nse, *shape)``) is the coefficient of the monomial :math:`\prod_k \varepsilon_{I_{rk}}` given by row ``r`` of ``indices`` (shape ``(nse, order)``, ``-1`` entries are padding). Row 0 is the distinguished constant monomial (all ``-1``), so .. math:: ub = d_0 + \sum_{r \ge 1} |d_r|, \qquad lb = d_0 - \sum_{r \ge 1} |d_r| : sound because every non-constant monomial ranges within :math:`[-1, 1]`; tight only when the monomials are independent. Linear primitives act on the element axes of ``data`` and never touch the monomial structure. """ data: torch.Tensor indices: torch.Tensor error_len: Dim def __post_init__(self): if self.indices.ndim != 2 or self.data.shape[0] != self.indices.shape[0]: raise ValueError("data and indices must have matching row counts") if not torch.compiler.is_compiling() and self.data.shape[0] == 0: raise ValueError("a generator must contain its constant row") if not torch.compiler.is_compiling(): assert bool(torch.all(self.indices[0] == -1)) @property @override def shape_dtype(self) -> ShapeDtype: return ShapeDtype( self.data.shape[1:], self.data.dtype ) @property def order(self) -> int: return self.indices.shape[-1] @property def nbytes(self): """Estimated memory of the stored data and indices, in bytes.""" return ( math.prod(self.data.shape) * self.data.dtype.itemsize + math.prod(self.indices.shape) * self.indices.dtype.itemsize )
[docs] @staticmethod def from_shape( shape_dtype: ShapeDtype, error_len: Dim, order: int = 1, ) -> "Generator": """Zero generator containing only its distinguished constant row.""" data = torch.zeros((1, *shape_dtype.shape), dtype=shape_dtype.dtype) indices = torch.full((1, order), -1, dtype=torch.int32) return Generator( data, indices, error_len )
[docs] @staticmethod def from_bias( bias: Bias ) -> "Generator": """Convert a Bias into a Generator with the same constant row.""" if not isinstance(bias, Expr): raise TypeError(f"expected Expr, got {type(bias)}") data = bias.arr.unsqueeze(0) indices = torch.full((1, 0), -1, dtype=torch.int32) return Generator( data, indices, 0 )
[docs] def with_data(self, data: torch.Tensor) -> "Generator": """Rebuild with new monomial coefficients; keeps the error indices.""" return Generator( data, self.indices, self.error_len )
[docs] @override def ub(self) -> torch.Tensor: return self.data[0] + self.data[1:].abs().sum(0)
[docs] @override def lb(self) -> torch.Tensor: return self.data[0] - self.data[1:].abs().sum(0)
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: c, hw = self.chw() return c - hw, c + hw
[docs] @override def chw(self) -> tuple[torch.Tensor, torch.Tensor]: c = self.data[0] hw = self.data[1:].abs().sum(0) return c, hw
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "Generator": if len(subscripts) != len(operands) + 2: raise ValueError("einsum requires one subscript per input and one output.") nse_label = ( max((label for part in subscripts for label in part), default=-1) + 1 ) inputs, output = subscripts[:-1], subscripts[-1] equation = [ (nse_label,) + inputs[0], *inputs[1:], (nse_label,) + output, ] return self.with_data( torch.einsum( einsum_formatter(equation), self.data, *operands, ) )
[docs] @override def add(self, other: Expr) -> "Generator": if not isinstance(other, Generator): return NotImplemented if other.indices is self.indices: return self.with_data(self.data + other.data) return sparse_poly.sum(self, other)
[docs] @override def reshape(self, *shape: Dim) -> "Generator": return self.with_data( self.data.reshape((self.data.shape[0], *shape)) )
[docs] @override def transpose(self, *perm: int) -> "Generator": return self.with_data( self.data.permute(0, *(axis + 1 for axis in perm)) )
[docs] @override def broadcast_to(self, *shape: Dim) -> "Generator": return self.with_data( torch.broadcast_to(self.data, (self.data.shape[0], *shape)) )
[docs] def map_indices( self, func: Callable[[torch.Tensor], torch.Tensor], error_len: Optional[int] = None, ) -> "Generator": """Rewrite the error indices, optionally into a new error length.""" error_len = self.error_len if error_len is None else error_len return Generator( self.data, func(self.indices), error_len )
[docs] def full_abs(self) -> "Generator": return self.with_data(self.data.abs())
[docs] def filter_nse(self, mask: torch.Tensor) -> tuple["Generator", torch.Tensor]: """Keep the masked monomial rows (the constant row always survives); return the kept generator and the summed magnitude of the dropped rows — a sound interval enclosure of what was removed.""" mask = torch.cat((torch.ones((1,), dtype=torch.bool), mask[1:])) kept, discarded = partition(mask) return Generator( self.data[kept], self.indices[kept], self.error_len ), self.data[discarded].abs().sum(0)
[docs] def nse_order(self) -> torch.Tensor: """Per-row monomial degree: the count of non-padding indices.""" return torch.count_nonzero(self.indices >= 0, dim=1)
[docs] def pad_order(self, order: int) -> "Generator": """Left-pad monomial indices to a larger maximum order.""" if order < self.order: raise ValueError(f"cannot reduce generator order {self.order} to {order}") if order == self.order: return self padding = torch.full( (self.indices.shape[0], order - self.order), -1, dtype=self.indices.dtype, ) return Generator( self.data, torch.cat((padding, self.indices), dim=1), self.error_len, )
@property def nterms(self) -> int: return self.data.shape[0]
__all__ = [ "Generator", ]