File size: 15,277 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
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
"""Real 32-recurrence Huginn calibration primitives for faithful LoopQ."""

from __future__ import annotations

from collections import defaultdict
from collections.abc import Mapping
from pathlib import Path
from typing import Any

import torch
import torch.nn.functional as F
from torch import nn

from adapters.huginn import HUGINN_LOOP_COUNT, HUGINN_PHYSICAL_LAYERS, PAPER_GROUPS
from .cta import CrossLoopTransitionAdapter
from .las import LoopAwareActivationScales
from .objective import AdaptiveMuCache, TrajectoryLoss, trajectory_aware_loss
from .ouro_calibration import (
    OuroCalibrationDataConfig,
    OuroLASStatisticsCollector,
    _extract_logits,
    _ste, _weight_ste, qdq_linear,
    load_pile_texts,
)
from .quantization import quantize_weight
from .sharing_gap import SharingGapStatistics
from .transforms import FlatQuantSVDKroneckerTransform, SharedKroneckerTransform


from loopq.paths import pinned_snapshot

PINNED_HUGINN_SNAPSHOT = pinned_snapshot("huginn")


def load_pinned_teacher_student(device: str):
    """Load independent pinned BF16 Huginn teacher and student models."""
    from transformers import AutoModelForCausalLM, AutoTokenizer

    common = dict(
        pretrained_model_name_or_path=str(PINNED_HUGINN_SNAPSHOT),
        trust_remote_code=True, local_files_only=True, torch_dtype=torch.bfloat16,
    )
    teacher = AutoModelForCausalLM.from_pretrained(**common).to(device).eval()
    student = AutoModelForCausalLM.from_pretrained(**common).to(device).eval()
    teacher.config.mean_recurrence = HUGINN_LOOP_COUNT
    student.config.mean_recurrence = HUGINN_LOOP_COUNT
    for parameter in teacher.parameters():
        parameter.requires_grad_(False)
    tokenizer = AutoTokenizer.from_pretrained(
        PINNED_HUGINN_SNAPSHOT, trust_remote_code=True, local_files_only=True
    )
    return teacher, student, tokenizer


def huginn_projection_sites(model: nn.Module) -> dict[str, tuple[nn.Module, ...]]:
    sites = {}
    for layer_idx, layer in enumerate(model.transformer.core_block):
        if layer_idx >= HUGINN_PHYSICAL_LAYERS:
            raise ValueError("pinned Huginn recurrent block has more than four layers")
        for group, spec in PAPER_GROUPS.items():
            modules = []
            for path in spec["hf_weights"]:
                current = layer
                for part in path.split("."):
                    current = getattr(current, part)
                modules.append(current)
            sites[f"transformer.core_block.{layer_idx}.{group}"] = tuple(modules)
    if len(sites) != HUGINN_PHYSICAL_LAYERS * len(PAPER_GROUPS):
        raise ValueError("Huginn LoopQ requires exactly 4x4 projection groups")
    return sites


class HuginnTrajectoryCapture:
    """Capture state after every recurrent block and insert all 31 CTA edges."""

    def __init__(self, final_norm: nn.Module, cta: CrossLoopTransitionAdapter | None):
        self.final_norm = final_norm
        self.cta = cta
        self.pre_cta: list[torch.Tensor] = []
        self.adapted: list[torch.Tensor] = []
        self._handle = None

    def __enter__(self):
        self.pre_cta.clear()
        self.adapted.clear()

        def hook(_module, _inputs, output):
            recurrence = len(self.pre_cta)
            if recurrence >= HUGINN_LOOP_COUNT:
                raise RuntimeError("Huginn recurrent block ran more than 32 times")
            self.pre_cta.append(output)
            if recurrence < HUGINN_LOOP_COUNT - 1 and self.cta is not None:
                output = self.cta(output, recurrence)
                self.adapted.append(output)
            return output

        self._handle = self.final_norm.register_forward_hook(hook)
        return self

    def __exit__(self, exception_type, _exception, _traceback):
        self._handle.remove()
        self._handle = None
        if exception_type is None and len(self.pre_cta) != HUGINN_LOOP_COUNT:
            raise RuntimeError(
                f"expected 32 recurrent states, captured {len(self.pre_cta)}"
            )


