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