Source code for boundlab.zonoq

r"""The quadratic zonotope domain: correlated linear plus quadratic terms.

A :class:`ZonoQ` pairs a linear zonotope with a quadratic form over the *same*
error symbols,

.. math::

   x \;=\; \sum_i L_i \varepsilon_i \;+\; \sum_{ij} Q_{ij}\, \varepsilon_i \varepsilon_j ,

so one matrix product of two zonotopes can be represented *exactly*
(:math:`L^{(1)} \otimes L^{(2)}` lands in :math:`Q`) instead of being
concretized on the spot.  Concretization of :math:`Q` still exploits
:math:`\varepsilon_i^2 \in [0, 1]` on the diagonal, and postponing it lets
later linear layers reshape the quadratic mass before it is ever collapsed.
"""

from __future__ import annotations

from collections import defaultdict
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 Expr, ExprGroup, Zeros
from boundlab import ibp
from boundlab import utils
from boundlab.error import Error
from boundlab.ibp.matmul import MatmulBiased
from boundlab.utils import Dim, ShapeDtype, TensorFormat, same_shape
from boundlab.ibp import Bias, Noise
from boundlab.interp import Interpreter, OpHandler, base
from boundlab.zono import Zono
from boundlab.error.alignment import SpanTable
from boundlab.zono.linearizers import LinearBounds
from boundlab.ibp.softmax import Softmax2ExpReciprocal
from boundlab.zono import linearizers
from boundlab.zonoq.generator import QGenerator
from boundlab.zonoq.quad import Quad

