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