Source code for boundlab.error.alignment

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