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