class HuginnDifferentiableQDQ(nn.Module):
    """LoopQ W4/A4-or-A8 hooks for Huginn's 16 shared projection groups."""

    def __init__(self, *, model, las, activation_bits, factor_by_width, checkpoint_linears=False, activation_ste="identity"):
        super().__init__()
        self.model, self.las, self.activation_bits = model, las, activation_bits
        if activation_ste not in {"identity", "rounding"}:
            raise ValueError("unknown activation STE")
        self.activation_ste = activation_ste
        self.checkpoint_linears = checkpoint_linears
        self.sites = huginn_projection_sites(model)
        self._encoded = {key: key.replace(".", "__") for key in self.sites}
        transforms = {}
        for key, modules in self.sites.items():
            width = modules[0].in_features
            factors = factor_by_width.get(width)
            if factors is None or factors[0] * factors[1] != width:
                raise ValueError(f"missing valid factors for Huginn width {width}")
            transforms[self._encoded[key]] = FlatQuantSVDKroneckerTransform(*factors)
        self.shared_transforms = nn.ModuleDict(transforms)
        self.selected_loop_transforms = nn.ModuleDict()
        self.statistics_loop_transforms = nn.ModuleDict()
        self._statistics_mode = False
        self._handles, self._calls = [], defaultdict(int)
        self._records = defaultdict(list)
        self._execution_views = {}
        for parameter in model.parameters():
            parameter.requires_grad_(False)

    def select_group(self, key: str) -> None:
        if key not in self.sites:
            raise KeyError(key)
        encoded = self._encoded[key]
        if encoded in self.selected_loop_transforms:
            return
        base = self.shared_transforms[encoded]
        loops = nn.ModuleList([
            base.fresh_copy() for _ in range(HUGINN_LOOP_COUNT)
        ])
        self.selected_loop_transforms[encoded] = loops

    def _base_transform_for(self, key: str, recurrence: int):
        encoded = self._encoded[key]
        if encoded in self.selected_loop_transforms:
            return self.selected_loop_transforms[encoded][recurrence]
        return self.shared_transforms[encoded]

    def transform_for(self, key: str, recurrence: int):
        encoded = self._encoded[key]
        if self._statistics_mode:
            return self.statistics_loop_transforms[encoded][recurrence]
        return self._base_transform_for(key, recurrence)

    def _build_statistics_transforms(self) -> None:
        transforms = {}
        for key in sorted(self.sites):
            copies = []
            for recurrence in range(HUGINN_LOOP_COUNT):
                base = self._base_transform_for(key, recurrence)
                copy = base.fresh_copy()
                copies.append(copy)
            transforms[self._encoded[key]] = nn.ModuleList(copies)
        self.statistics_loop_transforms = nn.ModuleDict(transforms)

    def begin(self, *, statistics_mode: bool = False) -> None:
        if self._handles:
            raise RuntimeError("QDQ hooks are already active")
        self._statistics_mode = statistics_mode
        if statistics_mode:
            self._build_statistics_transforms()
        self._calls.clear()
        self._records.clear()
        self._execution_views.clear()
        for key, modules in self.sites.items():
            for module in modules:
                def pre_hook(current, inputs, key=key):
                    recurrence = self._calls[id(current)]
                    self._calls[id(current)] += 1
                    if recurrence >= HUGINN_LOOP_COUNT:
                        raise RuntimeError("projection invoked more than 32 recurrences")
                    transform_module = self.transform_for(key, recurrence)
                    transform = self._execution_views.get(id(transform_module))
                    if transform is None:
                        transform = transform_module.materialize()
                        self._execution_views[id(transform_module)] = transform
                    current._loopq_override = qdq_linear(
                        inputs[0], current, transform, self.las, key, recurrence,
                        self.activation_bits, checkpoint=self.checkpoint_linears,
                        activation_ste=self.activation_ste,
                    )
                    # The original linear output is replaced below; its graph
                    # is unused. Keep its execution free of saved tensors.
                    return tuple(x.detach() if isinstance(x, torch.Tensor) else x for x in inputs)

                def post_hook(current, _inputs, _output):
                    replacement = current._loopq_override
                    del current._loopq_override
                    return replacement

                self._handles.append(module.register_forward_pre_hook(pre_hook))
                self._handles.append(module.register_forward_hook(post_hook))

    def end(self, *, validate: bool = True) -> None:
        for handle in self._handles:
            handle.remove()
        self._handles.clear()
        self._execution_views.clear()
        incomplete = [
            key for key, modules in self.sites.items() for module in modules
            if self._calls[id(module)] != HUGINN_LOOP_COUNT
        ]
        if validate and incomplete:
            raise RuntimeError(f"incomplete Huginn projection trajectories: {incomplete[:4]}")
        self._statistics_mode = False
        self.statistics_loop_transforms = nn.ModuleDict()

    def sharing_gap_statistics(self, loss: torch.Tensor):
        if not self._statistics_mode:
            raise RuntimeError("sharing-gap statistics require statistics_mode")
        result = {}
        transforms = [
            self.transform_for(key, recurrence)
            for key in sorted(self.sites)
            for recurrence in range(HUGINN_LOOP_COUNT)
        ]
        all_parameters = tuple(
            parameter for transform in transforms for parameter in transform.parameters()
        )
        all_gradients = torch.autograd.grad(
            loss, all_parameters, retain_graph=True, allow_unused=True
        )
        gradient_by_id = {
            id(parameter): gradient
            for parameter, gradient in zip(all_parameters, all_gradients)
        }
        for key in sorted(self.sites):
            per_recurrence = []
            for recurrence in range(HUGINN_LOOP_COUNT):
                transform = self.transform_for(key, recurrence)
                params = tuple(transform.parameters())
                accumulated = [
                    (gradient_by_id[id(parameter)].detach()
                     if gradient_by_id[id(parameter)] is not None
                     else torch.zeros_like(parameter))
                    for parameter in params
                ]
                per_recurrence.append(
                    torch.cat([item.flatten() for item in accumulated]).cpu()
                )
            gradients = torch.stack(per_recurrence)
            result[key] = SharingGapStatistics(
                gradients=gradients,
                # Square in FP64: finite FP32 VJPs can exceed sqrt(FP32_MAX).
                fisher_diagonal=gradients.to(torch.float64).square().mean(dim=0),
                # Selecting a group replaces one shared transform by 32 copies.
                parameter_count=gradients.shape[1] * (HUGINN_LOOP_COUNT - 1),
            )
        return result

    def export_shared(self):
        return {
            key: self.shared_transforms[self._encoded[key]].export_state()
            for key in sorted(self.sites)
        }

    def export_selected(self):
        return {
            key: {
                str(recurrence): transform.export_state()
                for recurrence, transform in enumerate(
                    self.selected_loop_transforms[self._encoded[key]]
                )
            }
            for key in sorted(self.sites)
            if self._encoded[key] in self.selected_loop_transforms
        }


