ZibinDong's picture
Upload pretrained ActionCodec2 artifact
fee0e43 verified
Raw History Blame Contribute Delete
8.68 kB
"""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