Source code for boundlab.ops.deept


"""CPU-friendly primitives used by the DeepT matrix-product bound."""

from __future__ import annotations

import torch
from torch.nn import functional as F


[docs] def deept_precise_estimate( x: torch.Tensor, y: torch.Tensor, *, block_size: int = 128, ) -> tuple[torch.Tensor, torch.Tensor]: """Bound the quadratic error term in a zonotope matrix product. ``x`` and ``y`` contain the generators of an ``(m, k)`` and ``(k, n)`` matrix respectively, with the shared error-symbol dimension last. The straightforward DeepT implementation materializes ``einsum("mki,knj->mnij", x, y)``, which requires ``O(m*n*errors**2)`` memory. This implementation computes the same row and column absolute sums in blocks. It therefore retains BLAS-friendly matrix multiplications on CPU while using only ``O(m*n*errors*block_size)`` temporary memory. Args: x: Generator coefficients with shape ``(m, k, errors)``. y: Generator coefficients with shape ``(k, n, errors)``. block_size: Number of left error symbols processed at once. Returns: ``(center, halfwidth)``, both with shape ``(m, n)``. """ if x.ndim != 3 or y.ndim != 3: raise ValueError("x and y must both be rank-3 arrays") if x.shape[1] != y.shape[0]: raise ValueError( f"contracting dimensions do not match: {x.shape[1]} != {y.shape[0]}" ) if x.shape[2] != y.shape[2]: raise ValueError( "x and y must use the same number of aligned error symbols" ) if block_size <= 0: raise ValueError("block_size must be positive") error_len = x.shape[2] output_shape = (x.shape[0], y.shape[1], error_len) if error_len == 0: zeros = torch.zeros(output_shape[:2], dtype=torch.result_type(x, y)) return zeros, zeros # Padding makes every slice the same static size, which keeps this function # traceable with fixed-shape blocks. padded_len = ((error_len + block_size - 1) // block_size) * block_size x_padded = F.pad(x, (0, padded_len - error_len)) col_sums = torch.zeros(output_shape, dtype=torch.result_type(x, y)) row_blocks = [] cols = col_sums for block in range(padded_len // block_size): start = block * block_size x_block = x_padded[:, :, start : start + block_size] products = torch.einsum("mki,knj->mnij", x_block, y) absolute = products.abs() row_blocks.append(absolute.sum(3)) cols = cols + absolute.sum(2) row_sums = torch.cat(row_blocks, dim=2) col_sums = cols row_sums = row_sums[..., :error_len] diagonal = torch.einsum("mki,kni->mni", x, y) center = 0.5 * diagonal.abs().sum(-1) loose = 0.5 * (row_sums + col_sums) - diagonal.abs() halfwidth = loose.abs().sum(-1) + center return center, halfwidth
__all__ = ["deept_precise_estimate"]