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