CoolFace
Modelpublic

BorisTM/loss-guided-static-multi

sourceHugging Faceapache-2.0updated 23h agoView on Hugging Face
0likes8downloads
collapse.py1931 linesDownload Raw Back to root
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)

Showing the first 1,200 of 1931 lines. Download the file for the rest.