loss-guided-static-multi / recursive_cascade.py
BorisTM's picture
Add files using upload-large-folder tool
309d3a9 verified
Raw History Blame Contribute Delete
47.3 kB
"""Recursive criterion-grown token vocabulary primitives.
This module is intentionally small and CPU-testable. It establishes the Q32
mechanism contract before the recursive tokenizer is connected to the streaming
trainer: learned tokens are emitted as discrete ids, may parent later tokens,
receive gradients, and resume with exact optimizer state.
"""
from __future__ import annotations
import copy
from dataclasses import dataclass
import hashlib
import json
import math
from typing import Any, Iterable, Sequence
import torch
from torch import nn
from .collapse import _weighted_index_add_
@dataclass(frozen=True)
class CascadeToken:
token_id: int
span: tuple[int, ...]
generation: int
left_parent: int | None
right_parent: int | None
class RecursiveCascadeVocabulary(nn.Module):
"""A monotone DAG vocabulary over canonical base-token spans."""
def __init__(
self,
base_embeddings: torch.Tensor,
*,
max_span_length: int = 32,
) -> None:
super().__init__()
if base_embeddings.ndim != 2 or not base_embeddings.shape[0]:
raise ValueError("base_embeddings must be a nonempty matrix")
if not base_embeddings.is_floating_point() or not bool(
torch.isfinite(base_embeddings).all()
):
raise ValueError("base_embeddings must be finite floating point")
if max_span_length < 2:
raise ValueError("max_span_length must be at least two")
self.dim = int(base_embeddings.shape[1])
self.max_span_length = int(max_span_length)
self.base_size = int(base_embeddings.shape[0])
self.embeddings = nn.ParameterList(
nn.Parameter(row.detach().clone()) for row in base_embeddings
)
self._tokens = tuple(
CascadeToken(index, (index,), 0, None, None)
for index in range(self.base_size)
)
@property
def tokens(self) -> tuple[CascadeToken, ...]:
return self._tokens
def _span_index(self) -> dict[tuple[int, ...], int]:
return {token.span: token.token_id for token in self._tokens}
def tokenize(self, base_ids: Sequence[int]) -> tuple[int, ...]:
"""Greedily emit longest canonical spans with stable id tie-breaking."""
sequence = tuple(int(value) for value in base_ids)
if any(value < 0 or value >= self.base_size for value in sequence):
raise ValueError("tokenize input must contain only base token ids")
spans = self._span_index()
by_first: dict[int, list[tuple[tuple[int, ...], int]]] = {}
for span, token_id in spans.items():
by_first.setdefault(span[0], []).append((span, token_id))
for candidates in by_first.values():
candidates.sort(key=lambda item: (-len(item[0]), item[1]))
emitted: list[int] = []
position = 0
while position < len(sequence):
chosen_span = (sequence[position],)
chosen_id = sequence[position]
for span, token_id in by_first.get(sequence[position], ()):
if sequence[position : position + len(span)] == span:
chosen_span, chosen_id = span, token_id
break
emitted.append(chosen_id)
position += len(chosen_span)
return tuple(emitted)
@staticmethod
def adjacent_pairs(token_ids: Sequence[int]) -> tuple[tuple[int, int], ...]:
emitted = tuple(int(value) for value in token_ids)
return tuple(zip(emitted, emitted[1:]))
def lookup(self, token_ids: Sequence[int]) -> torch.Tensor:
ids = tuple(int(value) for value in token_ids)
if not ids:
raise ValueError("lookup requires at least one token")
if any(value < 0 or value >= len(self.embeddings) for value in ids):
raise ValueError("lookup token id is outside the vocabulary")
return torch.stack([self.embeddings[value] for value in ids])
def encode_mean(self, base_ids: Sequence[int]) -> tuple[torch.Tensor, tuple[int, ...]]:
emitted = self.tokenize(base_ids)
return self.lookup(emitted).mean(dim=0), emitted
@staticmethod
def _average_parent_state(
optimizer: torch.optim.Optimizer,
left: nn.Parameter,
right: nn.Parameter,
) -> dict[Any, Any]:
left_state = optimizer.state.get(left, {})
right_state = optimizer.state.get(right, {})
if not left_state or set(left_state) != set(right_state):
raise ValueError("both parents require aligned initialized optimizer state")
result: dict[Any, Any] = {}
for name in left_state:
left_value, right_value = left_state[name], right_state[name]
if torch.is_tensor(left_value) != torch.is_tensor(right_value):
raise ValueError(f"optimizer parent state type differs: {name}")
if not torch.is_tensor(left_value):
if left_value != right_value:
raise ValueError(f"optimizer parent scalar differs: {name}")
result[name] = left_value
elif left_value.ndim == 0:
if not torch.equal(left_value, right_value):
raise ValueError(f"optimizer parent step differs: {name}")
result[name] = left_value.detach().clone()
else:
if left_value.shape != left.shape or right_value.shape != right.shape:
raise ValueError(f"optimizer parent row shape differs: {name}")
result[name] = (0.5 * (left_value + right_value)).detach().clone()
return result
def promote_pairs(
self,
pairs: Iterable[tuple[int, int]],
*,
optimizer: torch.optim.Optimizer,
) -> tuple[CascadeToken, ...]:
"""Materialize every new canonical span and attach it to Adam in place."""
requested = tuple((int(left), int(right)) for left, right in pairs)
if not requested:
return ()
if len(set(requested)) != len(requested):
raise ValueError("promotion pairs must be unique")
existing = self._span_index()
additions: list[CascadeToken] = []
for left_id, right_id in requested:
if not (0 <= left_id < len(self._tokens) and 0 <= right_id < len(self._tokens)):
raise ValueError("promotion parent id is outside the vocabulary")
left_token, right_token = self._tokens[left_id], self._tokens[right_id]
span = left_token.span + right_token.span
if len(span) > self.max_span_length:
raise ValueError("promotion exceeds max_span_length")
if span in existing:
raise ValueError("promotion span already exists")
generation = max(left_token.generation, right_token.generation) + 1
token_id = len(self._tokens) + len(additions)
parameter = nn.Parameter(
0.5
* (
self.embeddings[left_id].detach()
+ self.embeddings[right_id].detach()
)
)
state = self._average_parent_state(
optimizer, self.embeddings[left_id], self.embeddings[right_id]
)
self.embeddings.append(parameter)
optimizer.param_groups[0]["params"].append(parameter)
optimizer.state[parameter] = state
token = CascadeToken(token_id, span, generation, left_id, right_id)
additions.append(token)
existing[span] = token_id
self._tokens = (*self._tokens, *additions)
return tuple(additions)
def snapshot(self, optimizer: torch.optim.Optimizer) -> dict[str, Any]:
return {
"version": 1,
"protocol": "recursive-cascade-micro-state-v1",
"base_size": self.base_size,
"dim": self.dim,
"max_span_length": self.max_span_length,
"tokens": [
{
"token_id": token.token_id,
"span": list(token.span),
"generation": token.generation,
"left_parent": token.left_parent,
"right_parent": token.right_parent,
}
for token in self._tokens
],
# PyTorch state_dict values alias live parameter/optimizer storage.
# A resume snapshot must be immutable while the source run keeps
# training, so clone the complete trees at the snapshot boundary.
"model": copy.deepcopy(self.state_dict()),
"optimizer": copy.deepcopy(optimizer.state_dict()),
}
@classmethod
def from_snapshot(
cls,
payload: dict[str, Any],
*,
optimizer_kwargs: dict[str, Any],
) -> tuple["RecursiveCascadeVocabulary", torch.optim.AdamW]:
if (
not isinstance(payload, dict)
or payload.get("version") != 1
or payload.get("protocol") != "recursive-cascade-micro-state-v1"
):
raise ValueError("recursive cascade snapshot identity differs")
tokens = payload.get("tokens")
if not isinstance(tokens, list) or len(tokens) < int(payload["base_size"]):
raise ValueError("recursive cascade snapshot tokens are invalid")
model_state = payload.get("model")
if not isinstance(model_state, dict):
raise ValueError("recursive cascade snapshot model is invalid")
base_rows = torch.stack(
[model_state[f"embeddings.{index}"] for index in range(int(payload["base_size"]))]
)
restored = cls(base_rows, max_span_length=int(payload["max_span_length"]))
restored_tokens: list[CascadeToken] = []
for index, row in enumerate(tokens):
token = CascadeToken(
int(row["token_id"]),
tuple(int(value) for value in row["span"]),
int(row["generation"]),
None if row["left_parent"] is None else int(row["left_parent"]),
None if row["right_parent"] is None else int(row["right_parent"]),
)
if token.token_id != index:
raise ValueError("recursive cascade token ids are not contiguous")
restored_tokens.append(token)
for index in range(restored.base_size, len(restored_tokens)):
restored.embeddings.append(
nn.Parameter(model_state[f"embeddings.{index}"].detach().clone())
)
restored._tokens = tuple(restored_tokens)
restored.load_state_dict(model_state, strict=True)
optimizer = torch.optim.AdamW(restored.parameters(), **optimizer_kwargs)
optimizer.load_state_dict(payload["optimizer"])
return restored, optimizer
def recursive_pair_key(left: int, right: int) -> int:
"""Return a collision-free Cantor identity for two non-negative token ids."""
left, right = int(left), int(right)
if left < 0 or right < 0:
raise ValueError("recursive pair parents must be non-negative")
total = left + right
key = total * (total + 1) // 2 + right
if key > torch.iinfo(torch.int64).max:
raise OverflowError("recursive pair identity exceeds int64")
return key
def decode_recursive_pair_key(key: int) -> tuple[int, int]:
"""Invert :func:`recursive_pair_key` exactly using integer arithmetic."""
key = int(key)
if key < 0:
raise ValueError("recursive pair key must be non-negative")
diagonal = (math.isqrt(8 * key + 1) - 1) // 2
diagonal_start = diagonal * (diagonal + 1) // 2
right = key - diagonal_start
left = diagonal - right
if recursive_pair_key(left, right) != key:
raise ValueError("recursive pair key is not canonical")
return left, right
def _recursive_pair_keys(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
if left.dtype != torch.long or right.dtype != torch.long:
raise ValueError("recursive token ids must use torch.long")
if bool((left < 0).any()) or bool((right < 0).any()):
raise ValueError("recursive token ids must be non-negative")
total = left + right
# Cantor pairing is exact while the triangular term fits signed int64.
if total.numel() and int(total.max()) > 3_037_000_498:
raise OverflowError("recursive pair identity exceeds int64")
return total * (total + 1) // 2 + right
def _proposal_hash(
left: torch.Tensor, right: torch.Tensor, buckets: int
) -> torch.Tensor:
"""Preserve the Q15 proposal-row mapping at the exact Q32 launch."""
return ((left * 2_654_435_761 + right * 40_503).abs()) % int(buckets)
class ExactRecursiveDiscovery:
"""Unbounded exact pair/language statistics for one finite half-window."""
def __init__(self, *, n_languages: int) -> None:
if n_languages <= 0:
raise ValueError("recursive discovery needs a positive language count")
self.n_languages = int(n_languages)
self.records: dict[int, list[float | int]] = {}
self.observed_occurrences = 0
self.transferred_records = 0
@torch.no_grad()
def observe(
self,
pair_key: torch.Tensor,
language: torch.Tensor,
probability: torch.Tensor,
gradient: torch.Tensor,
) -> None:
if not pair_key.numel():
return
utility = -probability.detach().float() * gradient.detach().float()
finite = (
torch.isfinite(utility)
& torch.isfinite(probability.detach().float())
& (language >= 0)
& (language < self.n_languages)
)
if not bool(finite.any()):
return
pair_key = pair_key.detach()[finite].to(torch.long)
language = language.detach()[finite].to(torch.long)
probability = probability.detach()[finite].float()
utility = utility[finite]
self.observed_occurrences += int(pair_key.numel())
composite = pair_key * self.n_languages + language
unique, inverse = torch.unique(composite, return_inverse=True)
utility_sum = torch.zeros(unique.numel(), device=utility.device)
probability_sum = torch.zeros(unique.numel(), device=utility.device)
support = torch.zeros(unique.numel(), dtype=torch.long, device=utility.device)
positive_utility_support = torch.zeros(
unique.numel(), dtype=torch.long, device=utility.device
)
utility_sum.index_add_(0, inverse, utility)
probability_sum.index_add_(0, inverse, probability)
support.index_add_(0, inverse, torch.ones_like(inverse))
positive_utility_support.index_add_(
0, inverse, (utility > 0).to(torch.long)
)
rows = zip(
unique.cpu().tolist(),
utility_sum.cpu().tolist(),
probability_sum.cpu().tolist(),
support.cpu().tolist(),
positive_utility_support.cpu().tolist(),
strict=True,
)
for composite_key, value, probability_value, count, positive_count in rows:
key = int(composite_key)
row = self.records.setdefault(key, [0.0, 0.0, 0, 0])
row[0] = float(row[0]) + float(value)
row[1] = float(row[1]) + float(probability_value)
row[2] = int(row[2]) + int(count)
row[3] = int(row[3]) + int(positive_count)
self.transferred_records += 1
def snapshot_records(self) -> tuple[dict[str, float | int], ...]:
result = []
for composite, row in sorted(self.records.items()):
pair_key, language = divmod(composite, self.n_languages)
utility, probability_sum, support, positive_utility_support = (
float(row[0]), float(row[1]), int(row[2]), int(row[3])
)
result.append(
{
"composite_key": composite,
"pair_key": pair_key,
"language_index": language,
"utility": utility,
"probability_sum": probability_sum,
"captured_support": support,
"positive_utility_support": positive_utility_support,
"mean_probability": probability_sum / support,
}
)
return tuple(result)
class DenseRecursiveCascade(nn.Module):
"""Batched recursive tokenizer plus one dense trainable learned-token table.
Accepted rules are applied once per structural generation. Consequently a
token born at boundary ``g`` can be emitted and used by discovery during the
following interval, but a transition cannot recursively consume rows it is
creating itself. This makes every structural update an atomic DAG layer.
"""
def __init__(
self,
*,
base_size: int,
dim: int,
n_languages: int,
residual_buckets: int,
max_span_length: int = 32,
restored_state: dict[str, torch.Tensor] | None = None,
) -> None:
super().__init__()
if base_size <= 0 or dim <= 0 or n_languages <= 0:
raise ValueError("recursive cascade dimensions must be positive")
if residual_buckets <= 0 or max_span_length < 2:
raise ValueError("recursive residual capacity/span limit is invalid")
self.base_size = int(base_size)
self.dim = int(dim)
self.n_languages = int(n_languages)
self.residual_buckets = int(residual_buckets)
self.max_span_length = int(max_span_length)
restored = dict(restored_state or {})
metadata_names = {
"rule_left", "rule_right", "rule_generation", "span_offsets",
"span_values", "rule_active", "learned",
}
legacy_metadata_names = metadata_names - {"rule_active"}
if restored and set(restored) not in (metadata_names, legacy_metadata_names):
raise ValueError("recursive cascade restored state has an invalid schema")
learned = restored.get("learned", torch.empty((0, self.dim)))
if learned.ndim != 2 or learned.shape[1] != self.dim:
raise ValueError("recursive learned table has an invalid shape")
self.learned = nn.Parameter(learned.detach().clone())
count = int(learned.shape[0])
defaults = {
"rule_left": torch.empty(0, dtype=torch.long),
"rule_right": torch.empty(0, dtype=torch.long),
"rule_generation": torch.empty(0, dtype=torch.long),
"span_offsets": torch.zeros(1, dtype=torch.long),
"span_values": torch.empty(0, dtype=torch.long),
}
for name, default in defaults.items():
value = restored.get(name, default).detach().clone().to(torch.long)
self.register_buffer(name, value, persistent=True)
active = restored.get(
"rule_active", torch.ones(count, dtype=torch.bool)
).detach().clone().to(torch.bool)
self.register_buffer("rule_active", active, persistent=True)
if not (
self.rule_left.numel() == count
and self.rule_right.numel() == count
and self.rule_generation.numel() == count
and self.rule_active.numel() == count
and self.span_offsets.numel() == count + 1
and int(self.span_offsets[0]) == 0
and int(self.span_offsets[-1]) == self.span_values.numel()
):
raise ValueError("recursive rule metadata does not align with learned rows")
if count and (
bool((self.rule_generation <= 0).any())
or bool((self.rule_generation[1:] < self.rule_generation[:-1]).any())
or int(self.rule_left.max()) >= self.base_size + count
or int(self.rule_right.max()) >= self.base_size + count
):
raise ValueError("recursive rule DAG metadata is invalid")
self._validate_metadata()
self.proposal_merged = nn.Parameter(torch.zeros(self.residual_buckets, self.dim))
self.proposal_score = nn.Parameter(torch.zeros(self.residual_buckets))
self.language_bias = nn.Parameter(torch.full((self.n_languages,), -2.0))
self.discovery: ExactRecursiveDiscovery | None = None
self.usage_counts: torch.Tensor | None = None
@property
def learned_count(self) -> int:
return int(self.learned.shape[0])
@property
def generation(self) -> int:
return int(self.rule_generation[-1]) if self.rule_generation.numel() else 0
@property
def active_count(self) -> int:
return int(self.rule_active.sum())
def configure_discovery(self) -> None:
self.discovery = ExactRecursiveDiscovery(n_languages=self.n_languages)
self.usage_counts = torch.zeros(
self.learned_count, dtype=torch.long, device=self.learned.device
)
def snapshot_usage_token_ids(self) -> tuple[int, ...]:
if self.usage_counts is None:
raise RuntimeError("recursive usage collection is not configured")
rows = torch.nonzero(self.usage_counts > 0, as_tuple=False).squeeze(1)
return tuple(
self.base_size + int(row) for row in rows.detach().cpu().tolist()
)
def active_pair_keys(self) -> tuple[int, ...]:
left = self.rule_left.detach().cpu().tolist()
right = self.rule_right.detach().cpu().tolist()
active = self.rule_active.detach().cpu().tolist()
return tuple(
recursive_pair_key(int(left[row]), int(right[row]))
for row, enabled in enumerate(active)
if enabled
)
def token_span(self, token_id: int) -> tuple[int, ...]:
token_id = int(token_id)
if 0 <= token_id < self.base_size:
return (token_id,)
row = token_id - self.base_size
if row < 0 or row >= self.learned_count:
raise ValueError("recursive token id is outside the vocabulary")
start, stop = int(self.span_offsets[row]), int(self.span_offsets[row + 1])
return tuple(int(value) for value in self.span_values[start:stop].tolist())
def _validate_metadata(self) -> None:
seen_pairs: set[tuple[int, int]] = set()
seen_spans: set[tuple[int, ...]] = set()
left_values = self.rule_left.detach().cpu().tolist()
right_values = self.rule_right.detach().cpu().tolist()
generation_values = self.rule_generation.detach().cpu().tolist()
active_values = self.rule_active.detach().cpu().tolist()
offsets = self.span_offsets.detach().cpu().tolist()
span_values = self.span_values.detach().cpu().tolist()
row_spans = [
tuple(int(value) for value in span_values[offsets[row] : offsets[row + 1]])
for row in range(self.learned_count)
]
generations = set(int(value) for value in generation_values)
if generations and generations != set(range(1, max(generations) + 1)):
raise ValueError("recursive structural generations are not contiguous")
for row in range(self.learned_count):
token_id = self.base_size + row
left = int(left_values[row])
right = int(right_values[row])
generation = int(generation_values[row])
if left >= token_id or right >= token_id:
raise ValueError("recursive rule references a non-earlier token")
left_generation = (
0 if left < self.base_size else int(generation_values[left - self.base_size])
)
right_generation = (
0 if right < self.base_size else int(generation_values[right - self.base_size])
)
if left_generation >= generation or right_generation >= generation:
raise ValueError("recursive rule parent is not from an earlier generation")
pair = (left, right)
span = row_spans[row]
expected = (
((left,) if left < self.base_size else row_spans[left - self.base_size])
+ ((right,) if right < self.base_size else row_spans[right - self.base_size])
)
if pair in seen_pairs or span in seen_spans:
raise ValueError("recursive rule or canonical span is duplicated")
if span != expected or not span or len(span) > self.max_span_length:
raise ValueError("recursive canonical span metadata is invalid")
if min(span) < 0 or max(span) >= self.base_size:
raise ValueError("recursive canonical span contains a non-base id")
seen_pairs.add(pair)
seen_spans.add(span)
for row in range(self.learned_count):
if not bool(active_values[row]):
continue
for parent in (int(left_values[row]), int(right_values[row])):
if parent >= self.base_size and not bool(
active_values[parent - self.base_size]
):
raise ValueError("active recursive rule has an inactive parent")
def metadata_sha256(self) -> str:
payload = {
name: getattr(self, name).detach().cpu().tolist()
for name in (
"rule_left", "rule_right", "rule_generation", "span_offsets",
"span_values", "rule_active",
)
}
return hashlib.sha256(
json.dumps(payload, separators=(",", ":"), sort_keys=True).encode("utf-8")
).hexdigest()
def _lookup(
self, token_ids: torch.Tensor, base_embedding: torch.Tensor
) -> torch.Tensor:
if base_embedding.ndim != 2 or tuple(base_embedding.shape) != (
self.base_size, self.dim
):
raise ValueError("base embedding does not match recursive cascade")
if token_ids.numel() and (
int(token_ids.min()) < 0
or int(token_ids.max()) >= self.base_size + self.learned_count
):
raise ValueError("recursive token id is outside the vocabulary")
result = torch.empty(
(token_ids.numel(), self.dim),
dtype=base_embedding.dtype,
device=base_embedding.device,
)
base = token_ids < self.base_size
if bool(base.any()):
result[base] = base_embedding[token_ids[base]]
if bool((~base).any()):
result[~base] = self.learned[token_ids[~base] - self.base_size].to(
dtype=base_embedding.dtype
)
return result
@staticmethod
def _nonoverlapping_leftmost(candidate: torch.Tensor) -> torch.Tensor:
if candidate.dtype != torch.bool or candidate.ndim != 1:
raise ValueError("recursive merge candidates must be a boolean vector")
if not candidate.numel():
return candidate
positions = torch.arange(candidate.numel(), device=candidate.device)
previous_false = torch.cat(
(torch.ones(1, dtype=torch.bool, device=candidate.device), ~candidate[:-1])
)
starts = candidate & previous_false
start_positions = torch.where(starts, positions, torch.full_like(positions, -1))
last_start = torch.cummax(start_positions, dim=0).values
return candidate & ((positions - last_start).remainder(2) == 0)
def retokenize(
self, content: torch.Tensor, segment: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply the complete rule DAG to one packed batch on its current device."""
if content.dtype != torch.long or segment.dtype != torch.long:
raise ValueError("recursive packed tokens and segments must use torch.long")
if content.ndim != 1 or segment.shape != content.shape:
raise ValueError("recursive packed tokens and segments must align")
if content.numel() and (
int(content.min()) < 0 or int(content.max()) >= self.base_size
):
raise ValueError("recursive tokenizer input must contain base ids only")
emitted, emitted_segment = content, segment
for generation in range(1, self.generation + 1):
if emitted.numel() < 2:
break
rows = torch.nonzero(
(self.rule_generation == generation) & self.rule_active,
as_tuple=False,
).squeeze(1)
if not rows.numel():
continue
rule_keys = _recursive_pair_keys(
self.rule_left[rows], self.rule_right[rows]
)
order = torch.argsort(rule_keys, stable=True)
rule_keys = rule_keys[order]
rule_ids = rows[order] + self.base_size
inside = emitted_segment[:-1] == emitted_segment[1:]
adjacent = _recursive_pair_keys(emitted[:-1], emitted[1:])
positions = torch.searchsorted(rule_keys, adjacent)
safe = positions.clamp(max=rule_keys.numel() - 1)
found = inside & (positions < rule_keys.numel()) & (
rule_keys[safe] == adjacent
)
chosen = self._nonoverlapping_leftmost(found)
if not bool(chosen.any()):
continue
chosen_positions = torch.nonzero(chosen, as_tuple=False).squeeze(1)
replacement = emitted.clone()
replacement[chosen_positions] = rule_ids[safe[chosen_positions]]
keep = torch.ones_like(emitted, dtype=torch.bool)
keep[chosen_positions + 1] = False
emitted = replacement[keep]
emitted_segment = emitted_segment[keep]
return emitted, emitted_segment
def pool(
self,
base_embedding: torch.Tensor,
content: torch.Tensor,
segment: torch.Tensor,
language: torch.Tensor,
n_sentences: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Retokenize, train emitted rows, and observe next-generation pairs."""
emitted, emitted_segment = self.retokenize(content, segment)
if self.training and self.usage_counts is not None and emitted.numel():
learned_rows = emitted[emitted >= self.base_size] - self.base_size
if learned_rows.numel():
self.usage_counts.index_add_(
0,
learned_rows,
torch.ones_like(learned_rows, dtype=self.usage_counts.dtype),
)
vectors = self._lookup(emitted, base_embedding)
weight = torch.ones(emitted.numel(), dtype=vectors.dtype, device=vectors.device)
numerator = torch.zeros(
n_sentences, self.dim, dtype=vectors.dtype, device=vectors.device
)
denominator = torch.zeros(n_sentences, dtype=vectors.dtype, device=vectors.device)
denominator.index_add_(0, emitted_segment, weight)
if emitted.numel() > 1:
inside = emitted_segment[:-1] == emitted_segment[1:]
if bool(inside.any()):
left, right = emitted[:-1][inside], emitted[1:][inside]
pair_segment = emitted_segment[:-1][inside]
pair_language = language[pair_segment]
pair_key = _recursive_pair_keys(left, right)
bucket = _proposal_hash(left, right, self.residual_buckets)
probability = torch.sigmoid(
self.proposal_score[bucket] + self.language_bias[pair_language]
)
if (
self.discovery is not None
and self.training
and probability.requires_grad
):
saved_key = pair_key.detach()
saved_language = pair_language.detach()
saved_probability = probability.detach()
def observe(gradient: torch.Tensor) -> None:
if self.discovery is not None:
self.discovery.observe(
saved_key, saved_language, saved_probability, gradient
)
probability.register_hook(observe)
positions = torch.nonzero(inside, as_tuple=False).squeeze(1)
half = 0.5 * probability
weight = weight.index_add(0, positions, -half)
weight = weight.index_add(0, positions + 1, -half)
_weighted_index_add_(
numerator,
pair_segment,
self.proposal_merged[bucket],
probability,
)
denominator.index_add_(0, pair_segment, -probability)
_weighted_index_add_(numerator, emitted_segment, vectors, weight)
return (
numerator / denominator.clamp_min(1e-3).unsqueeze(1),
emitted,
emitted_segment,
)
@staticmethod
def _optimizer_row(
optimizer: torch.optim.Optimizer,
parameter: nn.Parameter,
row: int,
) -> dict[Any, Any]:
state = optimizer.state.get(parameter, {})
if not state:
raise ValueError("recursive promotion requires initialized parent Adam state")
result: dict[Any, Any] = {}
for name, value in state.items():
if not torch.is_tensor(value):
result[name] = value
elif value.ndim == 0:
result[name] = value.detach().clone()
elif tuple(value.shape) == tuple(parameter.shape):
result[name] = value[row].detach().clone()
else:
raise ValueError(f"unknown recursive optimizer state shape: {name}")
return result
@staticmethod
def _average_states(left: dict[Any, Any], right: dict[Any, Any]) -> dict[Any, Any]:
if set(left) != set(right):
raise ValueError("recursive parent Adam states differ")
result: dict[Any, Any] = {}
for name in left:
a, b = left[name], right[name]
if torch.is_tensor(a) != torch.is_tensor(b):
raise ValueError(f"recursive parent Adam state type differs: {name}")
if not torch.is_tensor(a):
if a != b:
raise ValueError(f"recursive parent Adam scalar differs: {name}")
result[name] = a
elif a.ndim == 0:
if not torch.equal(a, b):
raise ValueError(f"recursive parent Adam step differs: {name}")
result[name] = a.detach().clone()
else:
if a.shape != b.shape:
raise ValueError(f"recursive parent Adam row differs: {name}")
result[name] = (0.5 * (a + b)).detach().clone()
return result
def _parent_state(
self,
optimizer: torch.optim.Optimizer,
base_embedding: nn.Parameter,
token_id: int,
) -> dict[Any, Any]:
if token_id < self.base_size:
return self._optimizer_row(optimizer, base_embedding, token_id)
return self._optimizer_row(
optimizer, self.learned, token_id - self.base_size
)
def promotable_pair_keys(
self, pair_keys: Iterable[int]
) -> tuple[tuple[int, ...], dict[str, int]]:
existing_pairs = {
recursive_pair_key(int(left), int(right))
for left, right in zip(
self.rule_left.detach().cpu().tolist(),
self.rule_right.detach().cpu().tolist(),
strict=True,
)
}
existing_spans = {
self.token_span(token_id)
for token_id in range(self.base_size, self.base_size + self.learned_count)
}
accepted: list[int] = []
rejections = {
"already_rule": 0,
"unknown_parent": 0,
"duplicate_span": 0,
"span_too_long": 0,
}
total = self.base_size + self.learned_count
for key in sorted({int(value) for value in pair_keys}):
left, right = decode_recursive_pair_key(key)
if key in existing_pairs:
rejections["already_rule"] += 1
continue
if left >= total or right >= total:
rejections["unknown_parent"] += 1
continue
span = self.token_span(left) + self.token_span(right)
if len(span) > self.max_span_length:
rejections["span_too_long"] += 1
continue
if span in existing_spans:
rejections["duplicate_span"] += 1
continue
accepted.append(key)
existing_spans.add(span)
return tuple(accepted), rejections
def deactivate_unobserved(
self, observed_token_ids: Iterable[int]
) -> dict[str, Any]:
"""Deactivate the exact active sub-DAG absent from both audit windows.
Every active ancestor of an observed learned token is retained. All
remaining active rules can be disabled without changing either observed
token stream: they were neither emitted nor needed to emit a descendant.
Physical rows and Adam state remain stable for cheap future reactivation.
"""
observed = sorted({int(value) for value in observed_token_ids})
if any(
token_id < self.base_size
or token_id >= self.base_size + self.learned_count
for token_id in observed
):
raise ValueError("recursive usage contains an unknown learned token")
active = self.rule_active.detach().cpu().tolist()
left = self.rule_left.detach().cpu().tolist()
right = self.rule_right.detach().cpu().tolist()
keep = [False] * self.learned_count
if observed:
rows = [token_id - self.base_size for token_id in observed]
if not all(active[row] for row in rows):
raise ValueError("recursive usage contains an inactive token")
for row in rows:
keep[row] = True
# Child ids are always larger than learned parent ids, so one reverse
# pass closes the retained set over every active learned ancestor.
for row in range(self.learned_count - 1, -1, -1):
if not keep[row]:
continue
for parent in (int(left[row]), int(right[row])):
if parent >= self.base_size:
keep[parent - self.base_size] = True
rows = [
row for row, enabled in enumerate(active) if enabled and not keep[row]
]
if rows:
device_rows = torch.tensor(
rows, dtype=torch.long, device=self.rule_active.device
)
self.rule_active[device_rows] = False
token_ids = [self.base_size + row for row in rows]
return {
"deactivated_rows": len(token_ids),
"active_rows": self.active_count,
"token_ids_sha256": hashlib.sha256(
json.dumps(token_ids, separators=(",", ":")).encode("utf-8")
).hexdigest(),
}
def reactivate_pair_keys(self, pair_keys: Iterable[int]) -> dict[str, Any]:
"""Reactivate inactive exact rules whose complete parent path is active."""
left_values = self.rule_left.detach().cpu().tolist()
right_values = self.rule_right.detach().cpu().tolist()
by_key = {
recursive_pair_key(int(left), int(right)): row
for row, (left, right) in enumerate(
zip(left_values, right_values, strict=True)
)
}
reactivated_keys: list[int] = []
reactivated_ids: list[int] = []
rejected_inactive_parent = 0
active = self.rule_active.detach().cpu().tolist()
for key in sorted({int(value) for value in pair_keys}):
row = by_key.get(key)
if row is None or bool(active[row]):
continue
parents = (int(left_values[row]), int(right_values[row]))
if any(
parent >= self.base_size
and not bool(active[parent - self.base_size])
for parent in parents
):
rejected_inactive_parent += 1
continue
active[row] = True
reactivated_keys.append(key)
reactivated_ids.append(self.base_size + row)
if reactivated_ids:
rows = torch.tensor(
[token_id - self.base_size for token_id in reactivated_ids],
dtype=torch.long,
device=self.rule_active.device,
)
self.rule_active[rows] = True
return {
"reactivated_rows": len(reactivated_ids),
"active_rows": self.active_count,
"pair_keys": reactivated_keys,
"token_ids": reactivated_ids,
"rejected_inactive_parent": rejected_inactive_parent,
}
def grow(
self,
pair_keys: Iterable[int],
*,
base_embedding: nn.Parameter,
optimizer: torch.optim.Optimizer,
initialization: str = "parent_mean",
) -> dict[str, Any]:
"""Atomically append one dense generation and migrate live Adam state."""
if initialization not in {"parent_mean", "parent_sum"}:
raise ValueError("unknown recursive row initialization")
accepted, rejections = self.promotable_pair_keys(pair_keys)
if not accepted:
return {
"generation": self.generation,
"added_rows": 0,
"learned_rows": self.learned_count,
"active_rows": self.active_count,
"rejections": rejections,
}
pairs = tuple(decode_recursive_pair_key(key) for key in accepted)
old_parameter = self.learned
old_count = self.learned_count
new_values = []
new_states = []
spans = []
for left, right in pairs:
parent_ids = torch.tensor(
[left, right], dtype=torch.long, device=base_embedding.device
)
parent_values = self._lookup(parent_ids, base_embedding)
new_values.append(
parent_values.sum(dim=0)
if initialization == "parent_sum"
else parent_values.mean(dim=0)
)
new_states.append(
self._average_states(
self._parent_state(optimizer, base_embedding, left),
self._parent_state(optimizer, base_embedding, right),
)
)
spans.append(self.token_span(left) + self.token_span(right))
new_parameter = nn.Parameter(
torch.cat((old_parameter.detach(), torch.stack(new_values)), dim=0),
requires_grad=True if not old_count else old_parameter.requires_grad,
)
occurrences = sum(
candidate is old_parameter
for group in optimizer.param_groups
for candidate in group["params"]
)
if occurrences not in ({0, 1} if not old_count else {1}):
raise ValueError("recursive learned table must occur once in the optimizer")
old_state = optimizer.state.get(old_parameter, {})
migrated: dict[Any, Any] = {}
for name in new_states[0]:
additions = [state[name] for state in new_states]
first = additions[0]
if not torch.is_tensor(first):
if any(value != first for value in additions[1:]):
raise ValueError(f"recursive added Adam scalar differs: {name}")
migrated[name] = first
elif first.ndim == 0:
if any(not torch.equal(value, first) for value in additions[1:]):
raise ValueError(f"recursive added Adam step differs: {name}")
if old_state and not torch.equal(old_state[name], first):
raise ValueError(f"recursive existing/added Adam step differs: {name}")
migrated[name] = first.detach().clone()
else:
added = torch.stack(additions)
if old_count:
old_value = old_state.get(name)
if (
not torch.is_tensor(old_value)
or tuple(old_value.shape) != tuple(old_parameter.shape)
):
raise ValueError(f"recursive existing Adam rows differ: {name}")
added = added.to(device=old_value.device, dtype=old_value.dtype)
migrated[name] = torch.cat((old_value.detach(), added), dim=0)
else:
migrated[name] = added
for group in optimizer.param_groups:
replaced: list[nn.Parameter] = []
for parameter in group["params"]:
replaced.append(
new_parameter if parameter is old_parameter else parameter
)
if occurrences == 0 and parameter is base_embedding:
replaced.append(new_parameter)
group["params"] = replaced
optimizer.state.pop(old_parameter, None)
optimizer.state[new_parameter] = migrated
self.learned = new_parameter
device = self.rule_left.device
self.rule_left = torch.cat(
(self.rule_left, torch.tensor([left for left, _ in pairs], device=device))
)
self.rule_right = torch.cat(
(self.rule_right, torch.tensor([right for _, right in pairs], device=device))
)
new_generation = self.generation + 1
self.rule_generation = torch.cat(
(
self.rule_generation,
torch.full((len(pairs),), new_generation, dtype=torch.long, device=device),
)
)
self.rule_active = torch.cat(
(
self.rule_active,
torch.ones(len(pairs), dtype=torch.bool, device=device),
)
)
span_values = [value for span in spans for value in span]
lengths = torch.tensor([len(span) for span in spans], dtype=torch.long, device=device)
appended_offsets = self.span_offsets[-1] + lengths.cumsum(0)
self.span_offsets = torch.cat((self.span_offsets, appended_offsets))
self.span_values = torch.cat(
(self.span_values, torch.tensor(span_values, dtype=torch.long, device=device))
)
self.discovery = None
return {
"generation": new_generation,
"added_rows": len(pairs),
"learned_rows": self.learned_count,
"active_rows": self.active_count,
"token_ids": list(range(self.base_size + old_count, self.base_size + self.learned_count)),
"pairs": [[left, right] for left, right in pairs],
"rejections": rejections,
"initialization": initialization,
"metadata_sha256": self.metadata_sha256(),
}