diff --git a/modelopt/torch/export/convert_hf_config.py b/modelopt/torch/export/convert_hf_config.py index 8bf8cb8ef5d..c50e25cdf5a 100644 --- a/modelopt/torch/export/convert_hf_config.py +++ b/modelopt/torch/export/convert_hf_config.py @@ -19,7 +19,9 @@ from collections import defaultdict from typing import Any -from .quant_format import IQ_BLOCK_METADATA, IQ_FORMATS +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY + +from .quant_format import IQ_FORMATS def _quant_algo_to_group_config(quant_algo: str, group_size: int | None = None) -> dict[str, Any]: @@ -121,7 +123,9 @@ def _quant_algo_to_group_config(quant_algo: str, group_size: int | None = None) "weights": {"dynamic": False, "num_bits": 8, "type": "float", "group_size": gs}, } elif quant_algo.lower() in IQ_FORMATS: - block_size, payload_bytes, effective_bits = IQ_BLOCK_METADATA[quant_algo.lower()] + iq_format = IQ_FORMAT_REGISTRY[quant_algo.lower()] + block_size, payload_bytes = iq_format.block_size, iq_format.block_bytes + effective_bits = iq_format.effective_bits if group_size not in (None, block_size): raise ValueError(f"{quant_algo} requires group size {block_size}, got {group_size}") # IQ payloads are self-contained blocks, not compressed-tensors integer groups. diff --git a/modelopt/torch/export/quant_format.py b/modelopt/torch/export/quant_format.py index 44fb4cb76eb..ffe655c2d71 100644 --- a/modelopt/torch/export/quant_format.py +++ b/modelopt/torch/export/quant_format.py @@ -19,20 +19,7 @@ constants, for example, are in :mod:`modelopt.torch.export.trtllm.model_config`. """ -from modelopt.torch.quantization.ggml import ( - IQ1_S_BLOCK_BYTES, - IQ1_S_BLOCK_SIZE, - IQ1_S_EFFECTIVE_BITS, - IQ2_XS_BLOCK_BYTES, - IQ2_XS_BLOCK_SIZE, - IQ2_XS_EFFECTIVE_BITS, - IQ2_XXS_BLOCK_BYTES, - IQ2_XXS_BLOCK_SIZE, - IQ2_XXS_EFFECTIVE_BITS, - quantize_iq1_s, - quantize_iq2_xs, - quantize_iq2_xxs, -) +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY QUANTIZATION_NONE = None QUANTIZATION_FP8 = "fp8" @@ -55,34 +42,15 @@ QUANTIZATION_IQ2_XXS = "iq2_xxs" QUANTIZATION_IQ2_XS = "iq2_xs" -# Every GGML IQ format. They share the weight-only, 256-value-block, per-module-scale -# shape, so export treats them as one family; adding a format means adding it here -# rather than extending a tuple at each use site. -IQ_FORMATS = frozenset( - { - QUANTIZATION_IQ1_S, - QUANTIZATION_IQ2_XXS, - QUANTIZATION_IQ2_XS, - } -) - -# Block geometry per IQ format: (block size, packed bytes per block, bits per weight). Checkpoint -# metadata spells the algorithm in upper case, so consumers look up -# ``IQ_BLOCK_METADATA[algo.lower()]`` rather than carrying a second spelling of the family. -IQ_BLOCK_METADATA = { - QUANTIZATION_IQ1_S: (IQ1_S_BLOCK_SIZE, IQ1_S_BLOCK_BYTES, IQ1_S_EFFECTIVE_BITS), - QUANTIZATION_IQ2_XXS: (IQ2_XXS_BLOCK_SIZE, IQ2_XXS_BLOCK_BYTES, IQ2_XXS_EFFECTIVE_BITS), - QUANTIZATION_IQ2_XS: (IQ2_XS_BLOCK_SIZE, IQ2_XS_BLOCK_BYTES, IQ2_XS_EFFECTIVE_BITS), -} - - -# The packer each format's checkpoint weights are written with. Both exporters resolve through -# this one mapping so they cannot drift apart. -IQ_PACKERS = { - QUANTIZATION_IQ1_S: quantize_iq1_s, - QUANTIZATION_IQ2_XXS: quantize_iq2_xxs, - QUANTIZATION_IQ2_XS: quantize_iq2_xs, -} +# Every GGML IQ format, derived from the registry the quantization backend dispatches through, so +# export and dispatch cannot disagree about which formats exist. They share the weight-only, +# 256-value-block, per-module-scale shape, so export treats them as one family. A format's block +# geometry and packer are read from IQ_FORMAT_REGISTRY directly. +# +# Registering a format therefore declares it exportable, and that is intended rather than a side +# effect: fake quant is dequantize(quantize(w)), so a format cannot be dispatched without the +# packer and block geometry that are all export reads. +IQ_FORMATS = frozenset(IQ_FORMAT_REGISTRY) # Formats whose scales are purely per-module, so export never merges them across the q/k/v diff --git a/modelopt/torch/export/quant_utils.py b/modelopt/torch/export/quant_utils.py index fb25d23c5e4..5226ec94c3e 100755 --- a/modelopt/torch/export/quant_utils.py +++ b/modelopt/torch/export/quant_utils.py @@ -27,6 +27,7 @@ from modelopt import __version__ from modelopt.torch.models import get_spec, list_all_possible +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY from modelopt.torch.quantization.model_calib import ( enable_stats_collection, finish_stats_collection, @@ -51,7 +52,6 @@ from ..quantization.nn import NVFP4StaticQuantizer, SequentialQuantizer, TensorQuantizer from .model_utils import TiedWeightMap, get_language_model_from_vl from .quant_format import ( - IQ_BLOCK_METADATA, IQ_FORMATS, KV_CACHE_FP8, KV_CACHE_FP8_K_NVFP4_V, @@ -773,7 +773,9 @@ def process_layer_quant_config(layer_config_dict): "group_size": block_size_value, } elif v in IQ_FORMATS: - block_size, payload_bytes, effective_bits = IQ_BLOCK_METADATA[v] + iq_format = IQ_FORMAT_REGISTRY[v] + block_size, payload_bytes = iq_format.block_size, iq_format.block_bytes + effective_bits = iq_format.effective_bits if block_size_value != block_size: raise ValueError( f"{v.upper()} requires block size {block_size}, got {block_size_value}" diff --git a/modelopt/torch/export/unified_export_hf.py b/modelopt/torch/export/unified_export_hf.py index 1ec13cccb97..64e158c05bb 100644 --- a/modelopt/torch/export/unified_export_hf.py +++ b/modelopt/torch/export/unified_export_hf.py @@ -67,6 +67,7 @@ from modelopt.torch.opt.conversion import ModeloptStateManager, modelopt_state from modelopt.torch.opt.plugins.huggingface import _MODELOPT_STATE_SAVE_NAME from modelopt.torch.quantization import set_quantizer_by_cfg_context +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY from modelopt.torch.quantization.nn import SequentialQuantizer, TensorQuantizer from modelopt.torch.quantization.qtensor import MXFP8QTensor, NVFP4QTensor from modelopt.torch.quantization.qtensor.base_qtensor import QTensorWrapper @@ -101,7 +102,6 @@ from .quant_format import ( FUSION_FREE_FORMATS, IQ_FORMATS, - IQ_PACKERS, QUANTIZATION_FP8, QUANTIZATION_FP8_PB_REAL, QUANTIZATION_FP8_PC_PT, @@ -635,7 +635,7 @@ def _export_quantized_weight( "IQ unified export currently supports modules with a standard 'weight' " f"attribute, got {weight_name!r} on {type(sub_module).__name__}" ) - quantize_iq = IQ_PACKERS[quantization_format] + quantize_iq = IQ_FORMAT_REGISTRY[quantization_format].quantize packed_weight, _ = quantize_iq(weight.to(dtype)) setattr(sub_module, weight_name, nn.Parameter(packed_weight, requires_grad=False)) maybe_clear_cuda_cache() diff --git a/modelopt/torch/export/unified_export_megatron.py b/modelopt/torch/export/unified_export_megatron.py index 67564853ce8..9a87632d417 100644 --- a/modelopt/torch/export/unified_export_megatron.py +++ b/modelopt/torch/export/unified_export_megatron.py @@ -35,6 +35,7 @@ from safetensors.torch import save_file from modelopt import __version__ +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY from modelopt.torch.quantization.nn.modules.tensor_quantizer import GroupedQuantizer from modelopt.torch.utils import import_plugin, warn_rank_0 from modelopt.torch.utils.plugins.hf_checkpoint_utils import ( @@ -57,7 +58,6 @@ from .plugins.megatron_importer import GPTModelImporter, _get_mamba_conv1d from .quant_format import ( IQ_FORMATS, - IQ_PACKERS, KV_CACHE_FP8, KV_CACHE_NVFP4, QUANTIZATION_FP8, @@ -1185,7 +1185,7 @@ def _get_weight_scales(self, quantized_state: dict[str, Any], qformat: str): @staticmethod def _pack_iq_weight(weight: torch.Tensor, qformat: str) -> torch.Tensor: """Pack one ``[out, in]`` weight and return its CPU payload.""" - quantize_iq = IQ_PACKERS[qformat] + quantize_iq = IQ_FORMAT_REGISTRY[qformat].quantize packed_weight, _ = quantize_iq(weight) return packed_weight.detach().cpu() diff --git a/modelopt/torch/quantization/ggml/__init__.py b/modelopt/torch/quantization/ggml/__init__.py index 66065b19017..db8a3473067 100644 --- a/modelopt/torch/quantization/ggml/__init__.py +++ b/modelopt/torch/quantization/ggml/__init__.py @@ -23,9 +23,12 @@ from .iq2_xs import __all__ as _iq2_xs_all from .iq2_xxs import * from .iq2_xxs import __all__ as _iq2_xxs_all +from .registry import IQ_FORMAT_REGISTRY, IQFormat __all__ = [ # noqa: PLE0604 *_iq1_s_all, *_iq2_xs_all, *_iq2_xxs_all, + "IQ_FORMAT_REGISTRY", + "IQFormat", ] diff --git a/modelopt/torch/quantization/ggml/backend.py b/modelopt/torch/quantization/ggml/backend.py index d1d51e85607..1f4a83eeff6 100644 --- a/modelopt/torch/quantization/ggml/backend.py +++ b/modelopt/torch/quantization/ggml/backend.py @@ -18,16 +18,7 @@ import torch from ..nn.modules.tensor_quantizer import register_quant_backend -from .iq1_s import iq1_s_fake_quant -from .iq2_xs import iq2_xs_fake_quant -from .iq2_xxs import iq2_xxs_fake_quant - -# One entry per GGML IQ format; adding a format is adding a row here. -_FAKE_QUANTS = { - "iq1_s": iq1_s_fake_quant, - "iq2_xs": iq2_xs_fake_quant, - "iq2_xxs": iq2_xxs_fake_quant, -} +from .registry import IQ_FORMAT_REGISTRY def ggml_fake_quant(inputs: torch.Tensor, quantizer) -> torch.Tensor: @@ -39,11 +30,11 @@ def ggml_fake_quant(inputs: torch.Tensor, quantizer) -> torch.Tensor: raise ValueError(f"Unsupported ggml backend_extra_args: {sorted(unknown_args)}") # num_bits arrives untyped from the quantizer and is a tuple for scalar formats, # so narrow before the lookup rather than relying on the dict to reject it. - fake_quant = _FAKE_QUANTS.get(num_bits) if isinstance(num_bits, str) else None - if fake_quant is None: - supported = ", ".join(repr(name) for name in sorted(_FAKE_QUANTS)) + iq_format = IQ_FORMAT_REGISTRY.get(num_bits) if isinstance(num_bits, str) else None + if iq_format is None: + supported = ", ".join(repr(name) for name in sorted(IQ_FORMAT_REGISTRY)) raise ValueError(f"The ggml backend requires num_bits in ({supported})") - return fake_quant(inputs, quantizer, **extra_args) + return iq_format.fake_quant(inputs, quantizer, **extra_args) register_quant_backend("ggml", ggml_fake_quant) diff --git a/modelopt/torch/quantization/ggml/common.py b/modelopt/torch/quantization/ggml/common.py index 574db9bcac8..2b39d3e73b7 100644 --- a/modelopt/torch/quantization/ggml/common.py +++ b/modelopt/torch/quantization/ggml/common.py @@ -121,6 +121,59 @@ def fake_quantize_with_cache( return inputs + (reconstructed - inputs).detach() +@dataclass(frozen=True) +class IQFormat: + """Everything backend dispatch and export need to know about one IQ format. + + Each format module declares one of these beside its encoder and decoder, and + :data:`~modelopt.torch.quantization.ggml.registry.IQ_FORMAT_REGISTRY` lists them. The + per-format pieces -- codebook, search, payload layout -- stay in the format's module; what + lives here is the part every format does the same way. + """ + + name: str + block_size: int + block_bytes: int + quantize: Callable[..., tuple[torch.Tensor, torch.Tensor]] + dequantize: Callable[..., torch.Tensor] + # Encode and decode are chunked separately: packing runs once per weight and is bounded by + # its search temporaries, decoding runs every forward and is bounded by kernel launches. + block_chunk_size: int + decode_chunk_size: int + + @property + def effective_bits(self) -> float: + """Packed storage cost per weight.""" + return self.block_bytes * 8 / self.block_size + + def fake_quant( + self, + inputs: torch.Tensor, + quantizer, + *, + block_chunk_size: int | None = None, + decode_chunk_size: int | None = None, + ) -> torch.Tensor: + """TensorQuantizer backend for this format, with pass-through backward.""" + if getattr(quantizer, "num_bits", None) != self.name: + raise ValueError( + f"The ggml {self.name.upper()} backend requires num_bits={self.name!r}" + ) + return fake_quantize_with_cache( + inputs, + quantizer, + format_name=self.name, + block_chunk_size=( + self.block_chunk_size if block_chunk_size is None else block_chunk_size + ), + decode_chunk_size=( + self.decode_chunk_size if decode_chunk_size is None else decode_chunk_size + ), + quantize=self.quantize, + dequantize=self.dequantize, + ) + + def narrow_to_float32(blocks: torch.Tensor) -> torch.Tensor: """Narrow ``blocks`` to float32 the way the CUDA ``load_float`` helper does. diff --git a/modelopt/torch/quantization/ggml/iq1_s.py b/modelopt/torch/quantization/ggml/iq1_s.py index a43e1699db1..956c52458ee 100644 --- a/modelopt/torch/quantization/ggml/iq1_s.py +++ b/modelopt/torch/quantization/ggml/iq1_s.py @@ -38,7 +38,7 @@ from .codebooks import iq1_s_grid_bytes from .common import ( GGML_BLOCK_SIZE, - fake_quantize_with_cache, + IQFormat, narrow_to_float32, validate_block_chunk_size, validate_packed_weights, @@ -235,22 +235,17 @@ def dequantize_iq1_s( return decoded.reshape(shape) -def iq1_s_fake_quant( - inputs: torch.Tensor, - quantizer, - *, - block_chunk_size: int = _DEFAULT_BLOCK_CHUNK_SIZE, - decode_chunk_size: int = _DEFAULT_DECODE_CHUNK_SIZE, -) -> torch.Tensor: - """IQ1_S weight backend for TensorQuantizer, with pass-through backward.""" - if getattr(quantizer, "num_bits", None) != "iq1_s": - raise ValueError("The ggml IQ1_S backend requires num_bits='iq1_s'") - return fake_quantize_with_cache( - inputs, - quantizer, - format_name="iq1_s", - block_chunk_size=block_chunk_size, - decode_chunk_size=decode_chunk_size, - quantize=quantize_iq1_s, - dequantize=dequantize_iq1_s, - ) +IQ1_S_FORMAT = IQFormat( + name="iq1_s", + block_size=IQ1_S_BLOCK_SIZE, + block_bytes=IQ1_S_BLOCK_BYTES, + quantize=quantize_iq1_s, + dequantize=dequantize_iq1_s, + block_chunk_size=_DEFAULT_BLOCK_CHUNK_SIZE, + decode_chunk_size=_DEFAULT_DECODE_CHUNK_SIZE, +) + +# Kept for callers of the per-format entry point. The record captured quantize_iq1_s and +# dequantize_iq1_s when it was built, so patching those module functions changes neither backend +# dispatch nor this alias; substitute a format's encoder or decoder in IQ_FORMAT_REGISTRY. +iq1_s_fake_quant = IQ1_S_FORMAT.fake_quant diff --git a/modelopt/torch/quantization/ggml/iq2_xs.py b/modelopt/torch/quantization/ggml/iq2_xs.py index e2adc2fa44f..a4b7d505bc0 100644 --- a/modelopt/torch/quantization/ggml/iq2_xs.py +++ b/modelopt/torch/quantization/ggml/iq2_xs.py @@ -36,7 +36,7 @@ from .codebooks import iq2_xs_grid_bytes from .common import ( GGML_BLOCK_SIZE, - fake_quantize_with_cache, + IQFormat, narrow_to_float32, validate_block_chunk_size, validate_packed_weights, @@ -255,22 +255,17 @@ def dequantize_iq2_xs( return decoded.reshape(shape) -def iq2_xs_fake_quant( - inputs: torch.Tensor, - quantizer, - *, - block_chunk_size: int = _DEFAULT_BLOCK_CHUNK_SIZE, - decode_chunk_size: int = _DEFAULT_DECODE_CHUNK_SIZE, -) -> torch.Tensor: - """IQ2_XS weight backend for TensorQuantizer, with pass-through backward.""" - if getattr(quantizer, "num_bits", None) != "iq2_xs": - raise ValueError("The ggml IQ2_XS backend requires num_bits='iq2_xs'") - return fake_quantize_with_cache( - inputs, - quantizer, - format_name="iq2_xs", - block_chunk_size=block_chunk_size, - decode_chunk_size=decode_chunk_size, - quantize=quantize_iq2_xs, - dequantize=dequantize_iq2_xs, - ) +IQ2_XS_FORMAT = IQFormat( + name="iq2_xs", + block_size=IQ2_XS_BLOCK_SIZE, + block_bytes=IQ2_XS_BLOCK_BYTES, + quantize=quantize_iq2_xs, + dequantize=dequantize_iq2_xs, + block_chunk_size=_DEFAULT_BLOCK_CHUNK_SIZE, + decode_chunk_size=_DEFAULT_DECODE_CHUNK_SIZE, +) + +# Kept for callers of the per-format entry point. The record captured quantize_iq2_xs and +# dequantize_iq2_xs when it was built, so patching those module functions changes neither backend +# dispatch nor this alias; substitute a format's encoder or decoder in IQ_FORMAT_REGISTRY. +iq2_xs_fake_quant = IQ2_XS_FORMAT.fake_quant diff --git a/modelopt/torch/quantization/ggml/iq2_xxs.py b/modelopt/torch/quantization/ggml/iq2_xxs.py index 01d8bce0a76..a27c2679632 100644 --- a/modelopt/torch/quantization/ggml/iq2_xxs.py +++ b/modelopt/torch/quantization/ggml/iq2_xxs.py @@ -41,7 +41,7 @@ from .codebooks import iq2_xxs_grid_bytes from .common import ( GGML_BLOCK_SIZE, - fake_quantize_with_cache, + IQFormat, narrow_to_float32, validate_block_chunk_size, validate_packed_weights, @@ -265,22 +265,17 @@ def dequantize_iq2_xxs( return decoded.reshape(shape) -def iq2_xxs_fake_quant( - inputs: torch.Tensor, - quantizer, - *, - block_chunk_size: int = _DEFAULT_BLOCK_CHUNK_SIZE, - decode_chunk_size: int = _DEFAULT_DECODE_CHUNK_SIZE, -) -> torch.Tensor: - """IQ2_XXS weight backend for TensorQuantizer, with pass-through backward.""" - if getattr(quantizer, "num_bits", None) != "iq2_xxs": - raise ValueError("The ggml IQ2_XXS backend requires num_bits='iq2_xxs'") - return fake_quantize_with_cache( - inputs, - quantizer, - format_name="iq2_xxs", - block_chunk_size=block_chunk_size, - decode_chunk_size=decode_chunk_size, - quantize=quantize_iq2_xxs, - dequantize=dequantize_iq2_xxs, - ) +IQ2_XXS_FORMAT = IQFormat( + name="iq2_xxs", + block_size=IQ2_XXS_BLOCK_SIZE, + block_bytes=IQ2_XXS_BLOCK_BYTES, + quantize=quantize_iq2_xxs, + dequantize=dequantize_iq2_xxs, + block_chunk_size=_DEFAULT_BLOCK_CHUNK_SIZE, + decode_chunk_size=_DEFAULT_DECODE_CHUNK_SIZE, +) + +# Kept for callers of the per-format entry point. The record captured quantize_iq2_xxs and +# dequantize_iq2_xxs when it was built, so patching those module functions changes neither backend +# dispatch nor this alias; substitute a format's encoder or decoder in IQ_FORMAT_REGISTRY. +iq2_xxs_fake_quant = IQ2_XXS_FORMAT.fake_quant diff --git a/modelopt/torch/quantization/ggml/registry.py b/modelopt/torch/quantization/ggml/registry.py new file mode 100644 index 00000000000..5faca3c0377 --- /dev/null +++ b/modelopt/torch/quantization/ggml/registry.py @@ -0,0 +1,33 @@ +# 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. + +"""The GGML IQ formats, listed once for backend dispatch and export.""" + +from .common import IQFormat +from .iq1_s import IQ1_S_FORMAT +from .iq2_xs import IQ2_XS_FORMAT +from .iq2_xxs import IQ2_XXS_FORMAT + +__all__ = ["IQ_FORMAT_REGISTRY", "IQFormat"] + +# Every IQ format, keyed by the name a quantizer's num_bits carries, in increasing bits per +# weight. Backend dispatch and both exporters read this mapping and export's IQ_FORMATS is derived +# from it, so adding a format is one entry here rather than a row in several parallel tables. +# +# It is an explicit list rather than formats registering themselves on import, so its contents +# never depend on which modules happen to have been imported first. +IQ_FORMAT_REGISTRY: dict[str, IQFormat] = { + fmt.name: fmt for fmt in (IQ1_S_FORMAT, IQ2_XXS_FORMAT, IQ2_XS_FORMAT) +} diff --git a/tests/gpu/torch/quantization/test_iq_formats_cuda.py b/tests/gpu/torch/quantization/test_iq_formats_cuda.py index 70d85c1d0bd..50546b370e9 100644 --- a/tests/gpu/torch/quantization/test_iq_formats_cuda.py +++ b/tests/gpu/torch/quantization/test_iq_formats_cuda.py @@ -31,6 +31,7 @@ IQ1_S_BLOCK_BYTES, IQ2_XS_BLOCK_BYTES, IQ2_XXS_BLOCK_BYTES, + IQ_FORMAT_REGISTRY, ) # module, packer name, per-block payload size, whether the packer takes precomputed scales @@ -177,3 +178,8 @@ def test_cuda_float64_matches_pytorch_encoder(monkeypatch, name): monkeypatch.setattr(module, "get_cuda_ext_ggml", lambda: None) reference, _ = getattr(module, f"quantize_{name}")(weight) assert torch.equal(reference, packed) + + +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) diff --git a/tests/gpu_megatron/torch/export/test_unified_export_megatron.py b/tests/gpu_megatron/torch/export/test_unified_export_megatron.py index d3f23c7f4f8..561cd6ae050 100644 --- a/tests/gpu_megatron/torch/export/test_unified_export_megatron.py +++ b/tests/gpu_megatron/torch/export/test_unified_export_megatron.py @@ -92,7 +92,7 @@ def _verify_model_quant_config( # Every IQ format the exporter accepts. Only the list of formats comes from the export # tables; each test resolves what it expects from the codec module itself, so a wrong entry -# in IQ_PACKERS or IQ_BLOCK_METADATA cannot make both sides of an assertion agree. +# in IQ_FORMAT_REGISTRY cannot make both sides of an assertion agree. IQ_FORMAT_NAMES = sorted(IQ_FORMATS) diff --git a/tests/unit/torch/export/test_convert_hf_config.py b/tests/unit/torch/export/test_convert_hf_config.py index bc446c9e042..1d2a9e09e0c 100644 --- a/tests/unit/torch/export/test_convert_hf_config.py +++ b/tests/unit/torch/export/test_convert_hf_config.py @@ -15,10 +15,22 @@ import pytest +import modelopt.torch.export.quant_format as quant_format import modelopt.torch.quantization.ggml as ggml from modelopt.torch.export.convert_hf_config import convert_hf_quant_config_format -from modelopt.torch.export.quant_format import IQ_BLOCK_METADATA, IQ_FORMATS +from modelopt.torch.export.quant_format import IQ_FORMATS from modelopt.torch.export.unified_export_hf import _revert_hf_quant_config_names +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY + + +def _geometry(fmt): + """Block geometry from the codec's own constants, independent of the registry under test.""" + upper = fmt.upper() + return ( + getattr(ggml, f"{upper}_BLOCK_SIZE"), + getattr(ggml, f"{upper}_BLOCK_BYTES"), + getattr(ggml, f"{upper}_EFFECTIVE_BITS"), + ) def test_convert_mixed_kv_cache_config_preserves_layer_map(): @@ -121,7 +133,7 @@ def test_iq_config_carries_block_metadata(fmt): A consumer reads group_size and block_payload_bytes to walk the payload, so a format that falls through to the generic branch produces a checkpoint that cannot be decoded. """ - block_size, payload_bytes, effective_bits = IQ_BLOCK_METADATA[fmt] + block_size, payload_bytes, effective_bits = _geometry(fmt) converted = convert_hf_quant_config_format( { "producer": {"name": "modelopt", "version": "test"}, @@ -141,7 +153,7 @@ def test_iq_config_carries_block_metadata(fmt): @pytest.mark.parametrize("fmt", sorted(IQ_FORMATS)) def test_iq_config_rejects_mismatched_group_size(fmt): """A caller's group size is rejected rather than silently rewritten to the block size.""" - block_size, _, _ = IQ_BLOCK_METADATA[fmt] + block_size, _, _ = _geometry(fmt) with pytest.raises(ValueError, match=f"requires group size {block_size}"): convert_hf_quant_config_format( { @@ -151,15 +163,11 @@ def test_iq_config_rejects_mismatched_group_size(fmt): ) -@pytest.mark.parametrize("fmt", sorted(IQ_FORMATS)) -def test_iq_block_metadata_matches_the_codec(fmt): - """The exported geometry is the codec's own, so a checkpoint cannot claim a wrong layout.""" - block_size, payload_bytes, effective_bits = IQ_BLOCK_METADATA[fmt] - upper = fmt.upper() - assert block_size == getattr(ggml, f"{upper}_BLOCK_SIZE") - assert payload_bytes == getattr(ggml, f"{upper}_BLOCK_BYTES") - assert effective_bits == pytest.approx(getattr(ggml, f"{upper}_EFFECTIVE_BITS")) - assert effective_bits == pytest.approx(payload_bytes * 8 / block_size) +def test_export_formats_are_the_registered_formats(): + """Export and backend dispatch agree on which IQ formats exist, name constants included.""" + assert frozenset(IQ_FORMAT_REGISTRY) == IQ_FORMATS + constants = {v for k, v in vars(quant_format).items() if k.startswith("QUANTIZATION_IQ")} + assert constants == IQ_FORMATS @pytest.mark.parametrize("fmt", sorted(IQ_FORMATS)) @@ -169,7 +177,7 @@ def test_iq_mixed_precision_config_group_carries_block_metadata(fmt): Mixed exports route each distinct layer config through the same helper, so a format missing there loses its geometry for exactly the layers that use it. """ - block_size, payload_bytes, effective_bits = IQ_BLOCK_METADATA[fmt] + block_size, payload_bytes, effective_bits = _geometry(fmt) converted = convert_hf_quant_config_format( { "producer": {"name": "modelopt", "version": "test"}, @@ -197,7 +205,7 @@ def test_iq_mixed_precision_config_group_carries_block_metadata(fmt): @pytest.mark.parametrize("fmt", sorted(IQ_FORMATS)) def test_iq_mixed_precision_rejects_bad_per_layer_group_size(fmt): """A per-layer group size is validated, not silently rewritten to the block size.""" - block_size, _, _ = IQ_BLOCK_METADATA[fmt] + block_size, _, _ = _geometry(fmt) with pytest.raises(ValueError, match=f"requires group size {block_size}"): convert_hf_quant_config_format( { diff --git a/tests/unit/torch/quantization/test_ggml_backend.py b/tests/unit/torch/quantization/test_ggml_backend.py index a9f93a7aee5..cffea63cb74 100644 --- a/tests/unit/torch/quantization/test_ggml_backend.py +++ b/tests/unit/torch/quantization/test_ggml_backend.py @@ -13,22 +13,45 @@ # See the License for the specific language governing permissions and # limitations under the License. +import dataclasses +import importlib +import re from types import SimpleNamespace import pytest import torch import modelopt.torch.quantization as mtq -import modelopt.torch.quantization.ggml.backend as backend_module +import modelopt.torch.quantization.ggml as ggml import modelopt.torch.quantization.ggml.iq1_s as iq1_s_module import modelopt.torch.quantization.ggml.iq2_xs as iq2_xs_module from modelopt.torch.quantization.config import QuantizerAttributeConfig +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY from modelopt.torch.quantization.ggml.backend import ggml_fake_quant from modelopt.torch.quantization.ggml.common import narrow_to_float32 from modelopt.torch.quantization.nn import TensorQuantizer +# Every registered format, paired with its own module. The module -- not the registry -- supplies +# what a test expects, so a mis-wired registry entry cannot make both sides of an assertion agree. +FORMAT_NAMES = sorted(IQ_FORMAT_REGISTRY) +FORMAT_MODULES = [ + (name, importlib.import_module(f"modelopt.torch.quantization.ggml.{name}")) + for name in FORMAT_NAMES +] -@pytest.mark.parametrize("num_bits", ["iq1_s", "iq2_xs"]) + +def _substitute(monkeypatch, name, **changes): + """Replace parts of one registered format for the duration of a test. + + Dispatch reads IQ_FORMAT_REGISTRY, so that is where a substitute has to go: patching the + format module's function would not reach it. + """ + monkeypatch.setitem( + IQ_FORMAT_REGISTRY, name, dataclasses.replace(IQ_FORMAT_REGISTRY[name], **changes) + ) + + +@pytest.mark.parametrize("num_bits", FORMAT_NAMES) def test_ggml_backend_via_quantize(num_bits): torch.manual_seed(1234) model = torch.nn.Linear(256, 2, bias=False) @@ -57,7 +80,9 @@ def test_ggml_backend_via_quantize(num_bits): def test_ggml_backend_rejects_unknown_format(): - with pytest.raises(ValueError, match="requires num_bits"): + # Match the dispatcher's own wording: a bare "requires num_bits" would also accept the + # per-format guard's message and so could pass on the wrong error. + with pytest.raises(ValueError, match="requires num_bits in"): ggml_fake_quant(torch.ones(1, 256), SimpleNamespace(num_bits="unknown")) @@ -82,8 +107,8 @@ def fake_quant(inputs, _quantizer, **kwargs): received.update(kwargs) return inputs - # The dispatcher resolves through its registry, so that is the seam to patch. - monkeypatch.setitem(backend_module._FAKE_QUANTS, "iq1_s", fake_quant) + # The dispatcher resolves through the registry, so that is the seam to patch. + monkeypatch.setitem(IQ_FORMAT_REGISTRY, "iq1_s", SimpleNamespace(fake_quant=fake_quant)) inputs = torch.ones(1, 256) quantizer = SimpleNamespace(num_bits="iq1_s", backend_extra_args=extra_args) @@ -98,19 +123,11 @@ def test_ggml_backend_rejects_unknown_extra_arg(): ggml_fake_quant(torch.ones(1, 256), quantizer) -@pytest.mark.parametrize( - ("num_bits", "module", "fake_quant_name", "quantize_name"), - [ - ("iq1_s", iq1_s_module, "iq1_s_fake_quant", "quantize_iq1_s"), - ("iq2_xs", iq2_xs_module, "iq2_xs_fake_quant", "quantize_iq2_xs"), - ], -) -def test_ggml_backend_caches_packed_weight_and_invalidates_on_change( - monkeypatch, num_bits, module, fake_quant_name, quantize_name -): +@pytest.mark.parametrize(("num_bits", "module"), FORMAT_MODULES) +def test_ggml_backend_caches_packed_weight_and_invalidates_on_change(monkeypatch, num_bits, module): weight = torch.randn(1, 256) quantizer = SimpleNamespace(num_bits=num_bits, _quantizer_cache=None) - original_quantize = getattr(module, quantize_name) + original_quantize = getattr(module, f"quantize_{num_bits}") call_count = 0 def counted_quantize(*args, **kwargs): @@ -118,8 +135,8 @@ def counted_quantize(*args, **kwargs): call_count += 1 return original_quantize(*args, **kwargs) - monkeypatch.setattr(module, quantize_name, counted_quantize) - fake_quant = getattr(module, fake_quant_name) + _substitute(monkeypatch, num_bits, quantize=counted_quantize) + fake_quant = IQ_FORMAT_REGISTRY[num_bits].fake_quant fake_quant(weight, quantizer, block_chunk_size=1) fake_quant(weight, quantizer, block_chunk_size=1) @@ -149,9 +166,7 @@ def test_narrow_to_float32_matches_the_cuda_load_float_policy(): ) -@pytest.mark.parametrize( - ("num_bits", "module"), [("iq1_s", iq1_s_module), ("iq2_xs", iq2_xs_module)] -) +@pytest.mark.parametrize(("num_bits", "module"), FORMAT_MODULES) def test_ggml_weight_is_packed_once_across_forwards(monkeypatch, num_bits, module): """The packed weight is reused across forwards rather than re-encoded each time. @@ -167,7 +182,7 @@ def counting(weight, **kwargs): calls.append(tuple(weight.shape)) return original(weight, **kwargs) - monkeypatch.setattr(module, packer, counting) + _substitute(monkeypatch, num_bits, quantize=counting) quantizer = TensorQuantizer( QuantizerAttributeConfig(num_bits=num_bits, block_sizes={-1: 256}, backend="ggml") ) @@ -180,9 +195,7 @@ def counting(weight, **kwargs): assert calls == [(4, 256)], f"expected one pack, got {len(calls)}" -@pytest.mark.parametrize( - ("num_bits", "module"), [("iq1_s", iq1_s_module), ("iq2_xs", iq2_xs_module)] -) +@pytest.mark.parametrize(("num_bits", "module"), FORMAT_MODULES) def test_ggml_decode_chunk_is_sized_independently_of_the_encode_chunk( monkeypatch, num_bits, module ): @@ -198,7 +211,7 @@ def recording(packed_weights, weight_shape, **kwargs): seen["block_chunk_size"] = kwargs["block_chunk_size"] return original(packed_weights, weight_shape, **kwargs) - monkeypatch.setattr(module, f"dequantize_{num_bits}", recording) + _substitute(monkeypatch, num_bits, dequantize=recording) quantizer = TensorQuantizer( QuantizerAttributeConfig(num_bits=num_bits, block_sizes={-1: 256}, backend="ggml") ) @@ -208,18 +221,45 @@ def recording(packed_weights, weight_shape, **kwargs): assert module._DEFAULT_DECODE_CHUNK_SIZE > module._DEFAULT_BLOCK_CHUNK_SIZE -@pytest.mark.parametrize( - ("num_bits", "module"), [("iq1_s", iq1_s_module), ("iq2_xs", iq2_xs_module)] -) -def test_ggml_decode_is_invariant_to_chunk_size(num_bits, module): - """Chunking the decode is a memory bound, not a numerical choice.""" - torch.manual_seed(0) - weight = torch.randn(3, 1024, dtype=torch.bfloat16) - packed, shape = getattr(module, f"quantize_{num_bits}")(weight) - dequantize = getattr(module, f"dequantize_{num_bits}") - - reference = dequantize(packed, shape, dtype=weight.dtype, block_chunk_size=1) - for chunk in (2, 7, 4096): - assert torch.equal( - dequantize(packed, shape, dtype=weight.dtype, block_chunk_size=chunk), reference - ) +def test_registry_lists_every_exported_encoder(): + """An encoder the package exports but the registry omits would be unreachable by dispatch.""" + encoders = { + name.removeprefix("quantize_") for name in ggml.__all__ if name.startswith("quantize_") + } + assert encoders == set(IQ_FORMAT_REGISTRY) + + +@pytest.mark.parametrize(("num_bits", "module"), FORMAT_MODULES) +def test_registry_record_is_wired_to_its_own_codec(num_bits, module): + """Each record points at its own format's encoder, decoder, geometry and chunk defaults.""" + record = IQ_FORMAT_REGISTRY[num_bits] + upper = num_bits.upper() + + assert record.name == num_bits + assert record.quantize is getattr(module, f"quantize_{num_bits}") + assert record.dequantize is getattr(module, f"dequantize_{num_bits}") + assert record.block_size == getattr(module, f"{upper}_BLOCK_SIZE") + assert record.block_bytes == getattr(module, f"{upper}_BLOCK_BYTES") + assert record.effective_bits == pytest.approx(getattr(module, f"{upper}_EFFECTIVE_BITS")) + assert record.block_chunk_size == module._DEFAULT_BLOCK_CHUNK_SIZE + assert record.decode_chunk_size == module._DEFAULT_DECODE_CHUNK_SIZE + + +@pytest.mark.parametrize(("num_bits", "module"), FORMAT_MODULES) +def test_public_fake_quant_is_the_registered_record(num_bits, module): + """The per-format entry point and backend dispatch run the same code path.""" + assert getattr(module, f"{num_bits}_fake_quant").__self__ is IQ_FORMAT_REGISTRY[num_bits] + + +@pytest.mark.parametrize("num_bits", FORMAT_NAMES) +def test_format_fake_quant_rejects_another_formats_quantizer(num_bits): + """Calling one format's fake quant with a quantizer configured for another is refused. + + Dispatch picks the record by num_bits, so it never reaches this guard; it protects direct + callers of a record or of a per-format ``_fake_quant`` alias. + """ + other = next(name for name in FORMAT_NAMES if name != num_bits) + expected = f"The ggml {num_bits.upper()} backend requires num_bits={num_bits!r}" + + with pytest.raises(ValueError, match=re.escape(expected)): + IQ_FORMAT_REGISTRY[num_bits].fake_quant(torch.ones(1, 256), SimpleNamespace(num_bits=other)) diff --git a/tests/unit/torch/quantization/test_iq_formats.py b/tests/unit/torch/quantization/test_iq_formats.py index f59c4566531..113effb0585 100644 --- a/tests/unit/torch/quantization/test_iq_formats.py +++ b/tests/unit/torch/quantization/test_iq_formats.py @@ -33,6 +33,7 @@ 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 +from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY from modelopt.torch.quantization.nn import TensorQuantizer # name -> (module, packed bytes per block, codebook entries, bits per weight) @@ -272,3 +273,8 @@ def test_error_decreases_with_bit_width(): errors.append(float((quantizer(weight) - weight).square().mean())) assert errors == sorted(errors, reverse=True), dict(zip(sorted(NAMES), 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)