mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[3/5] Add the IQ2_S codec (#2512)
### What does this PR do? Type of change: new feature (not yet user-reachable) **First of two PRs adding IQ2_S**, the widest of the GGML IQ formats at one and two bits (2.5625 bits per weight). This one lands the **PyTorch codec**: the encoder, the decoder and the 1024-entry codebook. It is deliberately **not registered**, so no quantizer dispatches to it and the `ggml` package does not export it. #2565 adds the CUDA encoder, registers the format and adds its recipe. ### What's distinctive about it **IQ2_S is the one format llama.cpp's own tooling gives no head start on**, so the search is written against the GGML layout directly. The interesting difference from IQ2_XS and IQ2_XXS is sign handling. IQ2_S stores a **full 8-bit sign mask** per group rather than a 7-bit parity-coded index. The encoder therefore takes the input signs as they are instead of flipping the weakest element to fix parity, and the search compares magnitudes directly, which is simpler than its siblings. ### Why the codec lands before the kernel The CUDA encoder's tests use this codec as their reference. They compare against the PyTorch encoder byte for byte and draw the grid and scale predictor from it. So the kernel cannot be tested before the codec exists, and it follows in #2565 together with the registration. Every registered format therefore keeps a CUDA encoder. ### Test changes that make the split possible A codec can now land before it is registered, so two test contracts in `test_iq_formats.py` are stated precisely: - The two tests that go through `TensorQuantizer` (pass-through gradient, error falls with bit width) iterate `IQ_FORMAT_REGISTRY`. Every other battery test calls the codec directly and covers IQ2_S here. - The coverage check now asserts `set(IQ_FORMAT_REGISTRY) <= set(FORMATS)` instead of equality. That is what its docstring already said: a registered format must be listed, or it escapes the contract. - `test_registry_lists_every_exported_encoder` is unchanged, and it is why this PR leaves the package exports alone: an exported encoder must be registered. The error-by-bit-width failure message also labels errors by the order they were measured in; it previously zipped them with alphabetical names. ### Testing **The decoder is validated against llama.cpp's own output, not just round-tripped:** ``` IQ2_S: 9 tensors, 2,355,200 blocks → 0 mismatched, max|diff| 0.0 ``` The new codebook matches the `ggml-common.h` table entry for entry. Blocks from `unsloth/Qwen3.8-27B-GGUF` ship as conformance vectors, so CI keeps checking bytes we did not produce. - `tests/unit/torch/quantization/test_ggml_backend.py`, `test_iq_formats.py`, `tests/unit/torch/export/test_convert_hf_config.py`, `tests/unit/recipe/test_presets.py`: **121 passed**, 14 of them IQ2_S codec cases, including the llama.cpp conformance check - `tests/gpu/torch/quantization/test_iq_formats_cuda.py`, `test_iq1_s_cuda.py`, `test_iq2_xs_cuda.py`: **35 passed**, unchanged by this PR ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ The new codebook is a GGML table, carried in `codebooks.py` with the source revision recorded. No new dependencies. - Did you write any new necessary tests?: ✅ - Did you update Changelog?: N/A. Nothing is user-reachable yet; #2565 carries the entry. - Did you get Claude approval on this PR?: ❌ Not yet run. ### Additional Information Merge order: #2511 (IQ2_XXS, merged) → #2525 (format registry, merged) → **this** → #2565 (IQ2_S CUDA encoder and registration) → #2513 (IQ1_M). 🤖 Generated with [Claude Code](https://claude.com/claude-code) <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added GGML-compatible IQ2_S quantization and dequantization support, including access to its magnitude grid. <!-- end of auto-generated comment: release notes by coderabbit.ai --> Signed-off-by: Chenjie Luo <chenjiel@nvidia.com> Co-authored-by: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5.5
parent
57f929e358
commit
767ef5533e
@@ -178,6 +178,33 @@ _IQ2_XXS_GRID_ZLIB_B64 = (
|
||||
"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))
|
||||
|
||||
@@ -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)
|
||||
@@ -141,6 +141,49 @@ _VECTORS = {
|
||||
"/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=="
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -30,6 +30,7 @@ from _test_utils.torch.quantization.iq_llama_cpp_vectors import (
|
||||
)
|
||||
|
||||
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 @@ FORMATS = {
|
||||
"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)
|
||||
|
||||
Reference in New Issue
Block a user