Source code for boundlab.ibp.mul

r"""Interval elementwise products, by the same component split as matmul:
``(c1 + e1)(c2 + e2)`` has three exact linear terms and one error-error term
bounded by :math:`\pm(w_1 w_2)`."""

import torch

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


[docs] class MulBiased(OpHandler): """Peel the ``Bias`` centers off both factors; the three center-carrying terms are exact linear maps, the rest is re-dispatched.""" op = "mul"
[docs] def condition(self, x, y, **kwargs): del kwargs return isinstance(x, Expr) and isinstance(y, Expr) and ( Bias in x.classset() or 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 ( interp.mul(x_inner, y_inner) + (x_inner * y_bias.arr + y_inner * x_bias.arr) + x_bias * y_bias.arr )
[docs] class MulNoised(OpHandler): """Handle the remaining error components: interval-to-interval products via ``to_intervals``, and ``Noise * Noise`` by :meth:`prop_noise`.""" op = "mul"
[docs] def condition(self, x, y, **kwargs): del kwargs return ( isinstance(x, Expr) and isinstance(y, Expr) and # Zeros operands belong to MulSimple (the product is Zeros). not isinstance(x, Zeros) and not isinstance(y, Zeros) and Bias not in x.classset() and Bias not in y.classset() and ( Noise in x.classset() or Noise in y.classset() ) )
[docs] def handle( self, interp: Interpreter, x: Expr, y: Expr, **kwargs, ) -> Expr: del kwargs x_noise, x_other = x.split(Noise) y_noise, y_other = y.split(Noise) return ( interp.mul(x_other, y_other) + interp.mul(x_other.to_intervals(), y_noise) + interp.mul(x_noise, y_other.to_intervals()) + self.prop_noise(x_noise, y_noise) )
[docs] def prop_noise(self, x: Noise, y: Noise) -> Noise: r"""Elementwise :math:`[-a, a] \cdot [-b, b] \subseteq \pm(ab)`, blending reasons by each factor's share of the total mass.""" 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__ = [ "MulBiased", "MulNoised", ]