"""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 )