r"""Products of two zonotopes.
The product of two affine forms is quadratic in the symbols,
.. math::
\Big(\sum_i A_i \varepsilon_i\Big)\Big(\sum_j B_j \varepsilon_j\Big)
= \sum_{i} A_i B_i\, \varepsilon_i^2
+ \sum_{i \ne j} A_i B_j\, \varepsilon_i \varepsilon_j ,
which no affine form represents exactly; an *estimate* maps it to a center
plus interval noise. The payoff of doing this carefully is the diagonal:
:math:`\varepsilon_i^2 \in [0, 1]` (not :math:`[-1, 1]`), so its
contribution has center :math:`\tfrac12\sum_i |A_i B_i|` and only half the
naive width, while each cross term stays a plain :math:`[-1, 1]` factor.
"""
from __future__ import annotations
from collections.abc import Callable
import torch
from boundlab import Expr
from boundlab.ibp import Noise
from boundlab.utils import same_shape
from boundlab.interp import OpHandler
from boundlab.ops.deept import deept_precise_estimate
from boundlab.zono import Zono
from boundlab.zono.generator import Generator
[docs]
def basic_estimate(x: Generator, y: Generator) -> tuple[torch.Tensor, torch.Tensor]:
r"""Box bound: :math:`|x| \le ub_x,\ |y| \le ub_y \Rightarrow |xy| \le ub_x\,ub_y`
— zero center, width ``x.ub() @ y.ub()``."""
output_shape = (x.shape_dtype.shape[0], y.shape_dtype.shape[1])
return torch.zeros(output_shape, dtype=x.tensor.dtype), x.ub() @ y.ub()
[docs]
def square_estimate(x: Generator, y: Generator) -> tuple[torch.Tensor, torch.Tensor]:
"""Alias of :func:`basic_estimate` (kept as a named dispatch target)."""
return basic_estimate(x, y)
[docs]
def precise_estimate_orig(x: Generator, y: Generator) -> tuple[torch.Tensor, torch.Tensor]:
r"""Reference DeepT-precise estimate (materializes the full pair tensor).
Forms every product coefficient :math:`T_{mnij} = \sum_k A_{mki} B_{knj}`,
then bounds the diagonal (:math:`\varepsilon_i^2 \in [0, 1]`) exactly —
center :math:`\tfrac12\sum_i |T_{mnii}|` — and every cross term by
:math:`|T_{mnij}|`, halved-and-summed over both orderings. Memory is
:math:`O(mn \cdot \text{errors}^2)`; :func:`precise_estimate` computes the
same bound blockwise.
"""
products = torch.einsum("mki,knj->mnij", x.tensor, y.tensor)
# sums = products.abs().sum((0, 1))
# def log_sums(sums: torch.Tensor) -> None:
# import time
# import matplotlib
# matplotlib.use("agg")
# import matplotlib.pyplot as plt
# import numpy as np
# values = np.asarray(sums).ravel()
# positive = values[values > 0]
# figure, axis = plt.subplots()
# if positive.size:
# axis.hist(positive, bins=50)
# axis.set(
# xlabel="sums",
# ylabel="Count",
# title=f"Distribution of product sums ({values.size - positive.size} zeros)",
# )
# figure.tight_layout()
# figure.savefig(f"sums_distribution_{time.time_ns()}.png", dpi=150)
# plt.close(figure)
# ``log_sums`` can be called here when distribution diagnostics are needed.
diagonal = torch.diagonal(products, dim1=-2, dim2=-1)
loose = 0.5 * (
products.abs().sum(3)
+ products.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]
def precise_estimate(x: Generator, y: Generator, block_size: int = 128) -> tuple[torch.Tensor, torch.Tensor]:
r"""The DeepT-precise bound, computed in error-symbol blocks so the
:math:`mn \cdot e^2` pair tensor is never materialized (see
:func:`boundlab.ops.deept.deept_precise_estimate`)."""
return deept_precise_estimate(x.tensor, y.tensor, block_size=block_size)
[docs]
class Matmul(OpHandler):
"""Zonotope × zonotope matrix product: align symbol tables, run the
configured estimate, return ``center + Noise(width, "matmul")``."""
op = "matmul"
[docs]
def __init__(
self,
matmul_estimate: Callable[
[Generator, Generator],
tuple[torch.Tensor, torch.Tensor],
] = precise_estimate,
):
self.matmul_estimate = matmul_estimate
[docs]
def condition(self, x, y, **kwargs):
del kwargs
return (
isinstance(x, Zono)
and isinstance(y, Zono)
and x.ndim == 2
and y.ndim == 2
)
[docs]
def handle(self, interp, x: Zono, y: Zono, **kwargs) -> Expr:
del kwargs
assert isinstance(x, Zono) and isinstance(y, Zono)
x, y = x.align_fill_zeros(y)
c, hw = self.matmul_estimate(x.gen, y.gen)
return Noise(hw, "matmul") + c
__all__ = [
"Matmul",
"basic_estimate",
"precise_estimate",
"precise_estimate_orig",
"square_estimate",
]