"""Runtime spatial Set-BPE vocabulary and lossless row serialization.""" from __future__ import annotations import heapq from collections.abc import Iterable, Sequence from dataclasses import dataclass import numpy as np from ..common._validation import active_dimensions __all__ = ["SpatialVocabulary", "aggregate_records"] def aggregate_records( codes: np.ndarray, bins: Sequence[int] | None = None ) -> tuple[np.ndarray, np.ndarray]: """Lossless micro-aggregation of identical quantized actions. Replacing ``w`` identical records by one record of weight ``w`` leaves every pair frequency and therefore every greedy decision unchanged, because identical actions follow identical merges forever. """ values = np.asarray(codes, dtype=np.int64) if values.ndim != 2: raise ValueError("codes must have shape [U, D]") if bins is not None: dimensions = tuple(int(count) for count in bins) if len(dimensions) != values.shape[1]: raise ValueError("bins must have one entry per code coordinate") counts = np.asarray(dimensions, dtype=np.int64) valid = values >= 0 if np.any(values < -1) or np.any(values >= counts[None, :]): raise ValueError("codes must contain -1 or an in-range cell index") cardinality = int(np.prod(np.asarray(dimensions, dtype=object))) if np.all(valid) and cardinality <= np.iinfo(np.int64).max: keys = np.ravel_multi_index(values.T, dimensions) unique_keys, counts = np.unique(keys, return_counts=True) unique = np.stack(np.unravel_index(unique_keys, dimensions), axis=1) return unique.astype(np.int64, copy=False), counts.astype(np.int64) unique, counts = np.unique(values, axis=0, return_counts=True) return unique, counts.astype(np.int64) @dataclass class SpatialVocabulary: """Atoms, merge rules and the decoder metadata they induce.""" bins: tuple[int, ...] merges: tuple[tuple[int, int], ...] def __post_init__(self) -> None: if not self.bins or any(count < 1 for count in self.bins): raise ValueError("every dimension needs at least one bin") self.offsets = np.concatenate( ([0], np.cumsum(np.asarray(self.bins, dtype=np.int64))) ) self.atom_count = int(self.offsets[-1]) self.vocab_size = self.atom_count + len(self.merges) dimension = len(self.bins) support = np.zeros(self.vocab_size, dtype=np.int64) expansion = np.full((self.vocab_size, dimension), -1, dtype=np.int64) for index, count in enumerate(self.bins): for value in range(count): token = int(self.offsets[index]) + value support[token] = 1 << index expansion[token, index] = value for rank, (left, right) in enumerate(self.merges): token = self.atom_count + rank if left >= token or right >= token: raise ValueError("a merge rule must reference already-created tokens") if support[left] & support[right]: raise ValueError("a merge rule joins two overlapping supports") support[token] = support[left] | support[right] merged = np.maximum(expansion[left], expansion[right]) expansion[token] = merged self.support = support self.expansion = expansion self.anchor = np.asarray( [int(mask & -mask).bit_length() - 1 for mask in support], dtype=np.int64 ) self.full_support = (1 << dimension) - 1 self._rules: dict[tuple[int, int], tuple[int, int]] = {} for rank, (left, right) in enumerate(self.merges): key = (min(left, right), max(left, right)) self._rules[key] = (self.atom_count + rank, rank) @property def dimension(self) -> int: return len(self.bins) def atom(self, dimension: int, value: int) -> int: return int(self.offsets[dimension]) + int(value) def encode( self, code: Sequence[int], *, active_dims: Sequence[int] | None = None, ) -> list[int]: """Deterministic online Set-BPE encoding (Algorithm 8). Merge rules are applied in creation-rank order, which is what makes the encoder reproducible: the same quantized action always yields the same token set regardless of how the search that learned the rules ran. ``-1`` in a dense row, or a missing coordinate in ``active_dims``, means semantic absence: it emits no atom and can never participate in a merge. """ compact = [int(value) for value in code] if active_dims is None: if len(compact) != self.dimension: raise ValueError( f"a dense quantized action must have {self.dimension} slots" ) pairs = enumerate(compact) else: dimensions = active_dimensions(active_dims, self.dimension) if len(dimensions) != len(compact): raise ValueError( "active_dims must have one entry per compact coordinate" ) pairs = zip(dimensions, compact, strict=True) tokens = set() for index, value in pairs: if value == -1: if active_dims is not None: raise ValueError("compact coordinates must not contain -1") continue if not 0 <= value < self.bins[index]: raise ValueError(f"bin {value} is out of range for dimension {index}") tokens.add(self.atom(index, value)) if not tokens: raise ValueError( "a quantized action must contain at least one active coordinate" ) heap: list[tuple[int, int, int, int]] = [] ordered = sorted(tokens) for position, left in enumerate(ordered): for right in ordered[position + 1 :]: rule = self._rules.get((left, right)) if rule is not None: heapq.heappush(heap, (rule[1], rule[0], left, right)) while heap: _, child, left, right = heapq.heappop(heap) if left not in tokens or right not in tokens: continue tokens.discard(left) tokens.discard(right) for other in tokens: key = (min(child, other), max(child, other)) rule = self._rules.get(key) if rule is not None: heapq.heappush(heap, (rule[1], rule[0], key[0], key[1])) tokens.add(child) return self.serialize(tokens) def serialize(self, tokens: Iterable[int]) -> list[int]: """Canonical order for autoregressive supervision. Set semantics carry no order, but a policy must be trained on one fixed target. Tokens within an action have pairwise disjoint supports, so their anchor dimensions are distinct and sorting by anchor is a total order. Anchor order is preferred over raw token id because it is the one that keeps constrained decoding feasible: the next token's anchor is always the lowest uncovered dimension, and a primitive atom for that dimension always exists, so no prefix can paint itself into a corner. """ return sorted(tokens, key=lambda token: (int(self.anchor[token]), int(token))) def decode( self, tokens: Sequence[int], *, active_dims: Sequence[int] | None = None, ) -> np.ndarray: """Expand tokens to a dense row, using ``-1`` for absent coordinates.""" code = np.full(self.dimension, -1, dtype=np.int64) covered = 0 for token in tokens: value = int(token) if not 0 <= value < self.vocab_size: raise ValueError(f"token {value} is outside the vocabulary") if covered & int(self.support[value]): raise ValueError("token supports overlap; this is not a valid action") covered |= int(self.support[value]) local = self.expansion[value] code = np.where(local >= 0, local, code) if active_dims is None: expected = self.full_support else: dimensions = active_dimensions(active_dims, self.dimension) expected = sum(1 << dimension for dimension in dimensions) if covered != expected: raise ValueError( "token set does not exactly cover the requested active dimensions" ) return code