Source code for boundlab.ibp.matmul

r"""Interval matrix products.

The bilinear form splits over components:

.. math::

   (c_1 + e_1)(c_2 + e_2) = c_1 c_2 + c_1 e_2 + e_1 c_2 + e_1 e_2 .

The three terms with a constant factor are linear maps (exact); only the
error-error product needs an interval bound,
:math:`[-a, a] \cdot [-b, b] \subseteq [-(|a| \cdot |b|),\, |a| \cdot |b|]`.
"""

import torch

from boundlab import Expr, ExprGroup
from boundlab.ibp import Bias, Noise
from boundlab.interp import Interpreter, OpHandler


[docs] class MatmulBiased(OpHandler): """Peel the ``Bias`` centers off both operands. ``(c1 + i1) @ (c2 + i2)`` expands into two exact linear maps (``i1 @ c2``, ``c1 @ i2``), a constant ``c1 @ c2``, and the remaining inner product ``i1 @ i2`` re-dispatched to whichever handler covers the error components. """ op = "matmul"
[docs] def condition(self, x, y, **kwargs): del kwargs return isinstance(x, Expr) and isinstance(y, Expr) and Bias in x.classset() and Bias in y.classset()
[docs] def handle( self, interp: Interpreter, x: Expr, y: Expr, **kwargs, ) -> Expr: del kwargs x_bias, x_inner = x.split(Bias) y_bias, y_inner = y.split(Bias) return ( (x_inner.matmul(y_bias.arr) + y_inner.rmatmul(x_bias.arr)) + interp.matmul(x_inner, y_inner) + x_bias @ y_bias.arr )
[docs] class MatmulNoise(OpHandler): r"""Interval bound for the error-error product: :math:`[-a, a] @ [-b, b] \subseteq \pm(a @ b)` for :math:`a, b \ge 0`, with the reasons blended by mass.""" op = "matmul"
[docs] def condition(self, x, y, **kwargs): del kwargs return isinstance(x, Noise) and isinstance(y, Noise)
[docs] def handle( self, interp: Interpreter, x: Noise, y: Noise, **kwargs, ) -> Expr: del interp, kwargs # [-a, a] @ [-b, b] is contained in [-(a @ b), a @ b]. x_part = x.noise.sum() y_part = y.noise.sum() # Clamp the denominator so two zero-noise operands blend to zero # weights instead of 0/0 = NaN. denom = (x_part + y_part).clamp(min=torch.finfo(x.noise.dtype).tiny) return Noise( x.noise @ y.noise, x.reasons * (x_part / denom) + y.reasons * (y_part / denom), )
__all__ = [ "MatmulBiased", "MatmulNoise", ]