File size: 9,413 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
"""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"])