def huginn_trajectory_loss(
    *, teacher, student, student_qdq, cta, inputs, step, mu_cache,
    collect_statistics=True, offload_saved_tensors=False,
) -> tuple[TrajectoryLoss, dict[str, SharingGapStatistics]]:
    """Execute true 32-step teacher/student trajectories and LoopQ Eq. 8."""
    final_teacher_norm = teacher.transformer.core_block[-1].norm_4
    final_student_norm = student.transformer.core_block[-1].norm_4
    # Replay the same random initial recurrent state/noise for the student.
    # Teacher consumes no global RNG progress; student advances it once.
    devices = sorted({value.device.index for value in inputs.values()
                      if isinstance(value, torch.Tensor) and value.is_cuda})
    with torch.random.fork_rng(devices=devices), torch.no_grad(), HuginnTrajectoryCapture(final_teacher_norm, None) as teacher_trace:
        teacher_output = teacher(**inputs, use_cache=False, num_steps=32)
    teacher_hidden = torch.stack(teacher_trace.pre_cta)

    student_qdq.begin(statistics_mode=collect_statistics)
    failed = True
    try:
        saved_tensors = (
            # Huginn's 32-step graph exceeds practical pinned-host budgets;
            # pageable CPU storage retains exact autograd semantics without
            # routing every saved tensor through the CUDA host allocator.
            torch.autograd.graph.save_on_cpu(pin_memory=False, device_type="cuda")
            if offload_saved_tensors else torch.autograd.graph.saved_tensors_hooks(
                lambda tensor: tensor, lambda tensor: tensor
            )
        )
        with saved_tensors:
            with HuginnTrajectoryCapture(final_student_norm, cta) as student_trace:
                # The pinned Huginn interprets a pair as [no-grad, with-grad] steps.
                grad_schedule = torch.tensor([0, HUGINN_LOOP_COUNT])
                student_output = student(
                    **inputs, use_cache=False, num_steps=grad_schedule
                )
                student_hidden = torch.stack(student_trace.pre_cta)
                adapted = torch.stack(student_trace.adapted)
                mu = mu_cache.get(step, teacher_hidden, student_hidden)
                loss = trajectory_aware_loss(
                    teacher_logits=_extract_logits(teacher_output),
                    student_logits=_extract_logits(student_output),
                    teacher_hidden=teacher_hidden, student_hidden=student_hidden,
                    adapted_transitions=adapted,
                    teacher_next_inputs=teacher_hidden[:-1], mu=mu,
                    include_transition=cta.enabled,
                )
                statistics = (
                    student_qdq.sharing_gap_statistics(loss.total)
                    if collect_statistics else {}
                )
        failed = False
    finally:
        student_qdq.end(validate=not failed)
    return loss, statistics


__all__ = [
    "HuginnDifferentiableQDQ", "HuginnTrajectoryCapture",
    "OuroCalibrationDataConfig", "OuroLASStatisticsCollector",
    "huginn_projection_sites", "huginn_trajectory_loss",
    "load_pile_texts", "load_pinned_teacher_student",
]