r"""Numeric kernels shared by the abstract domains.
Closed-form extrema of scalar quadratics over the unit interval, matrix
splittings, and small custom operators. Submodules hold the heavier kernels:
:mod:`~boundlab.ops.deept` (blockwise zonotope matmul bound),
:mod:`~boundlab.ops.sparse_poly` (sparse pair products, reference + Triton),
:mod:`~boundlab.ops.legendre` (Legendre bases) and
:mod:`~boundlab.ops.unary_fn_opt` (spline-certified function ranges).
"""
from __future__ import annotations
import torch
def marked_idenity(x: torch.Tensor, **kwargs) -> torch.Tensor:
"""Emit a custom ``boundlab::MarkedIdentity`` ONNX node marking ``x``
as a special identity, so interpreters can treat it specially instead of
seeing an anonymous ``Identity``.
Outside ``torch.export`` tracing this is a plain identity —
``torch.onnx.ops.symbolic`` returns zeros when run eagerly, which would
silently corrupt concrete forward passes."""
if not torch.compiler.is_exporting():
return x
return torch.onnx.ops.symbolic(
"boundlab::MarkedIdentity",
(x,),
attrs=kwargs,
dtype=x.dtype,
shape=x.shape
)
[docs]
def definiteness_split(matrix: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
r"""Split ``matrix`` as ``pos + neg`` by shifting the shared diagonal.
Both halves keep the off-diagonal :math:`M/2`; their diagonals are shifted
by :math:`\pm(\sum_j M_{ij} + \sum_i M_{ij})/4`, pushing ``pos`` toward
diagonal dominance (hence positive semidefiniteness) and ``neg`` the
opposite way, so each half's quadratic form can be bounded one-sidedly.
"""
if len(matrix.shape) > 2:
m = matrix.reshape(-1, matrix.shape[-2], matrix.shape[-1])
positive, negative = zip(*(definiteness_split(item) for item in m))
return torch.stack(positive), torch.stack(negative)
diag_hw = (matrix.sum(-1) + matrix.sum(-2)) / 4
h_matrix = matrix / 2
pos = _with_diagonal(h_matrix, torch.diagonal(h_matrix) + diag_hw)
neg = _with_diagonal(h_matrix, torch.diagonal(h_matrix) - diag_hw)
return pos, neg
def _with_diagonal(matrix: torch.Tensor, values: torch.Tensor) -> torch.Tensor:
result = matrix.clone()
indices = torch.arange(min(result.shape[-2:]))
result[..., indices, indices] = values
return result
[docs]
def quadratic_shift(A: torch.Tensor, b: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
r"""Complete the square: :math:`x^T A x + b^T x = (x - c)^T A (x - c) + \beta`
with :math:`c = -A^{-1} b / 2` and :math:`\beta = c^T A c` returned as
``(c, bias)``."""
c = torch.linalg.solve(-2 * A, b[..., None]).squeeze(-1)
bias = torch.einsum("...i,...ij,...j->...", c, A, c)
return c, bias
[docs]
def residual_add(x: torch.Tensor, fx: torch.Tensor) -> torch.Tensor:
"""Emit a custom ``boundlab::ResidualAdd`` ONNX node marking ``x + fx``
as a residual connection, so interpreters can treat the skip path
specially instead of seeing an anonymous ``Add``."""
return torch.onnx.ops.symbolic(
"boundlab::ResidualAdd",
(x, fx),
dtype=x.dtype,
shape=x.shape
)
__all__ = [
"deept",
"legendre",
"sparse_poly",
"unary_fn_opt",
"definiteness_split",
"quadratic_shift",
"residual_add",
]