Download loopq_quantization/scripts/loopq/huginn_calibration.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/huginn_calibration.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/huginn_calibration.py
-
curl -L -o huginn_calibration.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/huginn_calibration.py
15.3 kB
| """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", | |
| ] | |