JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
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",
]