"""Real-data Ouro calibration primitives for the LQ7 driver.""" from __future__ import annotations from dataclasses import asdict, dataclass from pathlib import Path from collections import defaultdict from collections.abc import Mapping from typing import Any import torch import torch.nn.functional as F from torch import nn from adapters.ouro import PAPER_GROUPS from .cta import CrossLoopTransitionAdapter from .las import LoopAwareActivationScales from .objective import AdaptiveMuCache, TrajectoryLoss, trajectory_aware_loss from .quantization import quantize_weight from .sharing_gap import SharingGapStatistics from .transforms import FlatQuantSVDKroneckerTransform, SharedKroneckerTransform from loopq.paths import pinned_snapshot PINNED_OURO_SNAPSHOT = pinned_snapshot("ouro") PILE_DATASET = "mit-han-lab/pile-val-backup" PILE_DATASET_REVISION = "2f5e46ae6a69cf0dce4b12f78241c408936ca0e4" @dataclass(frozen=True) class OuroCalibrationDataConfig: samples: int = 1024 max_length: int = 256 split: str = "validation" text_field: str = "text" smoke_small: bool = False @classmethod def paper(cls) -> "OuroCalibrationDataConfig": return cls() @classmethod def smoke(cls, samples: int = 2) -> "OuroCalibrationDataConfig": if not 1 <= samples <= 2: raise ValueError("--smoke-small permits only 1 or 2 samples") # A smoke validates the real-model forward/backward/export path, not # sequence-length scaling. Keeping 256 tokens here produces roughly # 200 GiB of exact saved tensors and turns a pipeline check into an # hour-scale calibration. Paper runs retain the 256-token contract. return cls(samples=samples, max_length=16, smoke_small=True) def metadata(self) -> dict[str, Any]: value = asdict(self) value.update({ "dataset": PILE_DATASET, "dataset_revision": PILE_DATASET_REVISION, "streaming": True, "paper_calibration": not self.smoke_small and self.samples == 1024 and self.max_length == 256, "warning": ( None if not self.smoke_small else "Pipeline validation only; not a paper calibration result." ), }) return value class OuroLASStatisticsCollector: """Online per-module/loop/group absmax collector with bounded memory.""" def __init__(self, *, bits: int, loop_count: int = 4, group_size: int = 32) -> None: if bits not in (4, 8): raise ValueError("activation bits must be 4 or 8") self.bits = bits self.loop_count = loop_count self.group_size = group_size self._absmax: dict[str, torch.Tensor] = {} self._calls: dict[str, int] = {} self._sample_calls: dict[str, int] | None = None def begin_sample(self) -> None: if self._sample_calls is not None: raise RuntimeError("LAS sample already active") self._sample_calls = {} def end_sample(self, *, expected_modules: set[str]) -> None: calls, self._sample_calls = self._sample_calls, None if calls is None or set(calls) != expected_modules or any(n != self.loop_count for n in calls.values()): raise ValueError("LAS sample must visit every module exactly four times") def observe(self, module_key: str, value: torch.Tensor) -> int: if self._sample_calls is not None: n = self._sample_calls.get(module_key, 0) if n >= self.loop_count: raise ValueError("LAS sample loop routing cannot wrap") self._sample_calls[module_key] = n + 1 loop = self._calls.get(module_key, 0) % self.loop_count self._calls[module_key] = self._calls.get(module_key, 0) + 1 width = value.shape[-1] groups = (width + self.group_size - 1) // self.group_size if module_key not in self._absmax: self._absmax[module_key] = torch.zeros(self.loop_count, groups) elif self._absmax[module_key].shape[1] != groups: raise ValueError("module feature width changed during LAS collection") work = value.detach().abs() work = F.pad(work, (0, groups * self.group_size - width)) maximum = work.reshape(-1, groups, self.group_size).amax(dim=(0, 2)).cpu() self._absmax[module_key][loop] = torch.maximum(self._absmax[module_key][loop], maximum) return loop def build_las(self) -> LoopAwareActivationScales: if self._sample_calls is not None: raise ValueError("LAS sample still active") if not self._absmax: raise ValueError("no activation statistics were collected") expected_calls = {name: count for name, count in self._calls.items() if count % self.loop_count} if expected_calls: raise ValueError(f"incomplete recurrent trajectories for modules: {expected_calls}") # Section 4.1 and Appendix B.3 specify one LAS scalar per module and # loop (O(TL)), while activation quantization remains group-wise. The # observed tensors establish routing/shape coverage; runtime group # absmax values are dynamic and LAS learns their clipping multiplier. return LoopAwareActivationScales.dynamic_for_modules( sorted(self._absmax), loop_count=self.loop_count, group_size=self.group_size, initial_clip=1.0 ) def observe_activation_pre_hook( collector: OuroLASStatisticsCollector, module_key: str, inputs: tuple[torch.Tensor, ...], ) -> None: """Collect statistics while preserving the module's original inputs. PyTorch treats a non-``None`` forward-pre-hook return value as replacement inputs. ``observe`` intentionally returns the recurrence index for direct callers, so registering it as the callback would replace a tensor with an integer. """ collector.observe(module_key, inputs[0]) return None def load_pinned_teacher_student(device: str): """Load two independent BF16 Ouro models from the pinned local snapshot.""" from transformers import AutoModelForCausalLM, AutoTokenizer common = dict( pretrained_model_name_or_path=str(PINNED_OURO_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() for parameter in teacher.parameters(): parameter.requires_grad_(False) tokenizer = AutoTokenizer.from_pretrained( PINNED_OURO_SNAPSHOT, trust_remote_code=True, local_files_only=True ) return teacher, student, tokenizer def load_pile_texts(config: OuroCalibrationDataConfig) -> list[str]: from datasets import load_dataset dataset = load_dataset( PILE_DATASET, split=config.split, revision=PILE_DATASET_REVISION, streaming=True, ) texts = [] for row in dataset: text = row[config.text_field] if isinstance(text, str) and text.strip(): texts.append(text) if len(texts) == config.samples: break if len(texts) != config.samples: raise ValueError(f"requested {config.samples} texts but found {len(texts)}") return texts def _ste(original: torch.Tensor, quantized: torch.Tensor) -> torch.Tensor: """Quantized forward plus an identity activation-gradient path. Unlike ``original + (quantized-original).detach()``, this form deliberately retains gradients to learned LAS scales and transforms through the QDQ expression while adding the standard identity STE for rounding. """ return quantized + (original - original.detach()) def _weight_ste(original: torch.Tensor, quantized: torch.Tensor) -> torch.Tensor: """RTN weight forward with the standard identity transform gradient. Weight scales are inferred rather than learned LoopQ parameters. Detaching their QDQ graph avoids retaining full-matrix rounding intermediates while the identity path still optimizes every transform affecting ``original``. """ return original + (quantized - original).detach() def qdq_linear(value, module, transform, las, key, loop, bits, *, checkpoint=False, activation_ste="identity"): """Optionally recompute a deterministic linear during backward. Capture the exact loop transform now: routing counters and module hooks must never be replayed, especially after statistics-mode cleanup. """ def compute(x): transformed = transform(x) activation_q = las.quantize(key, loop, transformed, bits=bits, rounding_ste=activation_ste == "rounding").dequantized if activation_ste == "identity": activation_q = _ste(transformed, activation_q) folded = transform.fold_weight(module.weight) weight_q = quantize_weight(folded.detach()).dequantized return F.linear(activation_q, _weight_ste(folded, weight_q), module.bias) if checkpoint: from torch.utils.checkpoint import checkpoint as recompute return recompute(compute, value, use_reentrant=False, preserve_rng_state=False) return compute(value) def ouro_projection_sites(model: nn.Module) -> dict[str, tuple[nn.Module, ...]]: """Resolve all 24x4 paper groups without depending on concrete HF classes.""" sites: dict[str, tuple[nn.Module, ...]] = {} for layer_index, layer in enumerate(model.model.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"model.layers.{layer_index}.{group}"] = tuple(modules) return sites class OuroTrajectoryCapture: """Capture true recurrent pre-CTA states and feed CTA into the next loop.""" def __init__(self, norm: nn.Module, cta: CrossLoopTransitionAdapter | None) -> None: self.norm = norm self.cta = cta self.pre_cta: list[torch.Tensor] = [] self.adapted: list[torch.Tensor] = [] self._handle = None def __enter__(self) -> "OuroTrajectoryCapture": self.pre_cta.clear() self.adapted.clear() def hook(_module, _inputs, output): loop = len(self.pre_cta) if loop >= 4: raise RuntimeError("Ouro norm was invoked more than four recurrent loops") self.pre_cta.append(output) if loop < 3 and self.cta is not None: output = self.cta(output, loop) self.adapted.append(output) return output self._handle = self.norm.register_forward_hook(hook) return self def __exit__(self, exception_type, _exception, _traceback) -> None: self._handle.remove() self._handle = None # Do not hide the primary forward/backward failure with a secondary # trajectory-length assertion during cleanup. if exception_type is None and len(self.pre_cta) != 4: raise RuntimeError(f"expected four recurrent states, captured {len(self.pre_cta)}") class OuroDifferentiableQDQ(nn.Module): """Differentiable HF calibration hooks for every Ouro paper group. The backbone weights remain owned by the model and are never mutated. A selected group gets four independent transform copies; all other groups route through one shared transform. """ def __init__( self, *, model: nn.Module, las: LoopAwareActivationScales, activation_bits: int, factor_by_width: Mapping[int, tuple[int, int]], checkpoint_linears: bool = False, activation_ste: str = "identity", statistics_coordinate: str = "svd_parameters", ) -> None: super().__init__() self.model = model self.las = las self.activation_bits = activation_bits if activation_ste not in {"identity", "rounding"}: raise ValueError("unknown activation STE") self.activation_ste = activation_ste if statistics_coordinate not in {"svd_parameters", "effective_factors"}: raise ValueError("statistics coordinate must be svd_parameters or effective_factors") self.statistics_coordinate = statistics_coordinate self.checkpoint_linears = checkpoint_linears self.sites = ouro_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 feature 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: list[Any] = [] self._calls: dict[int, int] = defaultdict(int) self._records: dict[tuple[str, int], list[torch.Tensor]] = defaultdict(list) self._execution_views: dict[int, object] = {} 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(4)]) self.selected_loop_transforms[encoded] = loops def _base_transform_for(self, key: str, loop: int) -> SharedKroneckerTransform: encoded = self._encoded[key] if encoded in self.selected_loop_transforms: return self.selected_loop_transforms[encoded][loop] return self.shared_transforms[encoded] def transform_for(self, key: str, loop: int) -> SharedKroneckerTransform: encoded = self._encoded[key] if self._statistics_mode: return self.statistics_loop_transforms[encoded][loop] return self._base_transform_for(key, loop) def _build_statistics_transforms(self) -> None: transforms = {} for key in sorted(self.sites): copies = [] for loop in range(4): base = self._base_transform_for(key, loop) copy = (base.fresh_copy() if self.statistics_coordinate == "svd_parameters" else SharedKroneckerTransform.from_export_state(base.export_state()).to( device=base.left.device, dtype=base.left.dtype)) 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): call = self._calls[id(current)] self._calls[id(current)] += 1 if call >= 4: raise RuntimeError("projection invoked more than four times; loop routing cannot wrap") loop = call activation = inputs[0] transform_module = self.transform_for(key, loop) 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( activation, current, transform, self.las, key, loop, 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)] != 4] for modules in self.sites.values(): for module in modules: if hasattr(module, "_loopq_override"): del module._loopq_override self._statistics_mode = False self.statistics_loop_transforms = nn.ModuleDict() if validate and incomplete: raise RuntimeError(f"incomplete projection trajectories: {incomplete[:4]}") def sharing_gap_statistics(self, loss: torch.Tensor, *, fisher_loss: torch.Tensor | None = None) -> dict[str, SharingGapStatistics]: """Collect Eq.8 transform VJPs and their diagonal-Fisher estimate.""" if not self._statistics_mode: raise RuntimeError("sharing-gap statistics require statistics_mode") result = {} transforms = [ self.transform_for(key, loop) for key in sorted(self.sites) for loop in range(4) ] 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) } fisher_gradients = (all_gradients if fisher_loss is None else torch.autograd.grad( fisher_loss, all_parameters, retain_graph=True, allow_unused=True)) fisher_by_id = dict(zip(map(id, all_parameters), fisher_gradients)) for key in sorted(self.sites): per_loop = [] fisher_per_loop = [] for loop in range(4): transform = self.transform_for(key, loop) 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_loop.append(torch.cat([item.flatten() for item in accumulated]).cpu()) fisher_per_loop.append(torch.cat([ (fisher_by_id[id(parameter)].detach() if fisher_by_id[id(parameter)] is not None else torch.zeros_like(parameter)).flatten() for parameter in params ]).cpu().to(torch.float64)) gradients = torch.stack(per_loop) result[key] = SharingGapStatistics( gradients=gradients, # Square in FP64: finite FP32 VJPs can exceed sqrt(FP32_MAX). fisher_diagonal=torch.stack(fisher_per_loop).square().mean(dim=0), parameter_count=gradients.shape[1], ) return result def export_shared(self) -> dict[str, dict[str, object]]: return {key: self.shared_transforms[self._encoded[key]].export_state() for key in sorted(self.sites)} def export_selected(self) -> dict[str, dict[str, dict[str, object]]]: return { key: {str(loop): transform.export_state() for loop, 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 _extract_logits(output: Any) -> torch.Tensor: if hasattr(output, "logits"): return output.logits if isinstance(output, (tuple, list)): return output[0] raise TypeError("model output must expose .logits or place logits first") def ouro_trajectory_loss( *, teacher: nn.Module, student: nn.Module, student_qdq: OuroDifferentiableQDQ, cta: CrossLoopTransitionAdapter, inputs: Mapping[str, torch.Tensor] | torch.Tensor, step: int, mu_cache: AdaptiveMuCache, collect_statistics: bool = True, offload_saved_tensors: bool = False, saved_tensor_keep_bytes: int | None = None, asynchronous_offload: bool = False, kl_mode: str = "conditional_topk", fisher_estimator: str = "trajectory_gradient_square", ) -> tuple[TrajectoryLoss, dict[str, SharingGapStatistics]]: """Run one bounded true four-loop teacher/student calibration example.""" if fisher_estimator not in {"trajectory_gradient_square", "model_score_mc"}: raise ValueError("unknown Fisher estimator") call = (lambda model: model(inputs)) if isinstance(inputs, torch.Tensor) else ( lambda model: model(**inputs, use_cache=False, exit_at_step=3) ) with torch.no_grad(), OuroTrajectoryCapture(teacher.model.norm, None) as teacher_trace: teacher_output = call(teacher) teacher_hidden = torch.stack(teacher_trace.pre_cta) student_qdq.begin(statistics_mode=collect_statistics) failed = True try: saved_tensors = ( # Weight-QDQ internals are detached by _weight_ste, bounding this # exact saved-tensor set so pinned asynchronous restoration is # practical. The former differentiable full-weight quantizer graph # exceeded the pinned-host allocator and is not retained here. torch.autograd.graph.save_on_cpu(pin_memory=True, device_type="cuda") if offload_saved_tensors else torch.autograd.graph.saved_tensors_hooks( lambda tensor: tensor, lambda tensor: tensor ) ) if saved_tensor_keep_bytes is not None: if not offload_saved_tensors: raise ValueError("budgeted placement requires offload_saved_tensors") from .saved_tensors import BudgetedSavedTensors saved_tensors = BudgetedSavedTensors( (p for p in student.parameters() if not p.requires_grad), keep_bytes=saved_tensor_keep_bytes) if asynchronous_offload: if not offload_saved_tensors or saved_tensor_keep_bytes is not None: raise ValueError("asynchronous offload requires CPU placement without retention budget") from .saved_tensors import AsyncSavedTensors saved_tensors = AsyncSavedTensors() with saved_tensors: with OuroTrajectoryCapture(student.model.norm, cta) as student_trace: student_output = call(student) 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, kl_mode=kl_mode, ) fisher_loss = None if collect_statistics and fisher_estimator == "model_score_mc": # One Monte Carlo draw of the joint categorical outputs # at the fixed teacher-forced contexts. This is a Fisher # estimator, not the square of the trajectory-loss gradient. logits = _extract_logits(student_output).float() logp = logits.log_softmax(-1) labels = torch.multinomial(logp.detach().exp().reshape(-1, logits.shape[-1]), 1) fisher_loss = -logp.reshape(-1, logits.shape[-1]).gather(-1, labels).sum() statistics = ( student_qdq.sharing_gap_statistics(loss.total, fisher_loss=fisher_loss) if collect_statistics else {} ) failed = False finally: student_qdq.end(validate=not failed) return loss, statistics