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