Instructions to use BorisTM/loss-guided-static-multi with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use BorisTM/loss-guided-static-multi with sentence-transformers:
from sentence_transformers import SentenceTransformer model = SentenceTransformer("BorisTM/loss-guided-static-multi") sentences = [ "The weather is lovely today.", "It's so sunny outside!", "He drove to the stadium." ] embeddings = model.encode(sentences) similarities = model.similarity(embeddings, embeddings) print(similarities.shape) # [3, 3] - Notebooks
- Google Colab
- Kaggle
Download recursive_cascade.py from BorisTM/loss-guided-static-multi: direct link, hf CLI and curl.
- Browser
- Download file 47.3 kB
-
https://huggingface.co/BorisTM/loss-guided-static-multi/resolve/main/recursive_cascade.py
- Command line
-
hf download hf://BorisTM/loss-guided-static-multi/recursive_cascade.py
-
curl -L -o recursive_cascade.py https://huggingface.co/BorisTM/loss-guided-static-multi/resolve/main/recursive_cascade.py
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_ | |
| 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) | |
| ) | |
| 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) | |
| 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 | |
| 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()), | |
| } | |
| 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 | |
| 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 | |
| def learned_count(self) -> int: | |
| return int(self.learned.shape[0]) | |
| def generation(self) -> int: | |
| return int(self.rule_generation[-1]) if self.rule_generation.numel() else 0 | |
| 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 | |
| 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, | |
| ) | |
| 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 | |
| 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(), | |
| } | |