Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions modelopt/torch/quantization/ggml/codebooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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))
217 changes: 217 additions & 0 deletions modelopt/torch/quantization/ggml/iq2_s.py
Original file line number Diff line number Diff line change
@@ -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)
43 changes: 43 additions & 0 deletions tests/_test_utils/torch/quantization/iq_llama_cpp_vectors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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=="
),
},
}


Expand Down
14 changes: 10 additions & 4 deletions tests/unit/torch/quantization/test_iq_formats.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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"}

Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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)
Loading