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