"""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"])