BorisTM/loss-guided-static-multi
08
1"""Learned collapsing of adjacent tokens into single units.2 3The earlier n-gram experiment *added* bigram rows next to the unigrams, so a4Georgian word cut into six pieces became six unigrams plus five bigrams — eleven5items instead of six. That raised fragmentation instead of lowering it, which is6why well-tokenised languages lost ground: an English sentence of four tokens7gained three bigram rows and half its pooled mass went to them.8 9Collapsing is a soft *replacement*: a learned merged row enters the numerator10while the neighbouring unigram rows lose mass. That changes the direction toward11units the tokenizer should have produced and targets tokens-per-word directly —12the quantity whose rank correlation with held-out Recall@1 is -0.607.13 14For adjacent pair i, a merge probability p_i gates a learned merged row M_i:15 16 numerator = sum_i (1 - (p_{i-1} + p_i)/2) * E[t_i] + sum_i p_i * M_i17 denominator = n - sum_i p_i18 19At p = 0 this is exactly the plain mean, so the control is recovered by20construction. Increasing p rotates the pooled direction from neighbouring21unigrams toward the learned merged row; overlapping candidates share the22reduction instead of needing a discrete matching. Under cosine scoring the23positive scalar denominator cancels, so the mechanism is the relative numerator24composition, not the smaller denominator.25 26The merge score carries a per-language bias, so a language whose words are27already single tokens can shut merging off globally rather than paying for a28mechanism it does not need — the failure mode of the ungated n-gram arm.29"""30 31from __future__ import annotations32 33from collections import defaultdict34from dataclasses import dataclass35import heapq36import math37import statistics38from typing import Any39 40import torch41from torch import nn42 43_MIX_A = 265443576144_MIX_B = 4050345_SKETCH_A = (6364136223846793005, 3202034522624059733, 3935559000370003845)46_SKETCH_B = (1442695040888963407, 2691343689449507681, 4768777513237032717)47_SKETCH_SIGN_A = (2862933555777941757, 7046029254386353131, 3935559000370003845)48_MAX_WEIGHTED_SOURCE_BYTES = 256 * 1024 * 102449 50 51def _weighted_index_add_(52 target: torch.Tensor,53 index: torch.Tensor,54 source: torch.Tensor,55 weight: torch.Tensor,56 *,57 max_chunk_bytes: int = _MAX_WEIGHTED_SOURCE_BYTES,58) -> torch.Tensor:59 """Index-add weighted rows without materialising an unbounded N x dim product."""60 if source.ndim != 2 or weight.ndim != 1 or source.shape[0] != weight.shape[0]:61 raise ValueError("weighted index-add needs aligned matrix rows and scalar weights")62 if index.ndim != 1 or index.shape[0] != source.shape[0]:63 raise ValueError("weighted index-add needs one destination per source row")64 if max_chunk_bytes <= 0:65 raise ValueError("weighted index-add chunk budget must be positive")66 bytes_per_row = max(1, source.shape[1] * source.element_size())67 rows_per_chunk = max(1, max_chunk_bytes // bytes_per_row)68 for start in range(0, source.shape[0], rows_per_chunk):69 stop = min(source.shape[0], start + rows_per_chunk)70 target.index_add_(71 0,72 index[start:stop],73 source[start:stop] * weight[start:stop].unsqueeze(1),74 )75 return target76 77 78@dataclass(frozen=True)79class PromotedVocabulary:80 """Deterministic result of bounded exact-pair promotion."""81 82 pair_keys: tuple[int, ...]83 records: tuple[dict[str, Any], ...]84 candidate_count: int85 reservation_memberships: tuple[tuple[int, int], ...]86 87 88class ContrastiveVocabularyDiscovery:89 """Bounded candidate table with signed CountSketch admission.90 91 The sketch receives ``-p * dL/dp`` after each backward pass. It transfers92 only the largest absolute pair/language aggregates from a forward. A fixed93 signed CountSketch retains cumulative admission evidence even for keys not94 currently in the exact table; the bounded exact table stores utility,95 probability and support only from admission onward.96 """97 98 admission_mode = "utility"99 100 def __init__(101 self,102 *,103 n_languages: int,104 candidate_capacity: int,105 top_per_forward: int,106 ) -> None:107 if n_languages <= 0:108 raise ValueError("n_languages must be positive")109 if candidate_capacity <= 0:110 raise ValueError("candidate_capacity must be positive")111 if top_per_forward <= 0:112 raise ValueError("top_per_forward must be positive")113 self.n_languages = int(n_languages)114 self.candidate_capacity = int(candidate_capacity)115 self.top_per_forward = int(top_per_forward)116 # composite pair/language key -> [signed utility, p sum, captured count]117 self.records: dict[int, list[float | int]] = {}118 self._sketch_width = max(64, self.candidate_capacity)119 self._utility_sketch = torch.zeros(120 (len(_SKETCH_A), self._sketch_width), dtype=torch.float64121 )122 # Lazy min-heap over retained rows. Updates append a new rank and leave123 # the previous entry stale; periodic rebuilding gives the metadata a124 # fixed bound without sorting the entire sketch on every forward.125 self._heap: list[tuple[float, int, int, int]] = []126 self.observed_occurrences = 0127 self.transferred_records = 0128 self.unique_records_before_transfer = 0129 self.dropped_by_top_per_forward = 0130 self.retained_updates = 0131 self.admitted_into_free_slots = 0132 self.capacity_replacements = 0133 self.capacity_rejections = 0134 135 @staticmethod136 def _rank(key: int, row: list[float | int]) -> tuple[float, int, int]:137 return float(row[0]), int(row[2]), -int(key)138 139 def _push_heap(self, key: int) -> None:140 heapq.heappush(self._heap, (*self._rank(key, self.records[key]), int(key)))141 142 def _discard_stale_heap_entries(self) -> None:143 while self._heap:144 utility, count, negative_key, key = self._heap[0]145 row = self.records.get(key)146 if row is not None and (utility, count, negative_key) == self._rank(key, row):147 return148 heapq.heappop(self._heap)149 150 def _rebuild_heap(self) -> None:151 self._heap = [(*self._rank(key, row), key) for key, row in self.records.items()]152 heapq.heapify(self._heap)153 154 def _update_utility_sketch(155 self, keys: torch.Tensor, values: torch.Tensor156 ) -> torch.Tensor:157 """Update fixed memory and return the median signed estimate per key."""158 159 keys = keys.to(device="cpu", dtype=torch.long)160 values = values.to(device="cpu", dtype=torch.float64)161 estimates = []162 for depth, (multiplier, offset, sign_multiplier) in enumerate(zip(163 _SKETCH_A, _SKETCH_B, _SKETCH_SIGN_A, strict=True164 )):165 bucket = torch.remainder(keys * multiplier + offset, self._sketch_width)166 sign = torch.where(167 torch.bitwise_and(keys * sign_multiplier + offset, 1) == 0,168 1.0,169 -1.0,170 ).to(torch.float64)171 self._utility_sketch[depth].index_add_(0, bucket, values * sign)172 estimates.append(self._utility_sketch[depth, bucket] * sign)173 return torch.stack(estimates).median(dim=0).values174 175 @torch.no_grad()176 def observe(177 self,178 pair_key: torch.Tensor,179 language: torch.Tensor,180 probability: torch.Tensor,181 gradient: torch.Tensor,182 ) -> None:183 if not pair_key.numel():184 return185 utility = -probability.detach().float() * gradient.detach().float()186 finite = torch.isfinite(utility)187 if not bool(finite.any()):188 return189 pair_key = pair_key.detach()[finite].to(torch.long)190 language = language.detach()[finite].to(torch.long)191 probability = probability.detach()[finite].float()192 utility = utility[finite]193 self.observed_occurrences += int(pair_key.numel())194 195 composite = pair_key * self.n_languages + language196 unique, inverse = torch.unique(composite, return_inverse=True)197 utility_sum = torch.zeros(unique.numel(), device=utility.device)198 probability_sum = torch.zeros(unique.numel(), device=utility.device)199 count = torch.zeros(unique.numel(), dtype=torch.long, device=utility.device)200 utility_sum.index_add_(0, inverse, utility)201 probability_sum.index_add_(0, inverse, probability)202 count.index_add_(0, inverse, torch.ones_like(inverse))203 204 self.unique_records_before_transfer += int(unique.numel())205 keep = min(self.top_per_forward, unique.numel())206 if keep < unique.numel():207 self.dropped_by_top_per_forward += int(unique.numel() - keep)208 chosen = torch.topk(utility_sum.abs(), keep, sorted=False).indices209 unique = unique[chosen]210 utility_sum = utility_sum[chosen]211 probability_sum = probability_sum[chosen]212 count = count[chosen]213 214 unique_cpu = unique.cpu()215 utility_cpu = utility_sum.cpu()216 admission_estimates = self._update_utility_sketch(unique_cpu, utility_cpu)217 keys = unique_cpu.tolist()218 utilities = utility_cpu.tolist()219 probabilities = probability_sum.cpu().tolist()220 counts = count.cpu().tolist()221 self.transferred_records += len(keys)222 incoming = {223 int(key): [float(value), float(p_sum), int(n)]224 for key, value, p_sum, n in zip(225 keys, utilities, probabilities, counts, strict=True226 )227 }228 admission = {229 int(key): float(estimate)230 for key, estimate in zip(keys, admission_estimates.tolist(), strict=True)231 }232 233 # First update retained rows, then consider unseen rows in descending234 # batch rank. This makes replacement deterministic and independent of235 # the device order returned by top-k.236 retained_keys = sorted(key for key in incoming if key in self.records)237 self.retained_updates += len(retained_keys)238 for key in retained_keys:239 value, p_sum, n = incoming.pop(key)240 row = self.records[key]241 row[0] = float(row[0]) + float(value)242 row[1] = float(row[1]) + float(p_sum)243 row[2] = int(row[2]) + int(n)244 self._push_heap(key)245 246 unseen = sorted(247 incoming.items(),248 key=lambda item: (admission[item[0]], int(item[1][2]), -item[0]),249 reverse=True,250 )251 for key, row in unseen:252 if len(self.records) < self.candidate_capacity:253 self.records[key] = row254 self._push_heap(key)255 self.admitted_into_free_slots += 1256 continue257 self._discard_stale_heap_entries()258 if not self._heap:259 self._rebuild_heap()260 admission_rank = (admission[key], int(row[2]), -key)261 if admission_rank <= self._heap[0][:3]:262 self.capacity_rejections += 1263 continue264 victim = heapq.heappop(self._heap)[3]265 del self.records[victim]266 self.records[key] = row267 self._push_heap(key)268 self.capacity_replacements += 1269 270 # The retained dictionary never exceeds candidate_capacity. Lazy heap271 # metadata is also bounded: after a call it has at most 2x as many272 # entries as the registered candidate table.273 if len(self._heap) > 2 * self.candidate_capacity:274 self._rebuild_heap()275 276 def _prune(self) -> None:277 ranked = sorted(278 self.records.items(),279 key=lambda item: self._rank(item[0], item[1]),280 reverse=True,281 )[: self.candidate_capacity]282 self.records = dict(ranked)283 self._rebuild_heap()284 285 def snapshot_records(self) -> tuple[dict[str, float | int], ...]:286 """Return the complete bounded pair/language table deterministically."""287 288 self._prune()289 rows = []290 for composite, row in sorted(self.records.items()):291 pair_key, language = divmod(int(composite), self.n_languages)292 utility, probability_sum, count = float(row[0]), float(row[1]), int(row[2])293 rows.append({294 "composite_key": int(composite),295 "pair_key": pair_key,296 "language_index": language,297 "utility": utility,298 "probability_sum": probability_sum,299 "captured_support": count,300 "mean_probability": probability_sum / count if count else 0.0,301 })302 return tuple(rows)303 304 def state_dict(self) -> dict[str, Any]:305 """Return the complete bounded discovery state for an exact resume."""306 307 self._prune()308 return {309 "version": 2,310 "n_languages": self.n_languages,311 "candidate_capacity": self.candidate_capacity,312 "top_per_forward": self.top_per_forward,313 "sketch_width": self._sketch_width,314 "records": {315 int(key): [float(row[0]), float(row[1]), int(row[2])]316 for key, row in sorted(self.records.items())317 },318 "utility_sketch": self._utility_sketch.clone(),319 "observed_occurrences": self.observed_occurrences,320 "transferred_records": self.transferred_records,321 "admission_pressure": self.admission_pressure(),322 }323 324 def load_state_dict(self, raw: dict[str, Any]) -> None:325 """Restore a state produced by :meth:`state_dict` after validation."""326 327 version = int(raw.get("version", -1))328 if version not in (1, 2):329 raise ValueError("unsupported discovery state version")330 if int(raw.get("n_languages", -1)) != self.n_languages:331 raise ValueError("discovery state language count does not match")332 if int(raw.get("candidate_capacity", -1)) != self.candidate_capacity:333 raise ValueError("discovery state candidate capacity does not match")334 if int(raw.get("top_per_forward", -1)) != self.top_per_forward:335 raise ValueError("discovery state top-per-forward does not match")336 if int(raw.get("sketch_width", -1)) != self._sketch_width:337 raise ValueError("discovery state sketch width does not match")338 sketch = raw.get("utility_sketch")339 if (340 not torch.is_tensor(sketch)341 or sketch.dtype != torch.float64342 or tuple(sketch.shape) != tuple(self._utility_sketch.shape)343 or not bool(torch.isfinite(sketch).all())344 ):345 raise ValueError("discovery state utility sketch is invalid")346 raw_records = raw.get("records")347 if not isinstance(raw_records, dict) or len(raw_records) > self.candidate_capacity:348 raise ValueError("discovery state records are invalid")349 records: dict[int, list[float | int]] = {}350 for raw_key, raw_row in raw_records.items():351 key = int(raw_key)352 if key < 0 or not isinstance(raw_row, (list, tuple)) or len(raw_row) != 3:353 raise ValueError("discovery state record is invalid")354 utility, probability, count = float(raw_row[0]), float(raw_row[1]), int(raw_row[2])355 if not (math.isfinite(utility) and math.isfinite(probability)) or count < 0:356 raise ValueError("discovery state record values are invalid")357 records[key] = [utility, probability, count]358 observed = int(raw.get("observed_occurrences", -1))359 transferred = int(raw.get("transferred_records", -1))360 if observed < 0 or transferred < 0:361 raise ValueError("discovery state counters are invalid")362 counter_names = tuple(self.admission_pressure())363 raw_pressure = raw.get("admission_pressure", {})364 if version == 2 and (365 not isinstance(raw_pressure, dict)366 or set(raw_pressure) != set(counter_names)367 ):368 raise ValueError("discovery state admission counters are invalid")369 pressure = {370 name: int(raw_pressure.get(name, 0)) for name in counter_names371 }372 if any(value < 0 for value in pressure.values()):373 raise ValueError("discovery state admission counters are invalid")374 375 self.records = records376 self._utility_sketch.copy_(sketch)377 self.observed_occurrences = observed378 self.transferred_records = transferred379 for name, value in pressure.items():380 setattr(self, name, value)381 self._rebuild_heap()382 383 def admission_pressure(self) -> dict[str, int]:384 return {385 "unique_records_before_transfer": self.unique_records_before_transfer,386 "dropped_by_top_per_forward": self.dropped_by_top_per_forward,387 "retained_updates": self.retained_updates,388 "admitted_into_free_slots": self.admitted_into_free_slots,389 "capacity_replacements": self.capacity_replacements,390 "capacity_rejections": self.capacity_rejections,391 }392 393 394 def promote(395 self,396 *,397 budget: int,398 min_support: int,399 per_language_quota: int,400 per_language_min_support: int,401 min_utility: float = 0.0,402 ) -> PromotedVocabulary:403 """Select utility-thresholded exact rows under a safety L0 ceiling."""404 405 if budget <= 0:406 raise ValueError("budget must be positive")407 if min_support <= 0 or per_language_min_support <= 0:408 raise ValueError("support thresholds must be positive")409 if per_language_quota < 0:410 raise ValueError("per_language_quota must be non-negative")411 if not math.isfinite(min_utility) or min_utility < 0:412 raise ValueError("min_utility must be finite and non-negative")413 self._prune()414 415 by_pair: dict[int, dict[str, Any]] = defaultdict(416 lambda: {"utility": 0.0, "p_sum": 0.0, "count": 0, "languages": {}}417 )418 for composite, row in self.records.items():419 pair_key, language = divmod(composite, self.n_languages)420 utility, p_sum, count = float(row[0]), float(row[1]), int(row[2])421 target = by_pair[pair_key]422 target["utility"] += utility423 target["p_sum"] += p_sum424 target["count"] += count425 target["languages"][language] = {426 "utility": utility,427 "p_sum": p_sum,428 "count": count,429 }430 431 selected: set[int] = set()432 reservation_languages: dict[int, set[int]] = defaultdict(set)433 if per_language_quota:434 candidates_by_language: list[list[tuple[float, int, int, int]]] = [435 [] for _ in range(self.n_languages)436 ]437 for pair_key, row in by_pair.items():438 # A language reservation protects multilingual coverage, but439 # it cannot override the registered global U(q) > 0440 # eligibility constraint. Build the per-language lists in441 # one pass over observed records; scanning every pair once per442 # marker is prohibitive for the 1,914-language run.443 if row["utility"] <= 0 or row["utility"] < min_utility:444 continue445 for language, language_row in row["languages"].items():446 if (language_row["utility"] > 0447 and language_row["count"] >= per_language_min_support):448 candidates_by_language[language].append((449 float(language_row["utility"]),450 int(language_row["count"]),451 -int(pair_key),452 int(pair_key),453 ))454 for language, candidates in enumerate(candidates_by_language):455 candidates.sort(reverse=True)456 for candidate in candidates[:per_language_quota]:457 pair_key = candidate[3]458 selected.add(pair_key)459 reservation_languages[pair_key].add(language)460 461 if len(selected) > budget:462 selected = set(sorted(463 selected,464 key=lambda key: (465 float(by_pair[key]["utility"]),466 int(by_pair[key]["count"]),467 -key,468 ),469 reverse=True,470 )[:budget])471 reservation_languages = {472 key: languages for key, languages in reservation_languages.items()473 if key in selected474 }475 476 global_candidates = [477 key for key, row in by_pair.items()478 if (479 row["utility"] > 0480 and row["utility"] >= min_utility481 and row["count"] >= min_support482 and key not in selected483 )484 ]485 global_candidates.sort(486 key=lambda key: (487 float(by_pair[key]["utility"]),488 int(by_pair[key]["count"]),489 -key,490 ),491 reverse=True,492 )493 selected.update(global_candidates[: max(0, budget - len(selected))])494 495 pair_keys = tuple(sorted(selected))496 rows = []497 for pair_key in pair_keys:498 row = by_pair[pair_key]499 reserved = sorted(reservation_languages.get(pair_key, ()))500 rows.append({501 "pair_key": pair_key,502 "utility": float(row["utility"]),503 "captured_support": int(row["count"]),504 "mean_probability": (505 float(row["p_sum"]) / int(row["count"]) if row["count"] else 0.0506 ),507 "language_count": len(row["languages"]),508 "reservation_language_count": len(reserved),509 "max_reservation_support": max(510 (int(row["languages"][language]["count"]) for language in reserved),511 default=0,512 ),513 "global_support_eligible": int(row["count"]) >= min_support,514 })515 return PromotedVocabulary(516 pair_keys=pair_keys,517 records=tuple(rows),518 candidate_count=len(by_pair),519 reservation_memberships=tuple(sorted(520 (pair_key, language)521 for pair_key, languages in reservation_languages.items()522 for language in languages523 )),524 )525 526 def audit_promotion(527 self,528 promoted: PromotedVocabulary,529 *,530 budget: int,531 min_support: int,532 per_language_quota: int,533 per_language_min_support: int,534 min_utility: float = 0.0,535 reservation_audit_quota: int | None = None,536 trainable_language_indices: tuple[int, ...] | None = None,537 ) -> dict[str, Any]:538 """Compare a quota selection with the same discovery under quota zero."""539 540 if reservation_audit_quota is None:541 reservation_audit_quota = per_language_quota542 if reservation_audit_quota < 0:543 raise ValueError("reservation_audit_quota must be non-negative")544 if trainable_language_indices is None:545 trainable_language_indices = tuple(range(self.n_languages))546 else:547 trainable_language_indices = tuple(548 int(index) for index in trainable_language_indices549 )550 if (551 not trainable_language_indices552 or len(set(trainable_language_indices)) != len(trainable_language_indices)553 or any(554 index < 0 or index >= self.n_languages555 for index in trainable_language_indices556 )557 ):558 raise ValueError("trainable language indices are invalid")559 trainable_set = set(trainable_language_indices)560 global_only = (561 promoted562 if per_language_quota == 0563 else self.promote(564 budget=budget,565 min_support=min_support,566 per_language_quota=0,567 per_language_min_support=per_language_min_support,568 min_utility=min_utility,569 )570 )571 reservation_selection = (572 promoted573 if per_language_quota == reservation_audit_quota574 else self.promote(575 budget=budget,576 min_support=min_support,577 per_language_quota=reservation_audit_quota,578 per_language_min_support=per_language_min_support,579 min_utility=min_utility,580 )581 )582 all_candidate_counts = [0] * self.n_languages583 aggregate_utility: dict[int, float] = defaultdict(float)584 for composite, row in self.records.items():585 pair_key, language = divmod(int(composite), self.n_languages)586 if language not in trainable_set:587 raise ValueError(588 "discovery contains a nontrainable language candidate"589 )590 all_candidate_counts[language] += 1591 aggregate_utility[pair_key] += float(row[0])592 all_eligible_counts = [0] * self.n_languages593 for composite, row in self.records.items():594 pair_key, language = divmod(int(composite), self.n_languages)595 if (596 aggregate_utility[pair_key] > 0597 and aggregate_utility[pair_key] >= min_utility598 and float(row[0]) > 0599 and int(row[2]) >= per_language_min_support600 ):601 all_eligible_counts[language] += 1602 candidate_counts = [603 all_candidate_counts[index] for index in trainable_language_indices604 ]605 eligible_counts = [606 all_eligible_counts[index] for index in trainable_language_indices607 ]608 all_reservation_counts = [0] * self.n_languages609 for _, language in reservation_selection.reservation_memberships:610 if language not in trainable_set:611 raise ValueError("reservation selected a nontrainable language")612 all_reservation_counts[language] += 1613 reservation_counts = [614 all_reservation_counts[index] for index in trainable_language_indices615 ]616 617 reservation_utility = math.fsum(618 float(record["utility"]) for record in reservation_selection.records619 )620 global_utility = math.fsum(621 float(record["utility"]) for record in global_only.records622 )623 utility_displaced = global_utility - reservation_utility624 membership_capacity = reservation_audit_quota * len(trainable_language_indices)625 reservation_keys = set(reservation_selection.pair_keys)626 global_keys = set(global_only.pair_keys)627 return {628 "deployed_selector": {629 "per_language_quota": per_language_quota,630 "matches_global_only": promoted.pair_keys == global_only.pair_keys,631 "reservation_audit_quota": reservation_audit_quota,632 },633 "candidate_table": {634 "capacity": self.candidate_capacity,635 "pair_language_records": len(self.records),636 "saturation": len(self.records) / self.candidate_capacity,637 "candidate_pairs": promoted.candidate_count,638 "initializer_languages_total": self.n_languages,639 "languages_total": len(trainable_language_indices),640 "trainable_language_indices": list(trainable_language_indices),641 "languages_with_candidates": sum(642 count > 0 for count in candidate_counts643 ),644 "records_per_language": candidate_counts,645 "reservation_eligible_records_per_language": eligible_counts,646 "min_records_per_language": min(candidate_counts),647 "median_records_per_language": statistics.median(candidate_counts),648 "max_records_per_language": max(candidate_counts),649 "admission_pressure": {650 **{651 "unique_records_before_transfer": (652 self.unique_records_before_transfer653 ),654 "dropped_by_top_per_forward": (655 self.dropped_by_top_per_forward656 ),657 "transferred_records": self.transferred_records,658 },659 "retained_updates": self.retained_updates,660 "admitted_into_free_slots": self.admitted_into_free_slots,661 "capacity_replacements": self.capacity_replacements,662 "capacity_rejections": self.capacity_rejections,663 },664 },665 "reservation": {666 "quota_per_language": reservation_audit_quota,667 "membership_capacity": membership_capacity,668 "memberships": len(reservation_selection.reservation_memberships),669 "membership_occupancy": (670 len(reservation_selection.reservation_memberships)671 / membership_capacity672 if membership_capacity673 else 0.0674 ),675 "unique_reserved_rows": len({676 pair_key677 for pair_key, _ in reservation_selection.reservation_memberships678 }),679 "languages_with_reservations": sum(680 count > 0 for count in reservation_counts681 ),682 "memberships_per_language": reservation_counts,683 },684 "global_only_counterfactual": {685 "global_selected_rows": len(global_only.pair_keys),686 "global_selected_utility": global_utility,687 "reservation_selected_rows": len(reservation_selection.pair_keys),688 "reservation_selected_utility": reservation_utility,689 "utility_displaced_by_reservations": utility_displaced,690 "relative_utility_displaced_by_reservations": (691 utility_displaced / global_utility if global_utility else None692 ),693 "reservation_only_rows": len(reservation_keys - global_keys),694 "global_only_rows": len(global_keys - reservation_keys),695 },696 }697 698 def select_external(699 self,700 pair_keys: tuple[int, ...] | list[int],701 *,702 minimum_support: int,703 ) -> PromotedVocabulary:704 """Validate and materialize a hash-pinned external promotion decision."""705 706 if minimum_support <= 0:707 raise ValueError("external-selection minimum support must be positive")708 keys = tuple(int(key) for key in pair_keys)709 if not keys or tuple(sorted(set(keys))) != keys:710 raise ValueError("external pair keys must be non-empty, sorted and unique")711 self._prune()712 by_pair: dict[int, dict[str, Any]] = defaultdict(713 lambda: {"utility": 0.0, "p_sum": 0.0, "count": 0, "languages": {}}714 )715 for composite, raw in self.records.items():716 pair_key, language = divmod(composite, self.n_languages)717 utility, p_sum, count = float(raw[0]), float(raw[1]), int(raw[2])718 row = by_pair[pair_key]719 row["utility"] += utility720 row["p_sum"] += p_sum721 row["count"] += count722 row["languages"][language] = {723 "utility": utility, "p_sum": p_sum, "count": count,724 }725 missing = [key for key in keys if key not in by_pair]726 if missing:727 raise ValueError(f"external pair keys absent from discovery: {missing[:8]}")728 ineligible = [729 key for key in keys730 if by_pair[key]["utility"] <= 0 or by_pair[key]["count"] < minimum_support731 ]732 if ineligible:733 raise ValueError(734 f"external pair keys fail utility/support eligibility: {ineligible[:8]}"735 )736 rows = []737 for pair_key in keys:738 row = by_pair[pair_key]739 rows.append({740 "pair_key": pair_key,741 "utility": float(row["utility"]),742 "captured_support": int(row["count"]),743 "mean_probability": (744 float(row["p_sum"]) / int(row["count"]) if row["count"] else 0.0745 ),746 "language_count": len(row["languages"]),747 "reservation_language_count": 0,748 "max_reservation_support": 0,749 "global_support_eligible": int(row["count"]) >= minimum_support,750 })751 return PromotedVocabulary(752 pair_keys=keys,753 records=tuple(rows),754 candidate_count=len(by_pair),755 reservation_memberships=(),756 )757 758 759class FrequencyVocabularyDiscovery:760 """Independent fixed-memory heavy hitters ranked only by occurrence count.761 762 Unlike :class:`ContrastiveVocabularyDiscovery`, this observer never reads763 the value or sign of ``-p * dL/dp``. A Count-Min sketch sees every pair764 occurrence, while a bounded exact key table retains the largest sketch765 estimates. The table is pair-level rather than pair/language-level so a766 globally frequent pair does not spend multiple candidate slots merely767 because it occurs in several languages.768 769 The sketch has the same depth and width rule as the utility collector. Its770 estimates can over-count because of sketch collisions; they cannot771 under-count. This is an independent fixed-memory frequency selector, not772 an exact unbounded corpus histogram.773 """774 775 admission_mode = "frequency"776 777 def __init__(778 self,779 *,780 n_languages: int,781 candidate_capacity: int,782 top_per_forward: int,783 ) -> None:784 if n_languages <= 0:785 raise ValueError("n_languages must be positive")786 if candidate_capacity <= 0:787 raise ValueError("candidate_capacity must be positive")788 if top_per_forward <= 0:789 raise ValueError("top_per_forward must be positive")790 self.n_languages = int(n_languages)791 self.candidate_capacity = int(candidate_capacity)792 self.top_per_forward = int(top_per_forward)793 # pair key -> [unused utility, unused probability sum, support estimate]794 self.records: dict[int, list[float | int]] = {}795 self._sketch_width = max(64, self.candidate_capacity)796 self._frequency_sketch = torch.zeros(797 (len(_SKETCH_A), self._sketch_width), dtype=torch.long798 )799 self._heap: list[tuple[int, int, int]] = []800 self.observed_occurrences = 0801 self.transferred_records = 0802 self.unique_records_before_transfer = 0803 self.dropped_by_top_per_forward = 0804 self.retained_updates = 0805 self.admitted_into_free_slots = 0806 self.capacity_replacements = 0807 self.capacity_rejections = 0808 809 @staticmethod810 def _rank(key: int, row: list[float | int]) -> tuple[int, int]:811 return int(row[2]), -int(key)812 813 def _push_heap(self, key: int) -> None:814 heapq.heappush(self._heap, (*self._rank(key, self.records[key]), int(key)))815 816 def _discard_stale_heap_entries(self) -> None:817 while self._heap:818 support, negative_key, key = self._heap[0]819 row = self.records.get(key)820 if row is not None and (support, negative_key) == self._rank(key, row):821 return822 heapq.heappop(self._heap)823 824 def _rebuild_heap(self) -> None:825 self._heap = [826 (*self._rank(key, row), key) for key, row in self.records.items()827 ]828 heapq.heapify(self._heap)829 830 def _update_frequency_sketch(831 self, keys: torch.Tensor, counts: torch.Tensor832 ) -> torch.Tensor:833 """Update Count-Min state for every key and return its current estimate."""834 835 keys = keys.to(dtype=torch.long)836 counts = counts.to(device=keys.device, dtype=torch.long)837 if self._frequency_sketch.device != keys.device:838 self._frequency_sketch = self._frequency_sketch.to(keys.device)839 estimates = []840 for depth, (multiplier, offset) in enumerate(841 zip(_SKETCH_A, _SKETCH_B, strict=True)842 ):843 bucket = torch.remainder(844 keys * multiplier + offset, self._sketch_width845 )846 self._frequency_sketch[depth].index_add_(0, bucket, counts)847 estimates.append(self._frequency_sketch[depth, bucket])848 return torch.stack(estimates).amin(dim=0)849 850 @staticmethod851 def _largest_with_deterministic_boundary(852 keys: torch.Tensor, estimates: torch.Tensor, keep: int853 ) -> torch.Tensor:854 """Keep the largest estimates, resolving the cutoff by smaller key."""855 856 if keep >= keys.numel():857 return torch.arange(keys.numel(), device=keys.device)858 boundary = torch.topk(estimates, keep, sorted=False).values.min()859 above = torch.nonzero(estimates > boundary, as_tuple=False).squeeze(1)860 ties = torch.nonzero(estimates == boundary, as_tuple=False).squeeze(1)861 remaining = keep - int(above.numel())862 if remaining <= 0:863 return above[:keep]864 tie_choice = torch.topk(865 -keys.index_select(0, ties), remaining, sorted=False866 ).indices867 return torch.cat((above, ties.index_select(0, tie_choice)))868 869 @torch.no_grad()870 def observe(871 self,872 pair_key: torch.Tensor,873 language: torch.Tensor,874 probability: torch.Tensor,875 gradient: torch.Tensor,876 ) -> None:877 del language, probability, gradient878 if not pair_key.numel():879 return880 pair_key = pair_key.detach().to(torch.long)881 self.observed_occurrences += int(pair_key.numel())882 unique, counts = torch.unique(pair_key, return_counts=True)883 estimates = self._update_frequency_sketch(unique, counts)884 885 self.unique_records_before_transfer += int(unique.numel())886 keep = min(self.top_per_forward, unique.numel())887 if keep < unique.numel():888 self.dropped_by_top_per_forward += int(unique.numel() - keep)889 chosen = self._largest_with_deterministic_boundary(890 unique, estimates, keep891 )892 unique = unique.index_select(0, chosen)893 estimates = estimates.index_select(0, chosen)894 895 keys = unique.cpu().tolist()896 supports = estimates.cpu().tolist()897 self.transferred_records += len(keys)898 incoming = {int(key): int(support) for key, support in zip(keys, supports)}899 900 retained_keys = sorted(key for key in incoming if key in self.records)901 self.retained_updates += len(retained_keys)902 for key in retained_keys:903 support = incoming.pop(key)904 self.records[key][2] = max(int(self.records[key][2]), support)905 self._push_heap(key)906 907 unseen = sorted(908 incoming.items(), key=lambda item: (item[1], -item[0]), reverse=True909 )910 for key, support in unseen:911 row: list[float | int] = [0.0, 0.0, int(support)]912 if len(self.records) < self.candidate_capacity:913 self.records[key] = row914 self._push_heap(key)915 self.admitted_into_free_slots += 1916 continue917 self._discard_stale_heap_entries()918 if not self._heap:919 self._rebuild_heap()920 admission_rank = (int(support), -int(key))921 if admission_rank <= self._heap[0][:2]:922 self.capacity_rejections += 1923 continue924 victim = heapq.heappop(self._heap)[2]925 del self.records[victim]926 self.records[key] = row927 self._push_heap(key)928 self.capacity_replacements += 1929 930 if len(self._heap) > 2 * self.candidate_capacity:931 self._rebuild_heap()932 933 def _prune(self) -> None:934 ranked = sorted(935 self.records.items(),936 key=lambda item: self._rank(item[0], item[1]),937 reverse=True,938 )[: self.candidate_capacity]939 self.records = dict(ranked)940 self._rebuild_heap()941 942 def snapshot_records(self) -> tuple[dict[str, float | int], ...]:943 """Return one deterministic row per retained global pair key."""944 945 self._prune()946 return tuple(947 {948 "composite_key": int(pair_key),949 "pair_key": int(pair_key),950 "language_index": -1,951 "utility": 0.0,952 "probability_sum": 0.0,953 "captured_support": int(row[2]),954 "mean_probability": 0.0,955 }956 for pair_key, row in sorted(self.records.items())957 )958 959 def state_dict(self) -> dict[str, Any]:960 """Return the complete fixed-memory frequency state for exact resume."""961 962 self._prune()963 return {964 "version": 1,965 "protocol": "independent-frequency-count-min-v1",966 "n_languages": self.n_languages,967 "candidate_capacity": self.candidate_capacity,968 "top_per_forward": self.top_per_forward,969 "sketch_width": self._sketch_width,970 "records": {971 int(key): [0.0, 0.0, int(row[2])]972 for key, row in sorted(self.records.items())973 },974 "frequency_sketch": self._frequency_sketch.detach().cpu().clone(),975 "observed_occurrences": self.observed_occurrences,976 "transferred_records": self.transferred_records,977 "admission_pressure": self.admission_pressure(),978 }979 980 def load_state_dict(self, raw: dict[str, Any]) -> None:981 """Restore a state produced by :meth:`state_dict` after validation."""982 983 if (984 int(raw.get("version", -1)) != 1985 or raw.get("protocol") != "independent-frequency-count-min-v1"986 ):987 raise ValueError("unsupported frequency discovery state")988 if int(raw.get("n_languages", -1)) != self.n_languages:989 raise ValueError("frequency discovery language count does not match")990 if int(raw.get("candidate_capacity", -1)) != self.candidate_capacity:991 raise ValueError("frequency discovery candidate capacity does not match")992 if int(raw.get("top_per_forward", -1)) != self.top_per_forward:993 raise ValueError("frequency discovery top-per-forward does not match")994 if int(raw.get("sketch_width", -1)) != self._sketch_width:995 raise ValueError("frequency discovery sketch width does not match")996 sketch = raw.get("frequency_sketch")997 if (998 not torch.is_tensor(sketch)999 or sketch.dtype != torch.long1000 or tuple(sketch.shape) != tuple(self._frequency_sketch.shape)1001 or bool((sketch < 0).any())1002 ):1003 raise ValueError("frequency discovery sketch is invalid")1004 raw_records = raw.get("records")1005 if not isinstance(raw_records, dict) or len(raw_records) > self.candidate_capacity:1006 raise ValueError("frequency discovery records are invalid")1007 records: dict[int, list[float | int]] = {}1008 for raw_key, raw_row in raw_records.items():1009 key = int(raw_key)1010 if key < 0 or not isinstance(raw_row, (list, tuple)) or len(raw_row) != 3:1011 raise ValueError("frequency discovery record is invalid")1012 support = int(raw_row[2])1013 if support < 0:1014 raise ValueError("frequency discovery support is invalid")1015 records[key] = [0.0, 0.0, support]1016 observed = int(raw.get("observed_occurrences", -1))1017 transferred = int(raw.get("transferred_records", -1))1018 if observed < 0 or transferred < 0:1019 raise ValueError("frequency discovery counters are invalid")1020 counter_names = tuple(self.admission_pressure())1021 pressure = raw.get("admission_pressure")1022 if not isinstance(pressure, dict) or set(pressure) != set(counter_names):1023 raise ValueError("frequency discovery admission counters are invalid")1024 pressure = {name: int(pressure[name]) for name in counter_names}1025 if any(value < 0 for value in pressure.values()):1026 raise ValueError("frequency discovery admission counters are invalid")1027 1028 self.records = records1029 self._frequency_sketch = sketch.detach().cpu().clone()1030 self.observed_occurrences = observed1031 self.transferred_records = transferred1032 for name, value in pressure.items():1033 setattr(self, name, value)1034 self._rebuild_heap()1035 1036 def admission_pressure(self) -> dict[str, int]:1037 return {1038 "unique_records_before_transfer": self.unique_records_before_transfer,1039 "dropped_by_top_per_forward": self.dropped_by_top_per_forward,1040 "retained_updates": self.retained_updates,1041 "admitted_into_free_slots": self.admitted_into_free_slots,1042 "capacity_replacements": self.capacity_replacements,1043 "capacity_rejections": self.capacity_rejections,1044 }1045 1046 1047 1048class CollapseTable(nn.Module):1049 def __init__(1050 self,1051 buckets: int,1052 dim: int,1053 n_languages: int,1054 bias_init: float = -2.0,1055 *,1056 pair_key_base: int | None = None,1057 pair_keys: list[int] | tuple[int, ...] | torch.Tensor | None = None,1058 pair_slots: list[int] | tuple[int, ...] | torch.Tensor | None = None,1059 reclaimed_token_ids: list[int] | tuple[int, ...] | torch.Tensor | None = None,1060 residual_buckets: int = 0,1061 ) -> None:1062 super().__init__()1063 self.buckets = int(buckets)1064 self.residual_buckets = int(residual_buckets)1065 if self.residual_buckets < 0:1066 raise ValueError("residual_buckets must be non-negative")1067 self.pair_key_base = int(pair_key_base or 0)1068 reclaimed = torch.as_tensor(1069 reclaimed_token_ids if reclaimed_token_ids is not None else [], dtype=torch.long1070 )1071 self.merged = nn.Parameter(torch.zeros(1072 0 if reclaimed.numel() else self.buckets, dim1073 ))1074 self.score = nn.Parameter(torch.zeros(self.buckets))1075 self.residual_merged = nn.Parameter(torch.zeros(self.residual_buckets, dim))1076 self.residual_score = nn.Parameter(torch.zeros(self.residual_buckets))1077 self.language_bias = nn.Parameter(torch.full((n_languages,), float(bias_init)))1078 keys = torch.as_tensor(pair_keys if pair_keys is not None else [], dtype=torch.long)1079 slots = torch.as_tensor(pair_slots if pair_slots is not None else [], dtype=torch.long)1080 if keys.numel() != slots.numel():1081 raise ValueError("pair_keys and pair_slots must have the same length")1082 if keys.numel() and not bool((keys[1:] > keys[:-1]).all()):1083 raise ValueError("pair_keys must be strictly increasing")1084 if slots.numel() and (int(slots.min()) < 0 or int(slots.max()) >= self.buckets):1085 raise ValueError("pair_slots must index the collapse table")1086 if reclaimed.numel() and reclaimed.numel() != keys.numel():1087 raise ValueError("reclaimed_token_ids must align one-to-one with pair_keys")1088 if reclaimed.unique().numel() != reclaimed.numel():1089 raise ValueError("reclaimed_token_ids must be unique")1090 self.register_buffer("pair_keys", keys, persistent=True)1091 self.register_buffer("pair_slots", slots, persistent=True)1092 self.register_buffer("reclaimed_token_ids", reclaimed, persistent=True)1093 self.discovery: (1094 ContrastiveVocabularyDiscovery | FrequencyVocabularyDiscovery | None1095 ) = None1096 self.diagnostic_discovery: (1097 ContrastiveVocabularyDiscovery | FrequencyVocabularyDiscovery | None1098 ) = None1099 self._vocabulary_telemetry: Any | None = None1100 self._loss_path_scale: float | None = None1101 self._loss_path_observer: Any | None = None1102 1103 def extra_repr(self) -> str:1104 return (f"buckets={self.buckets}, explicit_pairs={self.pair_keys.numel()}, "1105 f"residual_buckets={self.residual_buckets}, "1106 f"reclaimed_rows={self.reclaimed_token_ids.numel()}")1107 1108 def set_vocabulary_telemetry(self, window: Any | None) -> None:1109 """Attach detached run telemetry without making it module state."""1110 1111 self._vocabulary_telemetry = window1112 1113 def begin_loss_path_pass(self, scale: float, observer: Any) -> None:1114 """Scale every soft pair gate and capture occurrence gradients.1115 1116 This is intentionally transient (not module state): an audit pass may1117 read the graph but must never alter a checkpoint or discovery table.1118 """1119 1120 if self._loss_path_scale is not None or self._loss_path_observer is not None:1121 raise RuntimeError("a loss-path pass is already active")1122 if float(scale) not in (0.0, 0.5, 1.0):1123 raise ValueError("registered loss-path scale must be 0, 0.5, or 1")1124 if observer is None or not callable(observer):1125 raise TypeError("loss-path observer must be callable")1126 self._loss_path_scale = float(scale)1127 self._loss_path_observer = observer1128 1129 def end_loss_path_pass(self) -> None:1130 if self._loss_path_scale is None or self._loss_path_observer is None:1131 raise RuntimeError("no loss-path pass is active")1132 self._loss_path_scale = None1133 self._loss_path_observer = None1134 1135 def abort_loss_path_pass(self) -> None:1136 self._loss_path_scale = None1137 self._loss_path_observer = None1138 1139 def _hash(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:1140 return ((a * _MIX_A + b * _MIX_B).abs()) % self.buckets1141 1142 def _hash_residual(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:1143 if self.residual_buckets <= 0:1144 raise ValueError("residual hash requires a positive residual capacity")1145 return ((a * _MIX_A + b * _MIX_B).abs()) % self.residual_buckets1146 1147 def _exact_key(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:1148 if self.pair_key_base <= 0:1149 raise ValueError("pair_key_base is required for exact-pair vocabulary")1150 return a.to(torch.long) * self.pair_key_base + b.to(torch.long)1151 1152 def configure_discovery(1153 self,1154 *,1155 candidate_capacity: int,1156 top_per_forward: int,1157 admission_mode: str = "utility",1158 ) -> None:1159 if self.pair_keys.numel() and not self.residual_buckets:1160 raise ValueError("cannot discover after explicit vocabulary promotion")1161 discovery_class = {1162 "utility": ContrastiveVocabularyDiscovery,1163 "frequency": FrequencyVocabularyDiscovery,1164 }.get(str(admission_mode))1165 if discovery_class is None:1166 raise ValueError("discovery admission_mode must be utility or frequency")1167 self.discovery = discovery_class(1168 n_languages=self.language_bias.numel(),1169 candidate_capacity=candidate_capacity,1170 top_per_forward=top_per_forward,1171 )1172 1173 def configure_diagnostic_discovery(1174 self,1175 *,1176 candidate_capacity: int,1177 top_per_forward: int,1178 admission_mode: str = "utility",1179 ) -> None:1180 """Attach a resettable observer which cannot affect promotion state."""1181 1182 if self.pair_keys.numel() and not self.residual_buckets:1183 raise ValueError("cannot diagnose discovery after explicit promotion")1184 discovery_class = {1185 "utility": ContrastiveVocabularyDiscovery,1186 "frequency": FrequencyVocabularyDiscovery,1187 }.get(str(admission_mode))1188 if discovery_class is None:1189 raise ValueError("discovery admission_mode must be utility or frequency")1190 self.diagnostic_discovery = discovery_class(1191 n_languages=self.language_bias.numel(),1192 candidate_capacity=candidate_capacity,1193 top_per_forward=top_per_forward,1194 )1195 1196 def _explicit_lookup(self, key: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:1197 if not self.pair_keys.numel():1198 return torch.zeros_like(key, dtype=torch.bool), torch.zeros_like(key)1199 positions = torch.searchsorted(self.pair_keys, key)1200 safe = positions.clamp(max=self.pair_keys.numel() - 1)