"""Layout bookkeeping for error symbols stored in flat coefficient axes.
Zonotopes and sparse polynomials store one coefficient slice per
:class:`~boundlab.Error` symbol along a trailing axis. A :class:`SpanTable`
records which ``start..stop`` slice (:class:`Span`) belongs to which symbol,
ordered by symbol ``ID``. When two values built over different symbol sets
meet, :meth:`SpanTable.align_with` merges the tables and returns
:class:`Alignment` runs describing how each side's slices map into the merged
layout — the domains then concatenate, zero-fill, or reindex accordingly.
"""
from __future__ import annotations
from dataclasses import dataclass
import math
from typing import Optional
import torch
from boundlab import Expr
from boundlab import utils
from boundlab.utils import Dim
def _same_dim(left: Dim, right: Dim) -> bool:
return left is right or str(left) == str(right)
def _numel(expr: Expr) -> Dim:
return math.prod(expr.shape_dtype.shape)
def _ordered(exprs: list[Expr]) -> list[Expr]:
positions = {expr: index for index, expr in enumerate(exprs)}
return sorted(
exprs,
key=lambda expr: (
getattr(expr, "ID", math.inf),
positions[expr],
),
)
[docs]
@dataclass(frozen=True)
class Span:
"""A half-open ``start..stop`` slice of the flat error-coefficient axis."""
start: Dim
stop: Dim
@property
def size(self) -> Dim:
return self.stop - self.start
def __repr__(self) -> str:
return f"{self.start}..{self.stop}"
def __str__(self) -> str:
return f"{self.start}..{self.stop}"
[docs]
@dataclass
class Alignment:
"""One run of symbols in a merged table, with its source slices.
``span1`` / ``span2`` locate the run in the two input tables (``None``
when that side does not carry the symbols) and ``span_out`` locates it in
the merged table. Adjacent runs of the same kind are merged so the
consuming code does as few slice/concat operations as possible.
"""
exprs: list[Expr]
span1: Optional[Span]
span2: Optional[Span]
span_out: Span
[docs]
def kind(self) -> str:
if self.span1 is not None and self.span2 is not None:
return "both"
if self.span1 is not None:
return "left"
if self.span2 is not None:
return "right"
raise ValueError("Alignment has no spans.")
[docs]
def merge(self, other: "Alignment") -> bool:
if self.kind() != other.kind():
return False
self.exprs += other.exprs
if self.span1 is not None:
if not _same_dim(self.span1.stop, utils.unwarp(other.span1).start):
return False
self.span1 = Span(self.span1.start, utils.unwarp(other.span1).stop)
if self.span2 is not None:
if not _same_dim(self.span2.stop, utils.unwarp(other.span2).start):
return False
self.span2 = Span(self.span2.start, utils.unwarp(other.span2).stop)
if not _same_dim(self.span_out.stop, other.span_out.start):
return False
self.span_out = Span(self.span_out.start, other.span_out.stop)
return True
[docs]
def terms_out(self) -> list[tuple[Expr, Span]]:
index_out = self.span_out.start
result = []
for expr in self.exprs:
stop = index_out + _numel(expr)
result.append((expr, Span(index_out, stop)))
index_out = stop
return result
[docs]
def subset_of(self, other: "Alignment") -> bool:
return set(self.exprs) <= set(other.exprs)
[docs]
@dataclass(frozen=True)
class SpanTable[T: Expr](dict[T, Span]):
"""Maps each error symbol to its :class:`Span`, in ``ID`` order.
The ID ordering makes layouts canonical: two tables over the same symbol
set are identical, so alignment reduces to a sorted merge.
"""
[docs]
def __init__(self, *args: T):
argli = list(args)
argli.sort() # type: ignore
start = 0
d = {}
for expr in argli:
d[expr] = Span(start, start + expr.numel())
start += expr.numel()
super().__init__(d)
self.validate()
[docs]
def validate(self) -> Dim:
"""Check spans are contiguous and ID-ordered; return the total length."""
index: Dim = 0
for expr in _ordered(list(self)):
span = self[expr] # type: ignore
if not _same_dim(span.start, index):
raise ValueError(
f"Expr {expr} starts at {span.start}, expected {index}."
)
stop = index + expr.numel()
if not _same_dim(span.stop, stop):
raise ValueError(
f"Expr {expr} stops at {span.stop}, expected {stop}."
)
index = stop
return index
[docs]
def align_with(
self,
other: "SpanTable[T]",
) -> tuple["SpanTable[T]", list[Alignment]]:
"""Merge two tables and describe how each maps into the union.
Returns the union table (over ``self ∪ other``, ID-ordered) plus a
list of merged :class:`Alignment` runs; shared symbols land in the
same output span, which is precisely what keeps them correlated
across the two values being combined.
"""
exprs = _ordered(list(dict.fromkeys((*self, *other))))
result = []
index_out: Dim = 0
for expr in exprs:
stop = index_out + expr.numel()
new_alignment = Alignment(
exprs=[expr],
span1=self.get(expr), # type: ignore
span2=other.get(expr), # type: ignore
span_out=Span(index_out, stop),
)
index_out = stop
if not result or not result[-1].merge(new_alignment):
result.append(new_alignment)
table = SpanTable(*exprs)
return table, result # type: ignore
[docs]
def subset_of(self, other: "SpanTable") -> bool:
"""Whether every symbol of this table appears in ``other``."""
return set(self) <= set(other)
def __eq__(self, other: object) -> bool:
if not isinstance(other, SpanTable):
return NotImplemented
return set(self) == set(other) and all(
_same_dim(self[expr].start, other[expr].start)
and _same_dim(self[expr].stop, other[expr].stop)
for expr in self
)
def __str__(self) -> str:
return f'SpanTable({", ".join(str(e) + str(self[e]) for e in self)})'
[docs]
def tree_flatten(self):
values = tuple(sorted(self.keys())) # type: ignore
return values, None
[docs]
@classmethod
def tree_unflatten(cls, aux, data):
assert isinstance(data, tuple) and all(isinstance(e, Expr) for e in data)
return cls(*data)
__all__ = ["Alignment", "SpanTable", "Span"]