"""Loop-aware activation clipping (LAS) for the isolated LoopQ reproduction. LoopQ Section 4.1 replaces a shared activation range with one range scalar per loop and module. It clips dynamic per-token, per-group absmax ranges; the paper explicitly counts only ``O(TL)`` added scalars. """ from __future__ import annotations import math from collections.abc import Mapping, Sequence from typing import Any import torch from torch import nn from .quantization import GroupwiseQuantizationResult, fake_quantize_groupwise OURO_LOOP_COUNT = 4 HUGINN_LOOP_COUNT = 32 class LoopAwareActivationScales(nn.Module): """Positive activation scales routed by module name and loop index. Scales are represented by trainable log-scales so calibration cannot make them non-positive. Module names are kept in a deterministic sorted tuple; a ``ParameterList`` avoids imposing PyTorch ParameterDict naming rules on architecture module paths. """ FORMAT_VERSION = 2 def __init__( self, scales: Mapping[str, torch.Tensor], *, loop_count: int, group_size: int = 32, shared_across_loops: bool = False, dynamic_clip: bool = False, ) -> None: super().__init__() if loop_count <= 0: raise ValueError("loop_count must be positive") if group_size <= 0: raise ValueError("group_size must be positive") if not scales: raise ValueError("at least one module scale tensor is required") self.shared_across_loops = bool(shared_across_loops) self.dynamic_clip = bool(dynamic_clip) self.loop_count = int(loop_count) self.group_size = int(group_size) self.module_names = tuple(sorted(scales)) self._module_to_index = {name: index for index, name in enumerate(self.module_names)} parameters = [] self._group_counts: dict[str, int] = {} for name in self.module_names: if not name: raise ValueError("module names must be non-empty") value = torch.as_tensor(scales[name], dtype=torch.float32) if value.ndim != 2 or value.shape[0] != self.loop_count or value.shape[1] == 0: raise ValueError( f"scales[{name!r}] must have shape " f"({self.loop_count}, groups), got {tuple(value.shape)}" ) if not torch.isfinite(value).all() or (value <= 0).any(): raise ValueError(f"scales[{name!r}] must be finite and strictly positive") self._group_counts[name] = value.shape[1] if self.dynamic_clip and value.shape[1] != 1: raise ValueError("dynamic LAS stores one scalar per module and loop") if self.shared_across_loops: value = value.amax(dim=0, keepdim=True) parameters.append(nn.Parameter(value.log())) self.log_scales = nn.ParameterList(parameters) @classmethod def from_calibration( cls, calibration: Mapping[str, Sequence[torch.Tensor] | torch.Tensor], *, loop_count: int, bits: int, group_size: int = 32, ) -> "LoopAwareActivationScales": """Initialize scales from activations using LQ1's absmax estimator. A module value may be a sequence containing one tensor per loop, or a tensor whose first dimension is the loop dimension. All remaining leading dimensions are calibration observations; groups are formed on the last dimension and reduced across every observation. """ if bits not in (4, 8): raise ValueError("LoopQ activation bits must be 4 or 8") if not calibration: raise ValueError("calibration mapping must not be empty") qmax = 2 ** (bits - 1) - 1 initialized: dict[str, torch.Tensor] = {} for module_name, module_samples in calibration.items(): samples = cls._split_loop_samples(module_name, module_samples, loop_count) widths = {sample.shape[-1] for sample in samples if sample.ndim > 0} if len(widths) != 1 or any(sample.ndim == 0 for sample in samples): raise ValueError( f"calibration[{module_name!r}] loop tensors must be non-scalar " "and have a common last dimension" ) width = next(iter(widths)) if width == 0: raise ValueError(f"calibration[{module_name!r}] feature dimension is empty") per_loop = [] for sample in samples: if not sample.is_floating_point(): raise TypeError(f"calibration[{module_name!r}] must be floating point") work = sample.detach().to(torch.float32) if not torch.isfinite(work).all(): raise ValueError(f"calibration[{module_name!r}] must be finite") group_scales = [] for start in range(0, width, group_size): absmax = work[..., start : start + group_size].abs().amax() group_scales.append( torch.where(absmax == 0, torch.ones_like(absmax), absmax / qmax) ) per_loop.append(torch.stack(group_scales)) initialized[module_name] = torch.stack(per_loop) return cls(initialized, loop_count=loop_count, group_size=group_size) @classmethod def dynamic_for_modules( cls, module_names: Sequence[str], *, loop_count: int, group_size: int = 32, initial_clip: float = 1.0, ) -> "LoopAwareActivationScales": """Create paper-counted per-module/per-loop dynamic clipping factors.""" if not 0 < initial_clip <= 1: raise ValueError("initial clipping factor must be in (0, 1]") scales = { name: torch.full((loop_count, 1), initial_clip) for name in module_names } return cls(scales, loop_count=loop_count, group_size=group_size, dynamic_clip=True) @staticmethod def _split_loop_samples( module_name: str, value: Sequence[torch.Tensor] | torch.Tensor, loop_count: int, ) -> tuple[torch.Tensor, ...]: if isinstance(value, torch.Tensor): if value.ndim < 2 or value.shape[0] != loop_count: raise ValueError( f"calibration[{module_name!r}] tensor must start with loop " f"dimension {loop_count}, got {tuple(value.shape)}" ) return tuple(value.unbind(0)) samples = tuple(value) if len(samples) != loop_count: raise ValueError( f"calibration[{module_name!r}] must contain {loop_count} loop tensors, " f"got {len(samples)}" ) if not all(isinstance(sample, torch.Tensor) for sample in samples): raise TypeError(f"calibration[{module_name!r}] entries must be tensors") return samples @classmethod def for_ouro( cls, calibration: Mapping[str, Sequence[torch.Tensor] | torch.Tensor], *, bits: int ) -> "LoopAwareActivationScales": return cls.from_calibration(calibration, loop_count=OURO_LOOP_COUNT, bits=bits) @classmethod def for_huginn( cls, calibration: Mapping[str, Sequence[torch.Tensor] | torch.Tensor], *, bits: int ) -> "LoopAwareActivationScales": return cls.from_calibration(calibration, loop_count=HUGINN_LOOP_COUNT, bits=bits) def scales_for(self, module_name: str, loop_index: int) -> torch.Tensor: """Return the positive scale vector for exactly one recurrent loop.""" if module_name not in self._module_to_index: raise KeyError(f"unknown LAS module {module_name!r}") if not 0 <= loop_index < self.loop_count: raise IndexError( f"loop_index must be in [0, {self.loop_count}), got {loop_index}" ) raw = self.log_scales[self._module_to_index[module_name]][ 0 if self.shared_across_loops else loop_index ].exp() if not self.dynamic_clip: return raw clipped = raw.clamp(max=1.0) return raw + (clipped - raw).detach() def quantize( self, module_name: str, loop_index: int, activation: torch.Tensor, *, bits: int, rounding_ste: bool = False, ) -> GroupwiseQuantizationResult: """Route LAS scales and apply the LQ1 activation quantizer.""" scales = self.scales_for(module_name, loop_index) expected_groups = math.ceil(activation.shape[-1] / self.group_size) if not self.dynamic_clip and expected_groups != scales.numel(): raise ValueError( f"activation for {module_name!r} requires {expected_groups} groups, " f"but LAS stores {scales.numel()}" ) if self.dynamic_clip: work = activation.detach().to(torch.float32) padding = expected_groups * self.group_size - activation.shape[-1] grouped = torch.nn.functional.pad(work, (0, padding)).reshape( *activation.shape[:-1], expected_groups, self.group_size) qmax = 2 ** (bits - 1) - 1 absmax = grouped.abs().amax(dim=-1) base = torch.where(absmax == 0, torch.ones_like(absmax), absmax / qmax) expanded = base * scales.reshape((1,) * (activation.ndim - 1) + (1,)) else: expanded = scales.reshape((1,) * (activation.ndim - 1) + (-1,)).expand( activation.shape[:-1] + (expected_groups,)) return fake_quantize_groupwise( activation, bits=bits, group_size=self.group_size, axis=-1, scales=expanded, rounding_ste=rounding_ste, ) def export_state(self) -> dict[str, Any]: """Return a deterministic, torch-save-safe LAS artifact.""" return { "format": "loopq_las", "format_version": self.FORMAT_VERSION, "loop_count": self.loop_count, "group_size": self.group_size, "shared_across_loops": self.shared_across_loops, "dynamic_clip": self.dynamic_clip, "module_names": list(self.module_names), "log_scales": {name: self.log_scales[index].detach().cpu().clone() for index, name in enumerate(self.module_names)}, "scales": { name: torch.stack([ self.scales_for(name, loop).detach() for loop in range(self.loop_count) ]).cpu() for index, name in enumerate(self.module_names) }, } @classmethod def from_export_state(cls, state: Mapping[str, Any]) -> "LoopAwareActivationScales": required = {"format", "format_version", "loop_count", "group_size", "module_names", "scales"} missing = required.difference(state) if missing: raise ValueError(f"LAS export is missing fields: {sorted(missing)}") if state["format"] != "loopq_las" or state["format_version"] not in (1, cls.FORMAT_VERSION): raise ValueError("unsupported LAS export format or version") names = list(state["module_names"]) scales = state["scales"] if names != sorted(names) or set(names) != set(scales): raise ValueError("LAS export module_names must be sorted and match scales") module = cls(scales, loop_count=int(state["loop_count"]), group_size=int(state["group_size"]), shared_across_loops=bool(state.get("shared_across_loops", False)), dynamic_clip=bool(state.get("dynamic_clip", False))) # Preserve trained parameters exactly: log(exp(log_scale)) is not an # exact floating-point identity and can change RTN bin assignments. raw = state.get("log_scales") if raw is not None: if set(raw) != set(names): raise ValueError("LAS raw log-scale module names do not match") with torch.no_grad(): for index, name in enumerate(names): value = torch.as_tensor(raw[name], dtype=torch.float32) parameter = module.log_scales[index] if value.shape != parameter.shape or not torch.isfinite(value).all(): raise ValueError("LAS raw log-scale shape or finiteness mismatch") expected = torch.as_tensor(scales[name], dtype=torch.float32) applied = value.exp().clamp(max=1.0) if module.dynamic_clip else value.exp() if not torch.allclose(applied.expand_as(expected), expected, rtol=1e-6, atol=0): raise ValueError("LAS scale and raw log-scale payloads disagree") parameter.copy_(value) return module