"""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 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"\{\{|\}\}|\{([^{}:!]*)((?:[:!][^{}]*)?)\}")
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",
]