"""LoopQ sharing-gap scoring and progressive SLT selection (LQ4). Equation (6) scores a transform-weight group by summing cross-loop gradient variance after normalization by diagonal Fisher curvature. This module consumes saved calibration statistics; collection and trajectory calibration belong to LQ6. """ from __future__ import annotations import math from collections.abc import Callable, Mapping from dataclasses import dataclass from typing import Any import torch DEFAULT_FISHER_EPSILON = 1e-8 OURO_1_4B_GROUP_BUDGET = 4 HUGINN_PARAMETER_FRACTION = 0.05 def loss_value_gradients( loss: torch.Tensor, values: tuple[torch.Tensor, ...], ) -> tuple[torch.Tensor | None, ...]: """Evaluate loss-to-intermediate VJPs in one full-graph traversal.""" return torch.autograd.grad( loss, values, retain_graph=True, allow_unused=True ) def batched_chain_rule_gradients( values: tuple[torch.Tensor, ...], value_gradients: tuple[torch.Tensor | None, ...], parameters: tuple[torch.nn.Parameter, ...], ) -> tuple[torch.Tensor | None, ...]: """Sum exact intermediate-to-parameter VJPs for supplied loss gradients.""" active = tuple( (value, gradient) for value, gradient in zip(values, value_gradients) if gradient is not None ) if not active: return tuple(None for _ in parameters) return torch.autograd.grad( tuple(item[0] for item in active), parameters, grad_outputs=tuple(item[1] for item in active), retain_graph=True, allow_unused=True, ) @dataclass(frozen=True) class SharingGapStatistics: """Saved sufficient statistics for one selectable transform-weight group.""" gradients: torch.Tensor fisher_diagonal: torch.Tensor parameter_count: int def to_dict(self) -> dict[str, Any]: self.validate() return { "gradients": self.gradients.detach().cpu(), "fisher_diagonal": self.fisher_diagonal.detach().cpu(), "parameter_count": self.parameter_count, } @classmethod def from_dict(cls, state: Mapping[str, Any]) -> "SharingGapStatistics": required = {"gradients", "fisher_diagonal", "parameter_count"} missing = required.difference(state) if missing: raise ValueError(f"sharing-gap statistics are missing fields: {sorted(missing)}") item = cls( gradients=torch.as_tensor(state["gradients"]), fisher_diagonal=torch.as_tensor(state["fisher_diagonal"]), parameter_count=int(state["parameter_count"]), ) item.validate() return item def validate(self, *, expected_loop_count: int | None = None) -> None: if self.gradients.ndim < 2 or self.gradients.shape[0] < 2: raise ValueError("gradients must have shape (at least 2 loops, coordinates...)") if expected_loop_count is not None and self.gradients.shape[0] != expected_loop_count: raise ValueError( f"expected {expected_loop_count} loops, got {self.gradients.shape[0]}" ) if tuple(self.fisher_diagonal.shape) != tuple(self.gradients.shape[1:]): raise ValueError( "fisher_diagonal shape must equal the gradient coordinate shape, " f"got {tuple(self.fisher_diagonal.shape)} and {tuple(self.gradients.shape[1:])}" ) if not self.gradients.is_floating_point() or not self.fisher_diagonal.is_floating_point(): raise TypeError("gradients and fisher_diagonal must be floating point") if not torch.isfinite(self.gradients).all() or not torch.isfinite(self.fisher_diagonal).all(): raise ValueError("sharing-gap statistics must be finite") if (self.fisher_diagonal < 0).any(): raise ValueError("fisher_diagonal must be non-negative") if self.parameter_count <= 0: raise ValueError("parameter_count must be positive") def sharing_gap_score( statistics: SharingGapStatistics, *, epsilon: float = DEFAULT_FISHER_EPSILON, expected_loop_count: int | None = None, ) -> float: """Compute Eq. (6) with population variance (division by T).""" statistics.validate(expected_loop_count=expected_loop_count) if not math.isfinite(epsilon) or epsilon <= 0: raise ValueError("epsilon must be finite and positive") gradients = statistics.gradients.detach().to(device="cpu", dtype=torch.float64) fisher = statistics.fisher_diagonal.detach().to(device="cpu", dtype=torch.float64) variance = gradients.var(dim=0, unbiased=False) return float((variance / (fisher + epsilon)).sum().item()) @dataclass(frozen=True) class SelectionRound: round_index: int scores: dict[str, float] selected_group: str selected_parameter_count: int cumulative_parameter_count: int @dataclass(frozen=True) class ProgressiveSelectionArtifact: """Serializable audit trail proving selection and score recomputation.""" format_version: int loop_count: int epsilon: float stopping_rule: str requested_budget: float | int selected_groups: tuple[str, ...] rounds: tuple[SelectionRound, ...] def to_dict(self) -> dict[str, Any]: return { "format": "loopq_progressive_slt", "format_version": self.format_version, "loop_count": self.loop_count, "epsilon": self.epsilon, "stopping_rule": self.stopping_rule, "requested_budget": self.requested_budget, "selected_groups": list(self.selected_groups), "rounds": [ { "round_index": item.round_index, "scores": dict(sorted(item.scores.items())), "selected_group": item.selected_group, "selected_parameter_count": item.selected_parameter_count, "cumulative_parameter_count": item.cumulative_parameter_count, } for item in self.rounds ], } StatisticsCallback = Callable[ [tuple[str, ...], int], Mapping[str, SharingGapStatistics] ] def _validate_candidate_names(statistics: Mapping[str, SharingGapStatistics]) -> None: if not statistics: raise ValueError("at least one SLT candidate is required") if any(not isinstance(name, str) or not name for name in statistics): raise ValueError("SLT candidate names must be non-empty strings") def progressive_select_by_group_count( recompute_statistics: StatisticsCallback, *, budget: int, loop_count: int, epsilon: float = DEFAULT_FISHER_EPSILON, ) -> ProgressiveSelectionArtifact: """Progressively select exactly ``budget`` groups, recomputing every round.""" if budget <= 0: raise ValueError("group budget must be positive") return _progressive_select( recompute_statistics, loop_count=loop_count, epsilon=epsilon, group_budget=budget, parameter_target=None, requested_budget=budget, stopping_rule="group_count", ) def huginn_parameter_target(parameter_counts: Mapping[str, int]) -> int: """Return ceil(5% of all candidate transform parameters).""" if not parameter_counts or any(count <= 0 for count in parameter_counts.values()): raise ValueError("Huginn candidate parameter counts must be positive") return math.ceil(sum(parameter_counts.values()) * HUGINN_PARAMETER_FRACTION) def progressive_select_huginn( recompute_statistics: StatisticsCallback, *, loop_count: int = 32, epsilon: float = DEFAULT_FISHER_EPSILON, ) -> ProgressiveSelectionArtifact: """Select whole groups until their parameters first meet the 5% target.""" initial = recompute_statistics(tuple(), 0) _validate_candidate_names(initial) target = huginn_parameter_target( {name: item.parameter_count for name, item in initial.items()} ) def cached_first_callback(selected: tuple[str, ...], round_index: int): if round_index == 0 and not selected: return initial return recompute_statistics(selected, round_index) return _progressive_select( cached_first_callback, loop_count=loop_count, epsilon=epsilon, group_budget=None, parameter_target=target, requested_budget=HUGINN_PARAMETER_FRACTION, stopping_rule="parameter_fraction_ceil_then_whole_group_overshoot", ) def _progressive_select( recompute_statistics: StatisticsCallback, *, loop_count: int, epsilon: float, group_budget: int | None, parameter_target: int | None, requested_budget: float | int, stopping_rule: str, ) -> ProgressiveSelectionArtifact: if loop_count < 2: raise ValueError("loop_count must be at least 2") selected: list[str] = [] rounds: list[SelectionRound] = [] cumulative_parameters = 0 expected_names: set[str] | None = None expected_parameter_counts: dict[str, int] | None = None while True: if group_budget is not None and len(selected) >= group_budget: break if parameter_target is not None and cumulative_parameters >= parameter_target: break statistics = recompute_statistics(tuple(selected), len(rounds)) _validate_candidate_names(statistics) current_names = set(statistics) current_counts = {name: item.parameter_count for name, item in statistics.items()} if expected_names is None: expected_names = current_names expected_parameter_counts = current_counts elif current_names != expected_names or current_counts != expected_parameter_counts: raise ValueError( "progressive recomputation must preserve the candidate set and parameter counts" ) remaining = sorted(set(statistics).difference(selected)) if not remaining: raise ValueError("SLT budget cannot be satisfied by remaining candidates") scored = { name: sharing_gap_score( statistics[name], epsilon=epsilon, expected_loop_count=loop_count ) for name in remaining } # Descending score and then lexical name makes ties deterministic. chosen = min(remaining, key=lambda name: (-scored[name], name)) count = statistics[chosen].parameter_count cumulative_parameters += count selected.append(chosen) rounds.append( SelectionRound( round_index=len(rounds), scores=dict(sorted(scored.items())), selected_group=chosen, selected_parameter_count=count, cumulative_parameter_count=cumulative_parameters, ) ) return ProgressiveSelectionArtifact( format_version=1, loop_count=loop_count, epsilon=epsilon, stopping_rule=stopping_rule, requested_budget=requested_budget, selected_groups=tuple(selected), rounds=tuple(rounds), )