Source code for boundlab.zonoq.matmul

r"""Matrix products that keep the quadratic term symbolic.

For two linear zonotopes the product is exactly quadratic:

.. math::

   \Big(\sum_i A_i \varepsilon_i\Big)\Big(\sum_j B_j \varepsilon_j\Big)
   = \sum_{ij} (A_i B_j)\, \varepsilon_i \varepsilon_j ,

so :class:`Matmul` stores :math:`A_i B_j` in a
:class:`~boundlab.zonoq.quad.Quad` unchanged; only the linear-times-quadratic
and quadratic-times-quadratic cross terms (degree 3 and 4) are bounded into
interval noise.
"""

from dataclasses import dataclass

import torch

from boundlab import Expr
from boundlab.ibp import Bias, Noise
from boundlab.interp import OpHandler
from boundlab.zono import Zono
from boundlab.zono.matmul import precise_estimate
from boundlab.zonoq import ZonoQ
from boundlab.zonoq.generator import QGenerator
from boundlab.zonoq.quad import Quad

[docs] @dataclass class Matmul(OpHandler): r"""ZonoQ × ZonoQ product: :math:`L_1 L_2` exactly into the quadratic part; :math:`L Q`, :math:`Q L` and :math:`Q Q` terms box-bounded as noise (degree ≥ 3 has no representation here).""" op: str = "matmul"
[docs] def condition(self, x, y, **kwargs): del kwargs return ( isinstance(x, ZonoQ) and isinstance(y, ZonoQ) and x.ndim == 2 and y.ndim == 2 )
[docs] def handle(self, interp, x: ZonoQ, y: ZonoQ, **kwargs) -> Expr: del kwargs assert isinstance(x, ZonoQ) and isinstance(y, ZonoQ) x, y = x.align_fill_zeros(y) table = x.L.table noise = Noise(x.L.ub() @ y.Q.ub() + x.Q.ub() @ y.L.ub(), "mm_lq") + Noise(x.Q.ub() @ y.Q.ub(), "mm_qq") Q = Quad( table, QGenerator( torch.einsum( "mki,knj->mnij", x.L.gen.tensor, y.L.gen.tensor ) ) ) return ZonoQ(Zono.zeros(Q.shape_dtype, table), Q) + noise
[docs] @dataclass class MatmulExtraZono(OpHandler): """Matrix product for a ZonoQ with an independent Zono remainder.""" op: str = "matmul"
[docs] def condition(self, x, y, **kwargs): del kwargs return ( isinstance(x, Expr) and isinstance(y, Expr) and (Zono in x.classset() or Zono in y.classset()) and Bias not in x.classset() and Bias not in y.classset() and x.ndim == 2 and y.ndim == 2 )
[docs] def handle(self, interp, x: Expr, y: Expr, **kwargs) -> Expr: del kwargs x_extra, x_inner = x.split(Zono) y_extra, y_inner = y.split(Zono) x_extra = x_extra.to(Zono) y_extra = y_extra.to(Zono) cx, hwx = x_inner.chw() cy, hwy = y_inner.chw() inner = interp.matmul(x_inner, y_inner) noise, inner = inner.split(Noise) bias = cx @ y_extra + x_extra @ cy x_extra, y_extra = x_extra.align_fill_zeros(y_extra) center, error = precise_estimate(x_extra.gen, y_extra.gen) bias += center noise += Noise(error, "mm_zz") noise += Noise(hwx @ y_extra.ub() + x_extra.ub() @ hwy, "mm_zlzq") return inner + noise + bias
__all__ = [ "Matmul", "MatmulExtraZono", ]