From e40315909e24069b680dff0663b2a23df02e3463 Mon Sep 17 00:00:00 2001 From: Chenjie Luo Date: Mon, 28 Sep 2026 17:05:57 +0000 Subject: [PATCH] Add the IQ2_S codec First of two changes adding IQ2_S, the widest of the GGML IQ formats at one and two bits at 2.5625 bits per weight. This one lands the PyTorch codec: the encoder and decoder and the 1024-entry codebook. It is not registered yet, so no quantizer dispatches to it and the ggml package does not export it; the next change adds the CUDA encoder, registers the format and adds its recipe. The interesting difference from IQ2_XS and IQ2_XXS is the sign handling: IQ2_S stores a full eight-bit sign mask per group rather than a seven-bit parity-coded index, so the encoder takes the input signs as they are instead of flipping the weakest element to fix parity, and the search compares magnitudes directly. The decoder is validated against llama.cpp's own output over 2,355,200 blocks from unsloth/Qwen3.8-27B-GGUF, all bit-identical to dequantize_row_iq2_s, and the new codebook matches the ggml-common.h table entry for entry. Blocks from that checkpoint ship as conformance vectors. Because a codec can now land before it is registered, the two IQ tests that go through TensorQuantizer iterate the registry rather than every codec, and the battery's coverage check requires each registered format to be listed instead of requiring the two lists to be equal, which is what its docstring already said. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Chenjie Luo --- modelopt/torch/quantization/ggml/codebooks.py | 33 +++ modelopt/torch/quantization/ggml/iq2_s.py | 217 ++++++++++++++++++ .../quantization/iq_llama_cpp_vectors.py | 43 ++++ .../torch/quantization/test_iq_formats.py | 14 +- 4 files changed, 303 insertions(+), 4 deletions(-) create mode 100644 modelopt/torch/quantization/ggml/iq2_s.py diff --git a/modelopt/torch/quantization/ggml/codebooks.py b/modelopt/torch/quantization/ggml/codebooks.py index 5e015cd3765..22b14decb17 100644 --- a/modelopt/torch/quantization/ggml/codebooks.py +++ b/modelopt/torch/quantization/ggml/codebooks.py @@ -178,6 +178,33 @@ "Xf5hzoD21KEk6CpbkT/vBJbU" ) +# Compact byte representation of the canonical [1024, 8] IQ2_S grid. +_IQ2_S_GRID_ZLIB_B64 = ( + "eNqFmWF25DAIg//6CjoD979fG5sPI2bS7r48tttJYoOQhGet8ycySv4fkVFL9oHIqPAbIqPkD9DKB+SDIuPvx+zBkVHyF2mFvTAy" + "KnwBkfH3LluQlr4u7PfHc1/kz3mx4Mj43NU3EBmfZfYNacXXjUXG5+N9o1qyDUdGhSdAykRE/rwyISMxkVHKuOJrorieh/XEaelr" + "Ap/0RKalJzQyPtvtCX620RP9PO48xxMf6yRePflsjpe1gkTGnZ1WoL3aVqjIuN/SChcZd9laISOjlqywkXGnvxU6Mu40t6UBgMio" + "cEDstJx151Yi7wsDDLXhAkCRUUsGqMi4l9MAtl+nk3i11AG8yLg/1oAYGSUHplZ8BWhklBKwkT8vfQUwpaPWXAoHtuQAlxzoWnLA" + "JzPsVP3RAGCe3HIBktU3Ffny1igb3XH6IqGn3kCRcVe5NVRk3NmMzF7kbiN32Rpv/zayrSJ5qkF9NmRk3LBtDbrh2Bp0wyUSJpFw" + "sMYN9VbaC2sNrSVr7FhhDR6j0Z+7eysuIy/VHpaxqAqsgEXGQsrFXYKgFJIThgz9oBh0XiKJ1VEFii51LKvWJRzJiUeZpUlAQERy" + "QiKn/IWgIuOhmUtYkfHQySUw7iQFkfHQwCW2yHja+xIdb5yEB6RPO14CZGUQISsk9WCSmh7Y3x1MgjywvTuDMNnhG2HSYnAiPUxG" + "VicnSKQRK00BwZLB1Tddi7vEGxlPGi8Rk/lJyFAAt5403ApNwqZiq3Ehvd2JnNxyQexUehI8FMTSIXwQMYkfhLwJAFuHM0HSFASQ" + "9SYMD7915NGaoWY2mnCwVy6QunqTxi3JTlRD8hWYMGQXFTRKVispmkAHTGGiI6ZAPSp0nhd/CtWT3C5YSAJQQovouNVFt1zgFTbI" + "PtJ10aFT8AAloJjCd9TnCmCkANLhUwjp9cLickFE4qollr4KJJbqTShLg/LvWWikb42UmC6k+Yna0vkkTAS1h5q5bYJLT2FaEPu6" + "7hsv+TaBhkymUJ/kd8F2Zw1FUAKYcgo6zDmFPXbbXoHHil6hlwm+Nlyv8Bc1rWZGGzNPY3CTfH6AuUsadzridEf3DmAV43AfsaN8" + "LGUMZegog0FF5WMkvVXKsdzklBEBCfLxDvEpxVlOipBaGReQVGTg41YZGopFMUhmJcvHopIGWk8+xqDic0xBdVVY8bGiqE4+NhTy" + "5eNAKes0VHQIbAu7yu17GS25Da+t0WHSMGAyO1zSSBfJ7W0ZNLmNLYMGGuR2s6Aqt5OlZXKbWJ0vt4OV4+W2rlp5uV3DnpURlNuu" + "MoQwy7BX006VUZTbo5Iwuf0pppLbmuLsaSypgdx+1COm4dxq35zTmwGFIeUyX0uaxlQuz60FNtC2unZnRibl8liMjEfCQ+DkpsEF" + "e0PuprzdE4IXIwzzy+WlUl7FyUXK5eDjpCFpfNL3q6FGceQ0WaVeTnOT1kpr5TRWXLKcliYN1Zbk9FKmbRp4FBHzgLWQ00CR9fK2" + "LZJY3o4FWkAmb6tKBi+Rt0ddc0CQw/djUJDDrFppDg5yWFTuHoHrg4S8fGVJ5WUp6MjTXZ5sDh6wnDw9tec5kMi38zGYgBr5a4s6" + "5sAi//XH4MI/5gATW5E/BxmcETdiQXnAHHR44Bx4zjBzTTkv+m8A2q72vD/6QMQCl51+3wGJVmfhzDRsYA5ObAguYmPvA5WfzGFS" + "SACiTCLmwEVi3gYv7S19HcCaa82EtMGMxDISkOA5qJHwt4GNQrwNblhUigSlVjVOx1UhS2nS2nDvTl0b9LaXjeNJesGXf/3C1y01" + "CAIEKIOlMKsVkpd9XVGDItwAYMAqNeGaAEIhARJKiblafpxeQKO0AK7Ey4+1C4iQHICk6bkKXFnkI+d3UKVD83i2gAylAuhS8NQQ" + "AL76SN4Av7pktgYAsjTC8mO9agw0nwYpLswlsrc7KPsJMmpKQ6F2NFZJRG6VRnsbrCObiuLPQVun2z4G7kLNsmODatzaba6SSZCG" + "LseTZmGTYBsXD4ceDvQj5TACwDExvi0fz4og3gZ5JlGIg5VBIHPQB3NvA/8kGqQZwpmDf9rfIqK3AwCYYvUvLa/9qkxCYPOgAEJ7" + "OzA4hwGX8P47QIAIIZs4m9+02IlxHjDESWXNuDrVrOE7V1EnAkmLdQABwfKtLXRW31LmlxAQ7zyoKCgdZihzelIZ9ztJjhxWPyWK" + "PCXi5nm6cf8Jwb8deBykxb2yrCUlmYaoL1/rjOq8XmeEyR6tI6z6mK+2LS0/ln9+AIHIjsQ=" +) + @cache def iq1_s_grid_bytes() -> bytes: @@ -195,3 +222,9 @@ def iq2_xs_grid_bytes() -> bytes: def iq2_xxs_grid_bytes() -> bytes: """Decoded bytes of the [256, 8] IQ2_XXS magnitude table.""" return zlib.decompress(base64.b64decode(_IQ2_XXS_GRID_ZLIB_B64)) + + +@cache +def iq2_s_grid_bytes() -> bytes: + """Decoded bytes of the [1024, 8] IQ2_S magnitude table.""" + return zlib.decompress(base64.b64decode(_IQ2_S_GRID_ZLIB_B64)) diff --git a/modelopt/torch/quantization/ggml/iq2_s.py b/modelopt/torch/quantization/ggml/iq2_s.py new file mode 100644 index 00000000000..024965af048 --- /dev/null +++ b/modelopt/torch/quantization/ggml/iq2_s.py @@ -0,0 +1,217 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""IQ2_S fake quantization and GGML-compatible block packing. + +The encoder performs a single-pass squared-error grid search at a fixed, +empirically anchored super-block scale, mirroring :mod:`.iq2_xs`. Every 256 +logical values become one 82-byte block_iq2_s payload: + +* bytes 0..1: little-endian FP16 super-block scale d +* bytes 2..33: 32 low bytes of the grid index, four per sub-block +* bytes 34..65: 32 sign masks, four per sub-block +* bytes 66..73: eight bytes holding the grid index high 2 bits, four per byte +* bytes 74..81: 16 four-bit local scales, two per byte + +IQ2_S stores a full eight-bit sign mask per group rather than the seven-bit +parity-coded index used by IQ2_XS and IQ2_XXS, so the encoder can take the +input signs directly instead of flipping the weakest element to fix parity. + +The canonical 1024 x 8 magnitude grid lives in :mod:`.codebooks`, carried from +llama.cpp ggml-common.h revision 9b05354ec6fb58b4e665e9a39ebc40285c015638. +The matching dequantization formula is in ggml-quants.c at the same revision: +https://github.com/ggml-org/llama.cpp/blob/9b05354ec6fb58b4e665e9a39ebc40285c015638/ggml/src/ggml-quants.c#L2540-L2571 +""" + +import torch + +from .codebooks import iq2_s_grid_bytes +from .common import ( + GGML_BLOCK_SIZE, + narrow_to_float32, + validate_block_chunk_size, + validate_packed_weights, + validate_weight, +) + +__all__ = [ + "IQ2_S_BLOCK_BYTES", + "IQ2_S_BLOCK_SIZE", + "IQ2_S_EFFECTIVE_BITS", + "dequantize_iq2_s", + "iq2_s_grid", + "quantize_iq2_s", +] + +IQ2_S_BLOCK_SIZE = GGML_BLOCK_SIZE +IQ2_S_BLOCK_BYTES = 82 +IQ2_S_EFFECTIVE_BITS = IQ2_S_BLOCK_BYTES * 8 / IQ2_S_BLOCK_SIZE +_IQ2_S_GRID_ENTRIES = 1024 +_IQ2_S_LOCAL_SCALES = 16 +_IQ2_S_GROUPS = 32 +_IQ2_S_SUBBLOCKS = 8 +_IQ2_S_NATIVE_MAX = 43 * 31 / 8 +_IQ2_S_SCALE_ANCHOR_MIN = 0.65 +_IQ2_S_SCALE_ANCHOR_MAX = 0.92 +_IQ2_S_PEAK_TO_RMS_TAPER = 0.035 +# The grid is twice IQ2_XS's, so the same search tile costs twice the memory. +_DEFAULT_BLOCK_CHUNK_SIZE = 128 +_DEFAULT_DECODE_CHUNK_SIZE = 4096 + +_GRID_CACHE: dict[torch.device, torch.Tensor] = {} + + +def iq2_s_grid(device: torch.device | str | None = None) -> torch.Tensor: + """Return the canonical IQ2_S magnitude grid as float32.""" + resolved_device = torch.device(device or "cpu") + if resolved_device.type == "cuda" and resolved_device.index is None: + resolved_device = torch.device("cuda", torch.cuda.current_device()) + if resolved_device not in _GRID_CACHE: + values = torch.tensor(list(iq2_s_grid_bytes()), dtype=torch.float32) + _GRID_CACHE[resolved_device] = values.reshape(_IQ2_S_GRID_ENTRIES, 8).to( + device=resolved_device + ) + return _GRID_CACHE[resolved_device] + + +def _predict_iq2_s_scales(blocks: torch.Tensor) -> torch.Tensor: + """Predict one FP16 super-block scale for each flattened block.""" + x = narrow_to_float32(blocks) + amax = x.abs().amax(dim=1) + rms = x.square().mean(dim=1).sqrt() + peak_to_rms = torch.where(rms > 0, amax / rms, torch.zeros_like(rms)) + anchor_ratio = (1.0 - _IQ2_S_PEAK_TO_RMS_TAPER * peak_to_rms).clamp( + _IQ2_S_SCALE_ANCHOR_MIN, _IQ2_S_SCALE_ANCHOR_MAX + ) + return ((amax / _IQ2_S_NATIVE_MAX) * anchor_ratio).clamp(max=65504.0).to(torch.float16) + + +def _encode_blocks(blocks: torch.Tensor, grid: torch.Tensor) -> torch.Tensor: + """Encode a moderate-size batch of flattened 256-value blocks.""" + x = narrow_to_float32(blocks) + block_count = x.shape[0] + vectors = x.reshape(block_count, _IQ2_S_GROUPS, 8) + magnitudes = vectors.abs() + negative = vectors < 0 + + d = _predict_iq2_s_scales(x) + d_float = d.float() + + xnorm = vectors.square().sum(dim=-1) + qnorm = grid.square().sum(dim=-1) + shape = (block_count, _IQ2_S_GROUPS, _IQ2_S_LOCAL_SCALES) + best_error = torch.full(shape, torch.inf, dtype=torch.float32, device=x.device) + best_entry = torch.zeros(shape, dtype=torch.int64, device=x.device) + # All eight signs are storable, so the search compares magnitudes directly. + for entry_start in range(0, _IQ2_S_GRID_ENTRIES, 64): + grid_tile = grid[entry_start : entry_start + 64] + dot = (magnitudes.unsqueeze(2) * grid_tile.reshape(1, 1, -1, 8)).sum(dim=-1) + tile_qnorm = qnorm[entry_start : entry_start + 64].reshape(1, 1, -1) + + for local in range(_IQ2_S_LOCAL_SCALES): + scale = d_float.reshape(-1, 1, 1) * ((2 * local + 1) / 8.0) + error = ( + xnorm.unsqueeze(-1) - 2.0 * scale * dot + scale.square() * tile_qnorm + ).clamp_min_(0) + tile_error, tile_index = error.min(dim=-1) + replace = tile_error < best_error[:, :, local] + best_error[:, :, local] = torch.where(replace, tile_error, best_error[:, :, local]) + best_entry[:, :, local] = torch.where( + replace, tile_index + entry_start, best_entry[:, :, local] + ) + + # One local scale covers two groups (16 values), as in IQ2_XS. + pair_error = best_error.reshape(block_count, 16, 2, _IQ2_S_LOCAL_SCALES).sum(dim=2) + selected_local = pair_error.argmin(dim=-1) + group_local = selected_local.repeat_interleave(2, dim=1) + selected_entry = best_entry.gather(2, group_local.unsqueeze(-1)).squeeze(-1) + + sign_bits = torch.arange(8, dtype=torch.int64, device=x.device) + sign_mask = (negative.to(torch.int64) << sign_bits).sum(dim=-1) + + packed = torch.empty((block_count, IQ2_S_BLOCK_BYTES), dtype=torch.uint8, device=x.device) + packed[:, :2] = d.contiguous().view(torch.uint8).reshape(block_count, 2) + packed[:, 2:34] = (selected_entry & 0xFF).to(torch.uint8) + packed[:, 34:66] = sign_mask.to(torch.uint8) + high = (selected_entry >> 8).reshape(block_count, _IQ2_S_SUBBLOCKS, 4) + packed[:, 66:74] = ( + high[:, :, 0] | (high[:, :, 1] << 2) | (high[:, :, 2] << 4) | (high[:, :, 3] << 6) + ).to(torch.uint8) + packed[:, 74:] = (selected_local[:, 0::2] | (selected_local[:, 1::2] << 4)).to(torch.uint8) + return torch.where((d_float == 0).unsqueeze(1), 0, packed) + + +@torch.no_grad() +def quantize_iq2_s( + weight: torch.Tensor, *, block_chunk_size: int = _DEFAULT_BLOCK_CHUNK_SIZE +) -> tuple[torch.Tensor, torch.Tensor]: + """Pack a floating-point weight into GGML-compatible IQ2_S blocks. + + Returned shapes are ``[*weight.shape[:-1], weight.shape[-1] // 256, 82]`` + and ``[weight.ndim]``. + """ + validate_weight(weight, "IQ2_S") + validate_block_chunk_size(block_chunk_size) + + logical_shape = torch.tensor(weight.shape, dtype=torch.int64) + blocks = weight.contiguous().reshape(-1, IQ2_S_BLOCK_SIZE) + grid = iq2_s_grid(weight.device) + packed_shape = (*weight.shape[:-1], weight.shape[-1] // IQ2_S_BLOCK_SIZE, IQ2_S_BLOCK_BYTES) + chunks = [ + _encode_blocks(blocks[start : start + block_chunk_size], grid) + for start in range(0, blocks.shape[0], block_chunk_size) + ] + return torch.cat(chunks).reshape(packed_shape), logical_shape + + +@torch.no_grad() +def dequantize_iq2_s( + packed_weights: torch.Tensor, + weight_shape: torch.Tensor, + *, + dtype: torch.dtype = torch.bfloat16, + block_chunk_size: int = _DEFAULT_DECODE_CHUNK_SIZE, +) -> torch.Tensor: + """Decode GGML-compatible IQ2_S payload bytes.""" + shape = validate_packed_weights( + packed_weights, weight_shape, block_bytes=IQ2_S_BLOCK_BYTES, format_name="IQ2_S" + ) + validate_block_chunk_size(block_chunk_size) + + blocks = packed_weights.contiguous().reshape(-1, IQ2_S_BLOCK_BYTES) + bit_positions = torch.arange(8, dtype=torch.int64, device=blocks.device) + high_shifts = torch.tensor([0, 2, 4, 6], dtype=torch.int64, device=blocks.device) + grid = iq2_s_grid(blocks.device) + decoded = torch.empty((blocks.shape[0], IQ2_S_BLOCK_SIZE), dtype=dtype, device=blocks.device) + for start in range(0, blocks.shape[0], block_chunk_size): + stop = min(start + block_chunk_size, blocks.shape[0]) + block_chunk = blocks[start:stop] + count = block_chunk.shape[0] + d = block_chunk[:, :2].contiguous().view(torch.float16).reshape(-1).float() + low = block_chunk[:, 2:34].to(torch.int64).reshape(count, _IQ2_S_SUBBLOCKS, 4) + sign_mask = block_chunk[:, 34:66].to(torch.int64).reshape(count, _IQ2_S_SUBBLOCKS, 4) + qh = block_chunk[:, 66:74].to(torch.int64).reshape(count, _IQ2_S_SUBBLOCKS) + scale_bytes = block_chunk[:, 74:].to(torch.int64) + + entries = low | (((qh.unsqueeze(-1) >> high_shifts) & 0x3) << 8) + local = torch.empty((count, 16), dtype=torch.int64, device=blocks.device) + local[:, 0::2] = scale_bytes & 0x0F + local[:, 1::2] = scale_bytes >> 4 + scales = d.unsqueeze(-1) * (0.5 + local.float()) * 0.25 + signs = 1.0 - 2.0 * ((sign_mask.unsqueeze(-1) >> bit_positions) & 1).float() + values = grid[entries] * signs + chunk_decoded = values.reshape(count, 16, 2, 8) * scales.unsqueeze(-1).unsqueeze(-1) + decoded[start:stop] = chunk_decoded.reshape(-1, IQ2_S_BLOCK_SIZE) + return decoded.reshape(shape) diff --git a/tests/_test_utils/torch/quantization/iq_llama_cpp_vectors.py b/tests/_test_utils/torch/quantization/iq_llama_cpp_vectors.py index 6fe3a9ccd06..693a80179b4 100644 --- a/tests/_test_utils/torch/quantization/iq_llama_cpp_vectors.py +++ b/tests/_test_utils/torch/quantization/iq_llama_cpp_vectors.py @@ -141,6 +141,49 @@ "/qrzW/DbKMq7yFtqTG/oGeZ9U/1/TYe49w==" ), }, + "iq2_s": { + "source": "blk.35.attn_output.weight", + "block_bytes": 82, + "blocks": ( + "eNoB7AET/q4M+OCIlAAzLJjmBQAJuLQHGSwziwAmnlIDogF04URRrgOutsGuCoscGISHWf5Z7PnTZCXDydppD389LinF" + "F+K8HgxMCGQcABGwhFvXRv1Y11zFCvZOQw2g3hUdMb0vDvZmARjtM7DOQZ5kTxxzz/H+AAeVcDR0leZIoJDxcfhrr7Uv" + "kpvnsBWJYWuM7FOcANKeHe00igDNHJJDAYl2po2NiMp/XwiRcq1apgHtYw6m6g/Aao/ThGi7zgieTN49ZMg3FACEA/WD" + "VoudZA4dZAYThwSiFspPvOiNjuUxu11ujtJevf5VIolkIiygEQJrv56qq32nzw0KowLwPnr0CswZwQKXCL5JdIKYjwKX" + "nkrYcesAH00mlyNWkgcISZ9synXVqLyJ7IuNpgFD4XvoXfHtkou3P9FjihCFiMAeQwlxiWa7act/5XglDa3SmBz14UIO" + "vwsbhiMZ3oIAM4AAgmkAAXRzAAlEAHWHs6+UevSuXXqwriyP2frZUU5JYkM/84aPBbjS5XTjweCoMK6QHAYBoCQ0M2bP" + "hbc+UA3t0EhwrUstChnOAAkLEQRd7in2AFXGGaCDeNvvEgC4IkeqzF+vBg0/wJbAXSu4r42iDMm+Oqdr4+mcwHeTesEG" + "wEgAQgMAVoImQ/hldEZUNmfm3b4=" + ), + "expected": ( + "eNp1V22IXtURPjbWajDtVhJZxKQvRUStP9ZoNGZm263YmlIDi0K0fq62frIFNVASpGQtrYa0DYuCFcMmK5WgIrqiqNmZ" + "VxfxYzERV1EjiLBGhFhFFiWQH8F05tyZc+fevPnx8DzznDnnnjNzzsvuzMQinJlYxEM33cxp41kkGiMLwDh7mqta8jHP" + "OX0nCUAx+c6HOLN9Y4b5ZDErNLYctjmULt0HgjZH0LHGpx7r66YD68BAQ592WBG8Nmj+qq15XBiz96d/wdB116Iwjexi" + "5YbnkJjVM99zOL1wmARgoBA3/IXZs7thvMbVG2jm9F+iMJjmzub/akzmZZjPLa09YesLmubQr9JX6yO0+xj3ODn6BiqO" + "cY72OUHONJhO+SkrOk/vR2F0Hli4pOtjGhvautwDvzONu1EDgqf3rL531ndlA0iMPXpOPVD1sOpno9d+F4wpAAL3rIv1" + "e7BVq67VNueKru7O4S8zhlct7wqDs2N+74toOdFjnxfOG++/atRz6z0PXuN95Hk7x0DxxIdPZ33XobUsTMYae456KHms" + "Y8U/51YHrX78UVR2aDw3dz77uGgfL3Om/jaK4/edxOkHAyQA0ahaPfMh+paPNo4DT67AdOrvYfaeLUULKLL6DsljG1Of" + "7VsY4Hsh34d+X79r+2DRZb/p7G3Q/6NXUbmvcwMrRFOEeeAcxiC99SJFbNyzDU2DaG5psJysb99+kKXGHGvudfUaa81D" + "D0ofSm/qbzfWDvugVk7Zs34/fAO9v7qu34XY/8iWC7p/1743j21uvC/gc1rrUets3Dqzf4ePWr/qh/bQ+5b7pLFz6VXd" + "Yyze4JPwxEV/R2EKiF7JOXTH5+hxyfnkPRDQGUdOHlTd/538rpsXdQvgvt5nu9ts993vPrTfgubYOym5facu5/7+B+T3" + "+HeqM5tmhccCUGiO5UH2L/98zchDz5EBlIMHoqdV973Uj6ptTP01Ob7vARLA+O4r2TSJxquXye+eaR13bXHJK/uo942+" + "zzAG8Sx23irn/r/CVN9vlckYVBvAYxljjzXvxPVPVbmnnUmK4dePQ4XqhZU3oscCNk8ZdEzZcwf+/BMcf+4vmB6+iASg" + "UM9hMStbXol9jsVk+RTX8LXjWJzjfQ7AcA9Kv2P9TIPUkf0bCl0zfI8dYa+NM2o8d9tqfPnhHSzMafEoCSCAglfGNL+M" + "/XwRTa0aQGHofLFXQQqPdVzhWnLZcsDyslbfdBxTH33NsC7ZuqWWXk/vU6+6e40a9ajvdr7/cu853nd/H37nA7K38Mo5" + "KOA20rZryHH7vgfRNBzF39wCC4+dzwrVAjpwr/x2C9sYhjFwv+jpdTS8W/6mED60Wv6mnpa/LQziN+JeOQO3XsmduYOg" + "SFcsJ4Fy9pXFV480Nq/k6Jh4GPJ9jNp59g2y2D1Kv/inApRnjyO0uOErD2z4A3ssGjU2qMaY22Muytpsa1BcQ2ur9dY6" + "t2p7VK2tTyj5bGNkNeB2Pdp1aOXVNa736XtunNF1GINwJky/+Q8ZYO7aUVSofvnSPWweuxfgc+jjrQ9yeohIAMYZ6m9Z" + "chjNhzAO4nPJvX+VAqbmlrCi7+sNrLH7yuKjaffJfE4rfkgCmP/+XWUanl+JAu4Raw64ZzGlTc86IDCEOOORt79iAR6V" + "L2/C3oW/Idb3FN4Oxbdkb8396gz1WSGcWWuBdvZSI+Pav+DijMmpPlSonv9U/ge74GIIup3Hzmn2nTWG6ZFLFqOyxkOv" + "T4KAlNXzWHLYcrKn68g3qrXkm8bkHPfhOb4P5ZlTPqN0ch90xn/NpnOsLB6q9pwwXvI0R3ywOXme5qkX5/lcyyWbV9bw" + "ea5tHoR5cW/12Gv/oM6RYRQG0+za4Bqm+HkOfoH4aGs05thaHkNcS8fynHt/puCA6OExvKzXrls/mCYSDh+cZWNUFvCh" + "98/tujZgC5V33c3UC/P7BlBhmk2D6rllj2jMOfd/z4CADNAC3bX4Kw7j7RxK6x9VQOCM8T++i5E9Z+TMu1lR5uz6GgTO" + "UVMAbDmyotvKqXDeMlR8vGltNzAbsKWx5XG8536v4n0Md63XHafU2RoB6f2r0Jgzb9+NRXuOcpVH6c0fYzp+F6SxD6aF" + "q3jsgzU5dqjvqOIqX3NnllLG6IEKM0sh7b+MM1c+mO+6go4rbx7jNPQrEFC6U+pxp9Sniiuop2OOKp9K3uQLNV76FkNM" + "Gcv/jUXXPhR/6gw8BjgAxy6/vhvjwjN3kAHSkglOOy7EElcait5xITfylbfv5ozYw2afuPSq2dPK89pVtcJSs81jWGqr" + "2uvYru/IThLI/8YfYdYnbMIcn7CJy1jFbVR++z7FfdZnqc+me6/OC7Z/KvuPd6E+E9ve6z37HajvAzTqUN8lbtSk+S3M" + "OvbHtfdG+xV7tmQi9pvynOqsVM5bnw16vrt2PfxdVO+lehP7L6vfVdQxx/j/CFEl4w==" + ), + }, } diff --git a/tests/unit/torch/quantization/test_iq_formats.py b/tests/unit/torch/quantization/test_iq_formats.py index 113effb0585..8582b725232 100644 --- a/tests/unit/torch/quantization/test_iq_formats.py +++ b/tests/unit/torch/quantization/test_iq_formats.py @@ -30,6 +30,7 @@ ) import modelopt.torch.quantization.ggml.iq1_s as iq1_s_module +import modelopt.torch.quantization.ggml.iq2_s as iq2_s_module import modelopt.torch.quantization.ggml.iq2_xs as iq2_xs_module import modelopt.torch.quantization.ggml.iq2_xxs as iq2_xxs_module from modelopt.torch.quantization.config import QuantizerAttributeConfig @@ -41,8 +42,12 @@ "iq1_s": (iq1_s_module, 50, 2048, 1.5625), "iq2_xxs": (iq2_xxs_module, 66, 256, 2.0625), "iq2_xs": (iq2_xs_module, 74, 512, 2.3125), + "iq2_s": (iq2_s_module, 82, 1024, 2.5625), } NAMES = sorted(FORMATS) +# The formats backend dispatch can reach. A codec can land before it is registered, so the +# tests that go through TensorQuantizer iterate these rather than every codec above. +DISPATCHED = sorted(IQ_FORMAT_REGISTRY) # IQ1 grids are ternary; IQ2 grids hold the magnitudes 8, 25 and 43. TERNARY = {"iq1_s"} @@ -211,7 +216,7 @@ def test_rejects_scalar_packed_payload(name): dequantize(torch.tensor(0, dtype=torch.uint8), torch.tensor([1, 256])) -@pytest.mark.parametrize("name", NAMES) +@pytest.mark.parametrize("name", DISPATCHED) def test_fake_quant_has_pass_through_gradient(name): quantizer = TensorQuantizer( QuantizerAttributeConfig(num_bits=name, block_sizes={-1: 256}, backend="ggml") @@ -265,16 +270,17 @@ def test_error_decreases_with_bit_width(): """More bits must buy less error, or a format's scale handling is wrong.""" generator = torch.Generator().manual_seed(7) weight = torch.randn((4, 1024), generator=generator) + ordered = sorted(DISPATCHED, key=lambda n: FORMATS[n][3]) errors = [] - for name in sorted(NAMES, key=lambda n: FORMATS[n][3]): + for name in ordered: quantizer = TensorQuantizer( QuantizerAttributeConfig(num_bits=name, block_sizes={-1: 256}, backend="ggml") ) errors.append(float((quantizer(weight) - weight).square().mean())) - assert errors == sorted(errors, reverse=True), dict(zip(sorted(NAMES), errors)) + assert errors == sorted(errors, reverse=True), dict(zip(ordered, errors)) def test_every_registered_format_is_covered(): """A format registered for dispatch must also be listed here, or it escapes this contract.""" - assert sorted(IQ_FORMAT_REGISTRY) == sorted(FORMATS) + assert set(IQ_FORMAT_REGISTRY) <= set(FORMATS)