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