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