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