Source code for boundlab.zonoq.generator

r"""Quadratic coefficient storage for :class:`~boundlab.zonoq.quad.Quad`."""

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
from boundlab.zono.generator import Generator


[docs] @dataclass(frozen=True) @final class QGenerator(Expr): r"""Coefficients of :math:`\sum_{ij} T_{\cdot ij}\, \varepsilon_i \varepsilon_j`. Concretization splits diagonal from off-diagonal: each :math:`\varepsilon_i^2 \in [0, 1]` contributes one-sidedly, while a cross term :math:`\varepsilon_i \varepsilon_j \in [-1, 1]` is bounded by its magnitude, halved-and-summed over both orderings so a symmetric matrix is not double-counted: .. math:: |x - \tfrac{1}{2}\textstyle\sum_i |T_{ii}|| \;\le\; \tfrac{1}{2}\textstyle\sum_i |T_{ii}| + \Big( \tfrac{1}{2}\big(\textstyle\sum_j |T_{\cdot j}| + \sum_i |T_{i \cdot}|\big) - |T_{ii}| \Big) \cdot 1 . ``concat`` of disjoint symbol groups is block-diagonal — cross-group quadratic coefficients are zero by construction. """ tensor: torch.Tensor """``torch.Tensor[*shape, error_len, error_len]`` representing the quadratic generator tensor."""
[docs] def __init__(self, tensor: torch.Tensor): if tensor is None: raise ValueError("QGenerator tensor cannot be None.") if tensor.ndim < 2: raise ValueError( "QGenerator tensor must have two trailing error dimensions." ) if str(tensor.shape[-2]) != str(tensor.shape[-1]): raise ValueError( "QGenerator error dimensions must be square, got " f"{tensor.shape[-2:]}." ) object.__setattr__(self, "tensor", tensor)
@property @override def shape_dtype(self) -> ShapeDtype: return ShapeDtype( self.tensor.shape[:-2], self.tensor.dtype )
[docs] @staticmethod def from_shape( shape_dtype: ShapeDtype, error_len: Dim, ) -> "QGenerator": return QGenerator( torch.zeros((*shape_dtype.shape, error_len, error_len), dtype=shape_dtype.dtype) )
@property def error_len(self) -> Dim: return self.tensor.shape[-1]
[docs] @override def ub(self) -> torch.Tensor: return self.lbub()[1]
[docs] @override def lb(self) -> torch.Tensor: return self.lbub()[0]
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: diagonal = torch.diagonal(self.tensor, dim1=-2, dim2=-1) loose = 0.5 * ( self.tensor.abs().sum(-1) + self.tensor.abs().sum(-2) ) - diagonal.abs() # 𝜀^2 ∊ [0, 1], 𝜀1^2 𝜀2^2 ∊ [-1, 1] lb = Generator(loose).lb() ub = Generator(diagonal).ub() - lb return lb, ub
[docs] @override def chw(self) -> tuple[torch.Tensor, torch.Tensor]: diagonal = torch.diagonal(self.tensor, dim1=-2, dim2=-1) loose = 0.5 * ( self.tensor.abs().sum(-1) + self.tensor.abs().sum(-2) ) - diagonal.abs() # 𝜀^2 ∊ [0, 1], 𝜀1^2 𝜀2^2 ∊ [-1, 1] center = Generator(diagonal).ub() / 2 halfwidth = Generator(loose).ub() + center return center, halfwidth
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "QGenerator": if len(subscripts) != len(operands) + 2: raise ValueError("einsum requires one subscript per input and one output.") first_error_label = ( max((label for part in subscripts for label in part), default=-1) + 1 ) second_error_label = first_error_label + 1 inputs, output = subscripts[:-1], subscripts[-1] equation = [ inputs[0] + (first_error_label, second_error_label), *inputs[1:], output + (first_error_label, second_error_label), ] return QGenerator( torch.einsum( einsum_formatter(equation), self.tensor, *operands, ) )
[docs] @override def add(self, other: Expr) -> "QGenerator": if not isinstance(other, QGenerator): return NotImplemented if not same_shape(self.tensor.shape, other.tensor.shape): raise ValueError( f"QGenerator shapes {self.tensor.shape} and {other.tensor.shape} " "are not compatible for addition." ) return QGenerator(self.tensor + other.tensor)
[docs] @override def reshape(self, *shape: Dim) -> "QGenerator": return QGenerator( self.tensor.reshape((*shape, self.error_len, self.error_len)) )
[docs] @override def transpose(self, *perm: int) -> "QGenerator": error_axis = self.ndim return QGenerator( self.tensor.permute(*perm, error_axis, error_axis + 1) )
[docs] @override def broadcast_to(self, *shape: Dim) -> "QGenerator": return QGenerator( torch.broadcast_to(self.tensor, (*shape, self.error_len, self.error_len)) )
[docs] @staticmethod def concat(*generators: "QGenerator") -> "QGenerator": 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.") # Each input describes a disjoint group of error symbols. Combining # them therefore forms a block-diagonal matrix in the two error axes; # all cross-group quadratic coefficients are zero. rows = [] for row_index, row_generator in enumerate(generators): blocks = [] for column_index, column_generator in enumerate(generators): if row_index == column_index: block = row_generator.tensor else: block = torch.zeros( ( *shape.shape, row_generator.error_len, column_generator.error_len, ), dtype=torch.result_type( row_generator.tensor, column_generator.tensor ), ) blocks.append(block) rows.append(torch.cat(blocks, dim=-1)) return QGenerator( torch.cat( rows, dim=-2, ) )
[docs] def full_abs(self) -> "QGenerator": return QGenerator(self.tensor.abs())
[docs] def error_span(self, span: Span) -> "QGenerator": return QGenerator( self.tensor[ ..., span.start : span.stop, span.start : span.stop, ] )
__all__ = ["QGenerator"]