[docs] @dataclass(frozen=True) @final class ZonoQ(Expr): """Correlated linear and quadratic zonotope terms. Represents ``sum_i L_i eps_i + sum_ij Q_ij eps_i eps_j``. ``L`` and ``Q`` share one error table so both terms refer to the same symbols. """ L: Zono Q: Quad
[docs] def __init__(self, L: Zono, Q: Quad): if not isinstance(L, Zono): raise TypeError(f"L must be a Zono, got {type(L)}.") if not isinstance(Q, Quad): raise TypeError(f"Q must be a Quad, got {type(Q)}.") if not same_shape(L.shape_dtype, Q.shape_dtype): raise ValueError( f"L and Q must have the same shape, got " f"{L.shape_dtype} and {Q.shape_dtype}." ) if L.table != Q.table: raise ValueError("L and Q must use the same span table.") object.__setattr__(self, "L", L) object.__setattr__(self, "Q", Q)
@property @override def shape_dtype(self) -> ShapeDtype: return ShapeDtype( self.L.shape_dtype.shape, torch.promote_types(self.L.shape_dtype.dtype, self.Q.shape_dtype.dtype), ) @property def table(self) -> SpanTable: return self.L.table
[docs] @staticmethod def make( L: Zono | None = None, Q: Quad | None = None, copy: "ZonoQ | None" = None, ) -> "ZonoQ": if copy is not None: L = copy.L if L is None else L Q = copy.Q if Q is None else Q if L is None and Q is None: raise ValueError("At least one of L or Q must be provided.") if L is not None and Q is not None and not same_shape(L.shape_dtype, Q.shape_dtype): raise ValueError( f"L and Q must have the same shape, got {L.shape_dtype} and {Q.shape_dtype}." ) if L is None and Q is not None: L = Zono.zeros(Q.shape_dtype, Q.table) if Q is None and L is not None: Q = Quad.zeros(L.shape_dtype, L.table) assert isinstance(L, Zono) and isinstance(Q, Quad) table, _ = L.table.align_with(Q.table) return ZonoQ( L.expanded_to_table(table), Q.expanded_to_table(table), )
[docs] @staticmethod def error( err: Error, amplitude: torch.Tensor | float = 1.0, ) -> "ZonoQ": """Purely linear diagonal term; use ``Quad.error`` for a quadratic one.""" return ZonoQ.make(L=Zono.error(err, amplitude))
[docs] @classmethod @override def convert_from(cls, expr: Expr) -> Optional["ZonoQ"]: """Lift a linear or quadratic component into a correlated ``ZonoQ``. ``Error``/``Noise`` sources become purely linear; use ``Quad.convert_from`` first to obtain a quadratic one. """ if isinstance(expr, Zeros): return ZonoQ.make(L=Zono.convert_from(expr), Q=Quad.convert_from(expr)) if isinstance(expr, Error): return ZonoQ.make(L=Zono.convert_from(expr)) if isinstance(expr, Zono): return ZonoQ.make(L=expr) if isinstance(expr, Quad): return ZonoQ.make(Q=expr) if (linear := Zono.convert_from(expr)) is not None: return ZonoQ.make(L=linear) return None
[docs] @override def ub(self) -> torch.Tensor: return self.L.ub() + self.Q.ub()
[docs] @override def lb(self) -> torch.Tensor: return self.L.lb() + self.Q.lb()
[docs] @override def lbub(self) -> tuple[torch.Tensor, torch.Tensor]: linear_lower, linear_upper = self.L.lbub() quadratic_lower, quadratic_upper = self.Q.lbub() return ( linear_lower + quadratic_lower, linear_upper + quadratic_upper, )
[docs] @override def chw(self) -> tuple[torch.Tensor, torch.Tensor]: linear_center, linear_halfwidth = self.L.chw() quadratic_center, quadratic_halfwidth = self.Q.chw() return ( linear_center + quadratic_center, linear_halfwidth + quadratic_halfwidth, )
[docs] @override def einsum( self, subscripts: list[tuple[int, ...]], *operands: torch.Tensor, ) -> "ZonoQ": return ZonoQ( self.L.einsum(subscripts, *operands), self.Q.einsum(subscripts, *operands), )
[docs] def align_fill_zeros( self, other: "ZonoQ", ) -> tuple["ZonoQ", "ZonoQ"]: if not isinstance(other, ZonoQ): raise TypeError(f"other must be a ZonoQ, got {type(other)}.") if self.table == other.table: return self, other table, _ = self.table.align_with(other.table) return self.expanded_to_table(table), other.expanded_to_table(table)
[docs] def expanded_to_table(self, table: SpanTable) -> "ZonoQ": if not self.table.subset_of(table): raise ValueError("Cannot expand to a table that omits existing errors.") if self.table == table: return self return ZonoQ( self.L.expanded_to_table(table), self.Q.expanded_to_table(table), )
[docs] @override def add(self, other: Expr) -> "ZonoQ": other = other.to(ZonoQ) 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 ZonoQ( left.L.add(right.L), left.Q.add(right.Q), )
[docs] @override def reshape(self, *shape: Dim) -> "ZonoQ": return ZonoQ(self.L.reshape(*shape), self.Q.reshape(*shape))
[docs] @override def transpose(self, *perm: int) -> "ZonoQ": return ZonoQ(self.L.transpose(*perm), self.Q.transpose(*perm))
[docs] @override def broadcast_to(self, *shape: Dim) -> "ZonoQ": return ZonoQ(self.L.broadcast_to(*shape), self.Q.broadcast_to(*shape))
[docs] def reasons_breakdown(self) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: result = defaultdict( lambda: torch.zeros(self.shape_dtype.shape, dtype=self.shape_dtype.dtype) ) totalL, reasonsL = self.L.reasons_breakdown() totalQ, reasonsQ = self.Q.reasons_breakdown() for reason, value in reasonsL.items(): result[reason] += value for reason, value in reasonsQ.items(): result[reason] += value return totalL + totalQ, dict(result)
[docs] @override def torch_print( self, group: Literal["reason", "error_ref"] = "reason", ) -> TensorFormat: return TensorFormat( "{Q} + {L}", Q=self.Q.torch_print(group), L=self.L.torch_print(group), )
from .matmul import Matmul, MatmulExtraZono
[docs] def keep_zonoq_with_zono(x, name: str=""): """``after_each`` normalizer: fresh ``Noise`` becomes new ``Zono`` symbols (quadratic mass stays where the matmul handlers put it).""" del name if isinstance(x, Expr) and Noise in x.classset(): noise, other = x.split(Noise) # if all(r.startswith("mm_") for r in noise.reasons): # return noise.to(ZonoQ) + other return noise.to(Zono) + other assert isinstance(x, Expr) and x.classset().issubset({Bias, Zono, ZonoQ}) \ or isinstance(x, torch.Tensor) return x
interpret = Interpreter( base.interpret, MatmulBiased(), MatmulExtraZono(), Matmul(), linearizers.MaxWithConst2Relu(), linearizers.Relu(), linearizers.Tanh(), linearizers.Exp(), linearizers.Reciprocal(), Softmax2ExpReciprocal(), after_each=keep_zonoq_with_zono )
[docs] def precise_lbub(x: Expr) -> tuple[torch.Tensor, torch.Tensor]: """Return the built-in sound bounds without the optional Gurobi path.""" return x.lbub()
__all__ = [ "generator", "quad", "matmul", "Matmul", "MatmulExtraZono", "QGenerator", "Quad", "ZonoQ", "interpret", "keep_zonoq_with_zono", "precise_lbub", ]