JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
9.41 kB
"""Optimizer ownership, tracing, and checkpoint helpers for LoopQ LQ6."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import math
import torch
from torch import nn
from .objective import AdaptiveMuCache, TrajectoryLoss
PAPER_GRADIENT_CLIP_NORM = 1.0
@dataclass(frozen=True)
class CalibrationConfig:
"""Required choices omitted by the LoopQ paper; no silent defaults."""
optimizer: str
learning_rate: float
total_steps: int
las_learning_rate: float | None = None
transform_learning_rate: float | None = None
cta_learning_rate: float | None = None
mu_update_interval: int = 100
gradient_clip_norm: float = PAPER_GRADIENT_CLIP_NORM
weight_decay: float | None = None
betas: tuple[float, float] = (0.9, 0.999)
adam_epsilon: float = 1e-8
def resolved_optimizer(self) -> dict:
"""Persist resolved values rather than depending on library defaults."""
return dict(lr=self.learning_rate,
weight_decay=(0.01 if self.optimizer == "adamw" else 0.0)
if self.weight_decay is None else self.weight_decay,
betas=tuple(self.betas), eps=self.adam_epsilon)
def validate(self) -> None:
if self.optimizer not in {"adam", "adamw"}:
raise ValueError("optimizer must be explicitly 'adam' or 'adamw'")
if not math.isfinite(self.learning_rate) or self.learning_rate <= 0:
raise ValueError("learning_rate must be positive")
if self.total_steps <= 0:
raise ValueError("total_steps must be positive")
component_rates = (self.las_learning_rate, self.transform_learning_rate,
self.cta_learning_rate)
if any(rate is not None and (not math.isfinite(rate) or rate <= 0) for rate in component_rates):
raise ValueError("component learning rates must be positive when set")
if self.mu_update_interval <= 0:
raise ValueError("mu_update_interval must be positive")
if self.gradient_clip_norm != PAPER_GRADIENT_CLIP_NORM:
raise ValueError("LoopQ paper specifies gradient clipping 1.0")
resolved = self.resolved_optimizer()
if not math.isfinite(resolved["weight_decay"]) or resolved["weight_decay"] < 0:
raise ValueError("weight_decay must be finite and non-negative")
if len(self.betas) != 2 or any(not math.isfinite(b) or not 0 <= b < 1 for b in self.betas):
raise ValueError("Adam betas must lie in [0, 1)")
if not math.isfinite(self.adam_epsilon) or self.adam_epsilon <= 0:
raise ValueError("Adam epsilon must be finite and positive")
@dataclass(frozen=True)
class CalibrationParameterAllowlist:
named_parameters: dict[str, nn.Parameter]
@property
def parameters(self) -> tuple[nn.Parameter, ...]:
return tuple(self.named_parameters.values())
def build_parameter_allowlist(
*,
las: nn.Module,
shared_transforms: Mapping[str, nn.Module],
selected_loop_transforms: Mapping[str, nn.Module],
cta: nn.Module,
backbone: nn.Module,
) -> CalibrationParameterAllowlist:
"""Allow only paper-defined calibration parameters, never backbone weights."""
named: dict[str, nn.Parameter] = {}
seen: set[int] = set()
categories = {
"las": {"las": las},
"shared_transform": shared_transforms,
"selected_loop_transform": selected_loop_transforms,
"cta": {"cta": cta},
}
for category, modules in categories.items():
for module_name in sorted(modules):
for parameter_name, parameter in modules[module_name].named_parameters():
if not parameter.requires_grad:
continue
if id(parameter) in seen:
raise ValueError("a calibration parameter is owned by multiple allowlist entries")
seen.add(id(parameter))
named[f"{category}.{module_name}.{parameter_name}"] = parameter
backbone_ids = {id(parameter) for parameter in backbone.parameters()}
if seen.intersection(backbone_ids):
raise ValueError("backbone parameters must never enter the LoopQ calibration allowlist")
if not named:
raise ValueError("calibration parameter allowlist must not be empty")
return CalibrationParameterAllowlist(dict(sorted(named.items())))
def build_optimizer(
allowlist: CalibrationParameterAllowlist, config: CalibrationConfig
) -> torch.optim.Optimizer:
config.validate()
optimizer_type = torch.optim.AdamW if config.optimizer == "adamw" else torch.optim.Adam
rates = {
"las": config.las_learning_rate,
"shared_transform": config.transform_learning_rate,
"selected_loop_transform": config.transform_learning_rate,
"cta": config.cta_learning_rate,
}
groups: dict[float, list[nn.Parameter]] = {}
for name, parameter in allowlist.named_parameters.items():
category = name.split(".", 1)[0]
rate = rates[category] if rates[category] is not None else config.learning_rate
groups.setdefault(float(rate), []).append(parameter)
return optimizer_type(
[{"params": parameters, "lr": rate} for rate, parameters in sorted(groups.items())],
**config.resolved_optimizer(),
)
def verify_optimizer_allowlist(
optimizer: torch.optim.Optimizer, allowlist: CalibrationParameterAllowlist
) -> None:
actual = [parameter for group in optimizer.param_groups for parameter in group["params"]]
if len(actual) != len(set(map(id, actual))) or set(map(id, actual)) != set(
map(id, allowlist.parameters)
):
raise ValueError("optimizer parameters do not exactly match the LoopQ allowlist")
@torch.no_grad()
def clip_calibration_gradients(parameters: Sequence[nn.Parameter], max_norm: float) -> torch.Tensor:
"""Compute the global L2 norm in FP64 before applying paper clipping.
FP32 sum-of-squares can overflow even when every gradient is finite.
Preserve genuine nonfinite failures; never replace them with zeros.
"""
gradients = [parameter.grad for parameter in parameters if parameter.grad is not None]
if not gradients:
return torch.zeros((), dtype=torch.float64)
device = gradients[0].device
norms = [torch.linalg.vector_norm(gradient.to(torch.float64)).to(device)
for gradient in gradients]
norm = torch.linalg.vector_norm(torch.stack(norms))
if not torch.isfinite(norm):
raise RuntimeError("non-finite calibration gradient norm")
coefficient = (max_norm / (norm + 1e-6)).clamp(max=1.0)
for gradient in gradients:
gradient.mul_(coefficient.to(device=gradient.device, dtype=gradient.dtype))
return norm
def calibration_step(
*,
loss: TrajectoryLoss,
optimizer: torch.optim.Optimizer,
allowlist: CalibrationParameterAllowlist,
gradient_clip_norm: float = PAPER_GRADIENT_CLIP_NORM,
) -> dict[str, float | int | list[float]]:
if gradient_clip_norm != PAPER_GRADIENT_CLIP_NORM:
raise ValueError("LoopQ paper specifies gradient clipping 1.0")
verify_optimizer_allowlist(optimizer, allowlist)
optimizer.zero_grad(set_to_none=True)
if not torch.isfinite(loss.total):
raise ValueError("non-finite calibration loss")
loss.total.backward()
norm = clip_calibration_gradients(allowlist.parameters, gradient_clip_norm)
optimizer.step()
trace = loss.trace()
trace["gradient_norm_before_clip"] = float(norm.detach())
return trace
def save_calibration_checkpoint(
path: str | Path,
*,
step: int,
config: CalibrationConfig,
modules: Mapping[str, nn.Module],
optimizer: torch.optim.Optimizer,
mu_cache: AdaptiveMuCache,
traces: Sequence[Mapping[str, Any]],
) -> None:
config.validate()
artifact = {
"format": "loopq_trajectory_calibration",
"format_version": 1,
"step": int(step),
"config": config.__dict__,
"modules": {name: module.state_dict() for name, module in sorted(modules.items())},
"optimizer": optimizer.state_dict(),
"mu_cache": mu_cache.state_dict(),
"traces": [dict(item) for item in traces],
}
torch.save(artifact, Path(path))
def load_calibration_checkpoint(
path: str | Path,
*,
config: CalibrationConfig,
modules: Mapping[str, nn.Module],
optimizer: torch.optim.Optimizer,
mu_cache: AdaptiveMuCache,
) -> tuple[int, list[dict[str, Any]]]:
artifact = torch.load(Path(path), map_location="cpu", weights_only=True)
if artifact.get("format") != "loopq_trajectory_calibration" or artifact.get("format_version") != 1:
raise ValueError("unsupported calibration checkpoint format or version")
if artifact["config"] != config.__dict__:
raise ValueError("calibration checkpoint config does not match requested config")
if set(artifact["modules"]) != set(modules):
raise ValueError("calibration checkpoint module set does not match")
for name, module in modules.items():
module.load_state_dict(artifact["modules"][name])
optimizer.load_state_dict(artifact["optimizer"])
mu_cache.load_state_dict(artifact["mu_cache"])
return int(artifact["step"]), list(artifact["traces"])