JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
5.68 kB
"""Symmetric group-wise RTN fake quantization used by LoopQ.
The paper defines the scalar quantizer in Eq. (4) as
Q(x; c) = c * clip(round(x / c), q_min, q_max).
It specifies group-wise W4A4/W4A8 quantization with group size 32, but does
not specify several low-level choices. Those choices are deliberately kept
small and explicit in ``configs/loopq_compatibility_decisions.yaml``.
This module performs QDQ/fake quantization. Integer packing and kernels are
outside LQ1's scope.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
SUPPORTED_BITS = frozenset({4, 8})
@dataclass(frozen=True)
class GroupwiseQuantizationResult:
"""Fake-quantized tensor and the group parameters used to produce it."""
dequantized: torch.Tensor
integers: torch.Tensor
scales: torch.Tensor
bits: int
group_size: int
axis: int
qmin: int
qmax: int
def _round_half_away_from_zero(x: torch.Tensor) -> torch.Tensor:
"""Deterministic RTN with half-way values rounded away from zero."""
return torch.copysign(torch.floor(torch.abs(x) + 0.5), x)
def _canonical_axis(ndim: int, axis: int) -> int:
if ndim == 0:
raise ValueError("group-wise quantization requires a non-scalar tensor")
if not -ndim <= axis < ndim:
raise IndexError(f"axis {axis} is out of bounds for tensor rank {ndim}")
return axis % ndim
def fake_quantize_groupwise(
tensor: torch.Tensor,
*,
bits: int,
group_size: int = 32,
axis: int = -1,
scales: torch.Tensor | None = None,
rounding_ste: bool = False,
) -> GroupwiseQuantizationResult:
"""Apply symmetric signed group-wise RTN quantization and dequantization.
Groups are consecutive chunks on ``axis``. A final short group is
calibrated independently and never padded into the returned tensors.
Unless explicit learned ``scales`` are supplied, each group uses
``max(abs(group)) / (2**(bits - 1) - 1)``. All-zero groups use scale 1,
producing exact zeros without NaNs.
"""
if bits not in SUPPORTED_BITS:
raise ValueError(f"bits must be one of {sorted(SUPPORTED_BITS)}, got {bits}")
if group_size <= 0:
raise ValueError(f"group_size must be positive, got {group_size}")
if not tensor.is_floating_point():
raise TypeError("tensor must have a floating-point dtype")
canonical_axis = _canonical_axis(tensor.ndim, axis)
moved = tensor.movedim(canonical_axis, -1)
width = moved.shape[-1]
if width == 0:
raise ValueError("quantized axis must not be empty")
qmax = 2 ** (bits - 1) - 1
# Match the cited FlatQuant symmetric RTN implementation: the signed
# two's-complement range includes its most-negative code.
qmin = -(2 ** (bits - 1))
group_count = (width + group_size - 1) // group_size
expected_scale_shape = moved.shape[:-1] + (group_count,)
work = moved.to(torch.float32)
# Batch groups into one tensor; padding is excluded from returned values
# and zero padding cannot change an absmax, including the short tail.
padding = group_count * group_size - width
grouped = torch.nn.functional.pad(work, (0, padding)).reshape(
*moved.shape[:-1], group_count, group_size
)
if scales is not None:
group_scales = scales.to(device=tensor.device, dtype=torch.float32)
if tuple(group_scales.shape) != tuple(expected_scale_shape):
raise ValueError(
f"scales shape must be {tuple(expected_scale_shape)}, "
f"got {tuple(group_scales.shape)}"
)
if not torch.isfinite(group_scales).all() or (group_scales <= 0).any():
raise ValueError("all scales must be finite and strictly positive")
normalized = grouped / group_scales.unsqueeze(-1)
else:
absmax = grouped.abs().amax(dim=-1)
group_scales = torch.where(absmax == 0, torch.ones_like(absmax), absmax / qmax)
safe_absmax = torch.where(absmax == 0, torch.ones_like(absmax), absmax)
normalized = (grouped / safe_absmax.unsqueeze(-1)) * qmax
rounded = _round_half_away_from_zero(normalized)
if rounding_ste:
# Exact rounded forward; unit rounding derivative before clipping.
# This also differentiates x/scale, giving the scale gradient its
# quantization-residual term instead of just the integer code.
rounded = rounded.detach() + (normalized - normalized.detach())
integer = rounded.clamp(qmin, qmax)
integers = integer.flatten(-2)[..., :width].movedim(-1, canonical_axis).to(torch.int8)
dequantized = (integer * group_scales.unsqueeze(-1)).flatten(-2)[..., :width].movedim(-1, canonical_axis)
dequantized = dequantized.to(dtype=tensor.dtype)
return GroupwiseQuantizationResult(
dequantized=dequantized,
integers=integers,
scales=group_scales,
bits=bits,
group_size=group_size,
axis=canonical_axis,
qmin=qmin,
qmax=qmax,
)
def quantize_weight(weight: torch.Tensor, *, bits: int = 4) -> GroupwiseQuantizationResult:
"""LoopQ paper-contract weight QDQ: symmetric RTN, group size 32."""
if bits != 4:
raise ValueError("LoopQ's evaluated weight precision is 4 bits")
return fake_quantize_groupwise(weight, bits=bits, group_size=32, axis=-1)
def quantize_activation(
activation: torch.Tensor,
*,
bits: int,
scales: torch.Tensor | None = None,
) -> GroupwiseQuantizationResult:
"""LoopQ paper-contract activation QDQ for A4 or A8."""
return fake_quantize_groupwise(
activation, bits=bits, group_size=32, axis=-1, scales=scales
)