r"""The coefficient tensor behind a :class:`~boundlab.zono.Zono`."""
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from functools import partial
from typing import final, override
import torch
from boundlab import Expr
from boundlab.utils import Dim, ShapeDtype, einsum_formatter, same_shape
from boundlab.error.alignment import Span
[docs]
@dataclass(frozen=True)
@final
class Generator(Expr):
r"""Zonotope coefficients ``tensor[*shape, error_len]``.
Concretization is the dual-norm bound: over
:math:`\varepsilon \in [-1,1]^n` the affine form
:math:`G\varepsilon` attains exactly
.. math::
\sup_{\|\varepsilon\|_\infty \le 1} (G\varepsilon)_p
\;=\; \sum_k |G_{p k}| ,
so ``ub()`` is ``|tensor|.sum(-1)`` — tight, not just sound. Linear
primitives act on the leading axes and leave the error axis untouched.
"""
tensor: torch.Tensor
[docs]
def __init__(self, tensor: torch.Tensor):
if tensor is None:
raise ValueError("Generator tensor cannot be None.")
if tensor.ndim == 0:
raise ValueError("Generator tensor must have a trailing error dimension.")
object.__setattr__(self, "tensor", tensor)
@property
@override
def shape_dtype(self) -> ShapeDtype:
return ShapeDtype(self.tensor.shape[:-1], self.tensor.dtype)
[docs]
@staticmethod
def from_shape(
shape_dtype: ShapeDtype,
error_len: Dim,
) -> "Generator":
return Generator(torch.zeros((*shape_dtype.shape, error_len), dtype=shape_dtype.dtype))
@property
def error_len(self) -> Dim:
return self.tensor.shape[-1]
[docs]
@override
def ub(self) -> torch.Tensor:
if str(self.error_len) == "0":
return torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype)
return self.tensor.abs().sum(-1)
[docs]
@override
def lb(self) -> torch.Tensor:
return -self.ub()
[docs]
@override
def lbub(self) -> tuple[torch.Tensor, torch.Tensor]:
ub = self.ub()
return -ub, ub
[docs]
@override
def chw(self) -> tuple[torch.Tensor, torch.Tensor]:
return torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype), self.ub()
[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.")
error_label = (
max((label for part in subscripts for label in part), default=-1) + 1
)
inputs, output = subscripts[:-1], subscripts[-1]
equation = [
inputs[0] + (error_label,),
*inputs[1:],
output + (error_label,),
]
return Generator(
torch.einsum(
einsum_formatter(equation),
self.tensor,
*operands,
)
)
[docs]
@override
def add(self, other: Expr) -> "Generator":
if not isinstance(other, Generator):
return NotImplemented
if not same_shape(self.tensor.shape, other.tensor.shape):
raise ValueError(
f"Generator shapes {self.tensor.shape} and {other.tensor.shape} "
"are not compatible for addition."
)
return Generator(self.tensor + other.tensor)
[docs]
@override
def reshape(self, *shape: Dim) -> "Generator":
return Generator(self.tensor.reshape((*shape, self.error_len)))
[docs]
@override
def transpose(self, *perm: int) -> "Generator":
return Generator(self.tensor.permute(*perm, self.ndim))
[docs]
@override
def broadcast_to(self, *shape: Dim) -> "Generator":
return Generator(torch.broadcast_to(self.tensor, (*shape, self.error_len)))
[docs]
@staticmethod
def concat(*generators: "Generator") -> "Generator":
"""Concatenate along the error axis — the sum of zonotopes over
*disjoint* symbol groups."""
if not generators:
raise ValueError("At least one generator must be provided for concatenation.")
shape = generators[0].shape_dtype
if any(not same_shape(generator.shape_dtype, shape) for generator in generators):
raise ValueError("All generators must have the same shape for concatenation.")
return Generator(
torch.cat(
tuple(generator.tensor for generator in generators),
dim=-1,
)
)
[docs]
def full_abs(self) -> "Generator":
return Generator(self.tensor.abs())
[docs]
def error_span(self, span: Span) -> "Generator":
"""The coefficient slice belonging to one error symbol's span."""
return Generator(self.tensor[..., span.start : span.stop])
__all__ = ["Generator"]