Source code for boundlab.utils

"""Shared utilities: shape metadata, einsum plumbing, polynomials, formatting."""

from __future__ import annotations

from collections.abc import Callable, Sequence
from dataclasses import dataclass
import operator
from typing import Any, Callable, Iterable, Optional, final
import typing

from beartype.door import die_if_unbearable
import torch

from boundlab import utils


Dim = Any
"""A tensor dimension: a plain ``int`` or a symbolic size."""


[docs] @dataclass(frozen=True) class ShapeDtype: """Allocation-free shape and dtype metadata for an abstract expression.""" shape: tuple[Dim, ...] dtype: torch.dtype
[docs] def __init__(self, shape: Sequence[Dim], dtype: torch.dtype): object.__setattr__(self, "shape", tuple(shape)) object.__setattr__(self, "dtype", dtype)
[docs] @classmethod def from_value(cls, value: torch.Tensor | "ShapeDtype") -> "ShapeDtype": return value if isinstance(value, cls) else cls(value.shape, value.dtype)
[docs] def same_shape( left: Sequence[Dim] | ShapeDtype, right: Sequence[Dim] | ShapeDtype, ) -> bool: """Whether two shapes agree, comparing symbolic dimensions by string.""" left_shape = left.shape if isinstance(left, ShapeDtype) else left right_shape = right.shape if isinstance(right, ShapeDtype) else right return len(left_shape) == len(right_shape) and all( left_dim is right_dim or str(left_dim) == str(right_dim) for left_dim, right_dim in zip(left_shape, right_shape) )
[docs] def einsum_parser(subscripts: str) -> list[tuple[int, ...]]: """ Parse an einsum string into input tuples followed by one output tuple. """ subscripts = subscripts.replace(" ", "") if subscripts.count("->") != 1: raise ValueError("einsum subscripts must contain one explicit '->'.") inputs, output = subscripts.split("->") parts = inputs.split(",") + [output] labels: dict[str, int] = {} result = [] for part in parts: if any(not label.isalpha() or len(label) != 1 for label in part): raise ValueError(f"Unsupported einsum subscript {part!r}.") result.append( tuple(labels.setdefault(label, len(labels)) for label in part) ) input_labels = set().union(*(set(part) for part in result[:-1])) if any(label not in input_labels for label in result[-1]): raise ValueError("einsum output labels must appear in an input.") return result
[docs] def einsum_formatter(subscripts: list[tuple[int, ...]]) -> str: """ Format input tuples followed by one output tuple as an einsum string. """ if len(subscripts) < 2: raise ValueError("einsum subscripts require at least one input and one output.") alphabet = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" labels: dict[int, str] = {} for subscript in subscripts: for label in subscript: if label not in labels: if len(labels) == len(alphabet): raise ValueError("einsum supports at most 52 distinct labels.") labels[label] = alphabet[len(labels)] formatted = ["".join(labels[label] for label in part) for part in subscripts] return ",".join(formatted[:-1]) + "->" + formatted[-1]
[docs] def format_bytes(n) -> str: """Format a byte count as a human-readable string (e.g. ``1.5MiB``).""" if not isinstance(n, int): # Symbolic sizes (export shapes) cannot be scaled to a unit. return f"{n}B" size = float(n) for unit in ("B", "KiB", "MiB", "GiB"): if size < 1024: return f"{size:.3g}{unit}" size /= 1024 return f"{size:.3g}TiB"
[docs] def all_same(iterable: Iterable) -> bool: """Whether every element equals the first (vacuously true when empty).""" return len(set(iterable)) <= 1
[docs] def all_unique(iterable: Iterable) -> bool: """Whether no element repeats.""" """ Check if all elements in an iterable are unique. """ seen = set() for item in iterable: if item in seen: return False seen.add(item) return True
[docs] def unzip[T](iterable: Iterable[tuple[T, ...]]) -> list[list[T]]: """Transpose an iterable of equal-length tuples into per-slot lists.""" output = [] for tup in iterable: if len(output) == 0: output = [[] for _ in tup] else: assert len(output) == len(tup), "All tuples must have the same length." for i, item in enumerate(tup): output[i].append(item) return output
[docs] def is_statically_zero(x): """Return whether a concrete scalar/tensor contains only zeros.""" if isinstance(x, torch.Tensor): return bool(torch.all(x == 0)) return bool((x == 0).all()) if hasattr(x, "all") else x == 0
[docs] class TyDict[T](dict[type[T], T]): """Dictionary whose values are indexed by their concrete types."""
[docs] def __init__(self, *values: T) -> None: super().__init__() for value in values: key = value.__class__ if key in self: raise ValueError(f"Duplicate type {key} in TyDict.") self[key] = value
def __getitem__[U](self, key: type[U]) -> U: return cast(U, super().__getitem__(key)) # type: ignore def __setitem__[U](self, key: type[U], value: U) -> None: if value.__class__ is not key: raise TypeError( f"Value must be exactly {key.__name__}, got {type(value).__name__}" ) super().__setitem__(key, value) # type: ignore
[docs] def keyset(self) -> set[type[T]]: return set(super().keys()) # type: ignore
import re # One "{key}" or "{key:spec}" placeholder: group 1 is the key, group 2 the # spec. "{{" and "}}" are format-string escapes; matching them first keeps # their inner braces from being mistaken for a placeholder. _PLACEHOLDER = re.compile(r"\{\{|\}\}|\{([^{}:!]*)((?:[:!][^{}]*)?)\}")
[docs] @dataclass(frozen=True) @final class TensorFormat: """A format string plus the traced arrays that fill its placeholders. A kwarg may itself be a ``TensorFormat``; its format is spliced in place of the placeholder and its keys are renamed to stay unique. """ fmt: str kwargs: dict[str, Any]
[docs] def __init__(self, fmt: str, **kwargs: Any | TensorFormat): new_kwargs = {} def repl(match: re.Match) -> str: inner, spec = match.group(1), match.group(2) if inner is None: # "{{" or "}}" escape return match.group(0) if inner in kwargs and isinstance(kwargs[inner], TensorFormat): jaxfmt = cast(TensorFormat, kwargs[inner]).add_identifier(f"_{match.start()}") new_kwargs.update(jaxfmt.kwargs) return jaxfmt.fmt elif inner in kwargs and isinstance(kwargs[inner], str): # Splice literal strings into the format itself, escaped so # they survive the final str.format untouched. return cast(str, kwargs[inner]).replace("{", "{{").replace("}", "}}") elif inner in kwargs: new_kwargs[inner] = kwargs[inner] return f"{{{inner}{spec}}}" else: raise ValueError(f"Key {inner} not found in kwargs.") new_fmt = _PLACEHOLDER.sub(repl, fmt) object.__setattr__(self, "fmt", new_fmt) object.__setattr__(self, "kwargs", new_kwargs)
[docs] def add_identifier(self, identifier: str) -> TensorFormat: def repl(match: re.Match) -> str: inner, spec = match.group(1), match.group(2) if inner is None: # "{{" or "}}" escape return match.group(0) if inner in self.kwargs: return f"{{{inner}{identifier}{spec}}}" else: raise ValueError(f"Key {inner} not found in kwargs.") renamed = TensorFormat.__new__(TensorFormat) object.__setattr__(renamed, "fmt", _PLACEHOLDER.sub(repl, self.fmt)) object.__setattr__( renamed, "kwargs", {f"{k}{identifier}": v for k, v in self.kwargs.items()}, ) return renamed
[docs] def render(self) -> str: values = { key: value.detach().cpu().item() if isinstance(value, torch.Tensor) and value.numel() == 1 else value for key, value in self.kwargs.items() } return self.fmt.format(**values)
[docs] def print(self): print(self.render())
[docs] def __add__(self, other: str | TensorFormat) -> TensorFormat: return TensorFormat("{a}{b}", a=self, b=other)
def __radd__(self, other: str) -> TensorFormat: return TensorFormat("{a}{b}", a=other, b=self)
[docs] @staticmethod def join( iterable: Iterable[str | TensorFormat], sep: str | TensorFormat = "", ) -> TensorFormat: parts = list(iterable) if not parts: return TensorFormat("") fmt = "{sep}".join(f"{{part{i}}}" for i in range(len(parts))) kwargs = {"sep": sep} for i, item in enumerate(parts): kwargs[f"part{i}"] = item return TensorFormat(fmt, **kwargs)
type BinOp = Callable[[torch.Tensor| float, torch.Tensor| float], torch.Tensor | float] def _stack(coeffs: Sequence[torch.Tensor | float]) -> torch.Tensor: """Stack polynomial coefficients, lifting any plain Python scalars.""" return torch.stack([torch.as_tensor(coeff) for coeff in coeffs])
[docs] @dataclass(frozen=True) class Polynomial: r"""Dense univariate polynomial :math:`P(x) = \sum_k c_k x^k`. ``coeffs`` runs from the constant term up. Coefficients may be tensors, giving an independent polynomial per element. Besides ring arithmetic it supports composition (:meth:`chain`), argument substitution (:meth:`apply_add`, :meth:`apply_ax`), differentiation, and sound interval evaluation (:meth:`ibp`). """ coeffs: list[torch.Tensor | float]
[docs] def __init__(self, coeffs: Sequence[torch.Tensor | float]): object.__setattr__(self, "coeffs", list(coeffs))
@property def degree(self) -> int: return len(self.coeffs) - 1
[docs] @staticmethod def X(n: int=1) -> "Polynomial": """Polynomial representing x^n.""" return Polynomial([0.0] * n + [1.0])
[docs] def __add__(self, other: torch.Tensor | float | Polynomial) -> "Polynomial": if isinstance(other, Polynomial): new_coeffs = [ a + b for a, b in zip(self.coeffs, other.coeffs) ] if len(self.coeffs) > len(other.coeffs): new_coeffs += self.coeffs[len(other.coeffs):] else: new_coeffs += other.coeffs[len(self.coeffs) :] return Polynomial(new_coeffs) else: return Polynomial([self.coeffs[0] + other] + list(self.coeffs[1:]))
[docs] def __mul__(self, other: torch.Tensor | float| Polynomial) -> "Polynomial": return self.mul(other, operator.mul)
def __rmul__(self, other: torch.Tensor | float) -> "Polynomial": return self.mul(other, operator.mul)
[docs] def mul(self, other: torch.Tensor | float | Polynomial, mul: BinOp = operator.mul) -> "Polynomial": if isinstance(other, Polynomial): new_coeffs: list[torch.Tensor | float] = [0.0] * (self.degree + other.degree + 1) for i, a in enumerate(self.coeffs): for j, b in enumerate(other.coeffs): new_coeffs[i + j] += mul(a, b) return Polynomial(new_coeffs) else: return Polynomial([mul(coefficient, other) for coefficient in self.coeffs])
[docs] def apply_add(self, shift: torch.Tensor | float, mul: BinOp = operator.mul) -> "Polynomial": """P(x + shift)""" new_poly = Polynomial([self.coeffs[-1]]) for coefficient in reversed(self.coeffs[:-1]): new_poly = Polynomial([0.0, *new_poly.coeffs]) + new_poly.mul(shift, mul) + Polynomial([coefficient]) return new_poly
[docs] def apply_ax(self, a: torch.Tensor | float, mul: BinOp = operator.mul) -> "Polynomial": """P(a x)""" scale = 1.0 new_coeffs = [self.coeffs[0]] for coefficient in self.coeffs[1:]: scale = mul(scale, a) new_coeffs.append(mul(coefficient, scale)) return Polynomial(new_coeffs)
[docs] def apply_ax_reversed(self, a: Any, mul: BinOp = operator.mul) -> "Polynomial": """P(x / a) * a^n""" scale = 1.0 new_coeffs = [self.coeffs[-1]] for coefficient in reversed(self.coeffs[:-1]): scale = mul(scale, a) new_coeffs.append(mul(coefficient, scale)) return Polynomial(list(reversed(new_coeffs)))
[docs] def __call__(self, x: torch.Tensor | float) -> torch.Tensor | float: result = self.coeffs[-1] for coefficient in reversed(self.coeffs[:-1]): result = result * x + coefficient return result
[docs] def chain(self, other: Polynomial, mul: BinOp = operator.mul) -> Polynomial: """P(Q(x))""" result = Polynomial([0.0]) for coefficient in reversed(self.coeffs): result = result.mul(other, mul) + Polynomial([coefficient]) return result
[docs] def derivative(self, order: int =1) -> "Polynomial": if order < 0: raise ValueError("Derivative order must be non-negative.") if order > self.degree: return Polynomial([0.0]) new_coeffs = self.coeffs[:] for _ in range(order): new_coeffs = [ (i + 1) * new_coeffs[i + 1] for i in range(len(new_coeffs) - 1) ] return Polynomial(new_coeffs)
[docs] def ibp(self, c: torch.Tensor, hw: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: r"""Sound range of :math:`P` over :math:`[c - hw,\, c + hw]`. Substituting :math:`x = c + hw\,u` reduces to bounding over :math:`u \in [-1, 1]`, where each monomial is elementary: odd powers of ``u`` range over :math:`\pm|a_k|` and even powers over :math:`[\min(a_k, 0), \max(a_k, 0)]`. Summing those per-monomial ranges is sound (though not tight — monomial correlations are ignored). """ chained = self.chain(Polynomial([c, hw])) c0 = chained.coeffs[0] odd = chained.coeffs[1::2] even = chained.coeffs[2::2] # Degree <= 1 leaves odd and/or even empty; an empty stack is a zero term. odd_sum = _stack(odd).abs().sum(0) if odd else 0.0 even_lo = _stack(even).clamp(max=0.0).sum(0) if even else 0.0 even_hi = _stack(even).clamp(min=0.0).sum(0) if even else 0.0 return c0 - odd_sum + even_lo, c0 + odd_sum + even_hi
[docs] def sum[T](iter: Iterable[T]) -> T: """Sum with ``+`` and no zero start, so expression types keep their own ``__add__`` semantics; raises on an empty iterable.""" it = iter.__iter__() try: total = next(it) except StopIteration: raise ValueError("sum() of empty iterable with no initial value") for item in it: total += item # type: ignore return total
[docs] def unwarp[T](item: Optional[T]) -> T: """Assert an ``Optional`` is present and return it.""" if item is not None: return item raise ValueError("Expected a non-None value.")
[docs] def cast[T](ty: type[T], value: Any) -> T: """Runtime-checked cast: dies (via beartype) if ``value`` is not a ``ty``.""" die_if_unbearable(value, ty) return typing.cast(T, value)
__all__ = [ "Dim", "Polynomial", "ShapeDtype", "TensorFormat", "TyDict", "all_same", "all_unique", "cast", "einsum_formatter", "einsum_parser", "format_bytes", "is_statically_zero", "same_shape", "sum", "unwarp", "unzip", ]