Source code for boundlab.ops

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