Download loopq_quantization/scripts/loopq/sharing_gap.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/sharing_gap.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/sharing_gap.py
-
curl -L -o sharing_gap.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/sharing_gap.py
11.2 kB
| """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, | |
| ) | |
| 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, | |
| } | |
| 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()) | |
| class SelectionRound: | |
| round_index: int | |
| scores: dict[str, float] | |
| selected_group: str | |
| selected_parameter_count: int | |
| cumulative_parameter_count: int | |
| 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), | |
| ) | |