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