Source code for boundlab.zono.generator

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