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