Download loopq_quantization/scripts/loopq/calibrate.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 9.41 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/calibrate.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/calibrate.py
-
curl -L -o calibrate.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/calibrate.py
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 | |
| 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") | |
| class CalibrationParameterAllowlist: | |
| named_parameters: dict[str, nn.Parameter] | |
| 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") | |
| 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"]) | |