Source code for boundlab.zonoq.quad

r"""The pure-quadratic component :math:`\sum_{ij} Q_{ij} \varepsilon_i \varepsilon_j`."""

from __future__ import annotations

from collections.abc import Sequence
from dataclasses import dataclass
from functools import partial
from typing import Literal, Optional, final, override

import torch

from boundlab import Error, Expr, Zeros
from boundlab.ibp import Noise
from boundlab.utils import Dim, ShapeDtype, TensorFormat, same_shape
from boundlab.error.alignment import Alignment, SpanTable, Span
from boundlab.zono import Zono
from boundlab.zono.generator import Generator
from boundlab.zonoq.generator import QGenerator


[docs] @dataclass(frozen=True) @final class Quad(Expr): """Quadratic zonotope ``sum_ij gen[..., i, j] * eps_i * eps_j``.""" table: SpanTable gen: QGenerator
[docs] def __init__(self, table: SpanTable, gen: QGenerator): assert isinstance(table, SpanTable) if not isinstance(gen, QGenerator): raise TypeError(f"gen must be a QGenerator, got {type(gen)}.") expected_error_len = table.validate() if str(gen.error_len) != str(expected_error_len): raise ValueError( f"QGenerator error length {gen.error_len} does not match " f"sum of expression spans {expected_error_len}." ) object.__setattr__(self, "table", table) object.__setattr__(self, "gen", gen)
@property @override def shape_dtype(self) -> ShapeDtype: return self.gen.shape_dtype
[docs] @classmethod @override def convert_from(cls, expr: Expr) -> Optional["Quad"]: """Build elementwise squared terms from an ``Error`` symbol or ``Noise``. ``Bias`` has no quadratic counterpart, so an ``Intervals`` group keeps its center outside the ``Quad`` and only its ``Noise`` part converts here. """ if isinstance(expr, Zeros): return Quad.zeros(expr.shape_dtype)
[docs] @staticmethod def single( err: Expr, generator: QGenerator | torch.Tensor, ) -> "Quad": if not isinstance(generator, QGenerator): generator = QGenerator(torch.as_tensor(generator)) if str(generator.error_len) != str(err.numel()): raise ValueError( f"Generator error length {generator.error_len} does not match " f"Expr size {err.numel()}." ) return Quad( SpanTable(err), generator, )
[docs] @staticmethod def error( err: Error, amplitude: torch.Tensor | float = 1.0, ) -> "Quad": """Diagonal quadratic ``amplitude_k * eps_k**2`` for each component.""" amplitude = torch.as_tensor(amplitude) assert amplitude.ndim == 0 or same_shape(amplitude.shape, err.shape_dtype) size = err.numel() if amplitude.ndim == 0: amplitude = amplitude.to(err.shape_dtype.dtype).expand(size) else: amplitude = amplitude.reshape(-1) indices = torch.arange(size) tensor = torch.zeros((size, size, size), dtype=err.shape_dtype.dtype) tensor[indices, indices, indices] = amplitude return Quad( SpanTable(err), QGenerator(tensor.reshape((*err.shape_dtype.shape, size, size))), )
[docs] @staticmethod def zeros( shape_dtype: ShapeDtype, table: SpanTable | None = None, ) -> "Quad": gen = QGenerator.from_shape(shape_dtype, 0 if table is None else table.validate()) table = table if table is not None else SpanTable() return Quad(table, gen)
[docs] def full_abs(self) -> "Quad": return Quad(self.table, self.gen.full_abs())
[docs] @override def ub(self) -> torch.Tensor: return self.gen.ub()
[docs] @override def lb(self) -> torch.Tensor: return self.gen.lb()
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: return self.gen.lbub()
[docs] @override def chw(self) -> tuple[torch.Tensor, torch.Tensor]: return self.gen.chw()
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "Quad": return Quad( self.table, self.gen.einsum(subscripts, *operands), )
def __getitem__(self, item: Expr) -> QGenerator: return self.gen.error_span(self.table[item]) @staticmethod def _aligned_generator( generator: QGenerator, alignments: Sequence[Alignment], side: Literal["span1", "span2"], ) -> QGenerator: """Embed a generator into an aligned table without losing cross terms.""" rows = [] for row_alignment in alignments: row_span = getattr(row_alignment, side) blocks = [] for column_alignment in alignments: column_span = getattr(column_alignment, side) if row_span is None or column_span is None: block = torch.zeros( ( *generator.shape_dtype.shape, row_alignment.span_out.size, column_alignment.span_out.size, ), dtype=generator.tensor.dtype, ) else: block = generator.tensor[ ..., row_span.start : row_span.stop, column_span.start : column_span.stop, ] blocks.append(block) rows.append(torch.cat(blocks, dim=-1)) if not rows: return QGenerator.from_shape(generator.shape_dtype, 0) return QGenerator(torch.cat(rows, dim=-2))
[docs] def align_fill_zeros( self, other: "Quad", ) -> tuple["Quad", "Quad"]: if not isinstance(other, Quad): raise TypeError(f"other must be a QZono, got {type(other)}.") if self.table == other.table: return self, other table, alignments = self.table.align_with(other.table) return ( Quad(table, self._aligned_generator(self.gen, alignments, "span1")), Quad(table, self._aligned_generator(other.gen, alignments, "span2")), )
[docs] def expanded_to_table(self, table: SpanTable) -> "Quad": if not self.table.subset_of(table): raise ValueError("Cannot expand to a table that omits existing errors.") if self.table == table: return self output_table, alignments = self.table.align_with(table) return Quad( output_table, self._aligned_generator(self.gen, alignments, "span1"), )
[docs] @override def add(self, other: Expr) -> "Quad": if not isinstance(other, Quad): return NotImplemented if not same_shape(self.shape_dtype, other.shape_dtype): raise ValueError(f"ShapeDtype mismatch: {self.shape_dtype} and {other.shape_dtype}.") left, right = self.align_fill_zeros(other) return Quad(left.table, left.gen.add(right.gen))
[docs] @override def reshape(self, *shape: Dim) -> "Quad": return Quad(self.table, self.gen.reshape(*shape))
[docs] @override def transpose(self, *perm: int) -> "Quad": return Quad(self.table, self.gen.transpose(*perm))
[docs] @override def broadcast_to(self, *shape: Dim) -> "Quad": return Quad(self.table, self.gen.broadcast_to(*shape))
[docs] def reasons_breakdown(self) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: tensor = self.gen.tensor.abs().sum(-1) / 2 + self.gen.tensor.abs().sum(-2) / 2 return Zono(self.table, Generator(tensor)).reasons_breakdown()
[docs] def torch_print( self, group: Literal["reason", "error_ref"] = "reason", ) -> TensorFormat: if group == "reason": total, reasons = self.reasons_breakdown() total = 2 * total.mean() groups = { name: 2 * value.mean() for name, value in reasons.items() } else: groups = { str(error_ref): 2 * self.gen.error_span(span).ub().mean() for error_ref, span in self.table.items() } total = sum(groups.values()) fmt = ", ".join(f"{name}={{{name}:.2g}}" for name in groups) return TensorFormat(f"Quad({{total:.8g}}, {fmt})", total=total, **groups)
__all__ = ["Quad"]