JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
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,
)
@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),
)