File size: 5,683 Bytes
9118991 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 | """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
)
|