Source code for boundlab.zono.matmul

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