Files
Model-Optimizer/tests/unit/torch/quantization/test_quantize_cpu.py
T
Shengliang Xu 1cceb950d6 [OMNIML-3689] PTQ quant_cfg semantic correction. Design in doc _quant_cfg.rst (#1094)
### What does this PR do?

#### Summary

Redesigns the `quant_cfg` configuration format in ModelOpt's PyTorch
quantization stack, replacing the previous dict-based format with an
**ordered list of typed `QuantizerCfgEntry` dicts**.

##### Motivation

The old `quant_cfg` dict had several pain points:
- **Ambiguous precedence**: no explicit way to reason about which entry
wins when multiple keys match a quantizer
- **Mixed key namespaces**: wildcard paths and PyTorch class names lived
in the same dict level, requiring ad-hoc dispatch
- **Magic `"default"` key**: an implicit, undocumented catch-all that
was easy to misuse
- **Poor composability**: merging two configs required dict updates that
silently discarded keys
- **No YAML round-trip fidelity**: the nested structure couldn't be
expressed cleanly in YAML

##### New format

`quant_cfg` is now an ordered list of `QuantizerCfgEntry` TypedDicts.
Each entry has:
- `quantizer_name` *(required)*: `fnmatch` wildcard matched against
quantizer module names
- `cfg` *(optional)*: dict (or list of dicts) of
`QuantizerAttributeConfig` fields
- `enable` *(optional)*: toggles quantizer on/off independently of `cfg`
- `parent_class` *(optional)*: restricts match to quantizers whose
parent module is of the given PyTorch class (e.g. `"nn.BatchNorm2d"`)

Entries are applied in list order; later entries override earlier ones.
The canonical pattern is deny-all first (`_base_disable_all`), then
selectively re-enable and configure, then apply standard exclusions
(`_default_disabled_quantizer_cfg`).

##### Changes

**Core library (`modelopt/torch/quantization/`)**

- **`config.py`**:
- Added `QuantizerCfgEntry` TypedDict (line 163) and
`find_quant_cfg_entry_by_path()` helper for exact-match lookup of
entries by path.
- Added `normalize_quant_cfg_list()` (line 1539) that converts legacy
formats (flat dict, single-key dicts, `nn.*`-scoped dicts, `"default"`
key) to canonical `QuantizerCfgEntry` lists. After normalization every
entry is guaranteed to have explicit `quantizer_name`, `enable`, and
`cfg` keys.
- Converted `_default_disabled_quantizer_cfg` and
`_mamba_moe_disabled_quantizer_cfg` from dicts to lists of
`QuantizerCfgEntry`.
- Added `_base_disable_all` (line 205): canonical deny-all entry
(`[{"quantizer_name": "*", "enable": False}]`).
- Converted all ~30 built-in config constants (`INT8_DEFAULT_CFG`,
`FP8_DEFAULT_CFG`, `NVFP4_DEFAULT_CFG`, etc.) to list format using
`*_base_disable_all` and `*_default_disabled_quantizer_cfg` unpacking.
- KV-cache configs (`FP8_KV_CFG`, `NVFP4_KV_CFG`, etc.) are now minimal
lists designed to be concatenated with a primary config — they
intentionally omit `_base_disable_all` and `"algorithm"`.
- Added two `QuantizeConfig` Pydantic field validators: a
`mode="before"` validator that calls `normalize_quant_cfg_list()`, and a
`mode="after"` validator that validates `cfg` dicts against
`QuantizerAttributeConfig`.
- Updated `need_calibration()` to iterate the normalized list instead of
the old dict.
- Changed `QuantizeQuantCfgType` alias from `dict[str | Callable, ...]`
to `list[QuantizerCfgEntry]`.

- **`conversion.py`**:
- Rewrote `set_quantizer_by_cfg()` (line 217) to iterate the list
directly. Each entry's `parent_class` is resolved via
`QuantModuleRegistry[parent_class_name]` (the existing `_DMRegistryCls`
registry).
- Added `set_quantizer_attributes_full()` (line 314): full replacement
of quantizer attributes from a `QuantizerAttributeConfig`. Unspecified
fields revert to defaults, enforcing entry atomicity. Can also upgrade
`TensorQuantizer` → `SequentialQuantizer` or downgrade the reverse.
- Added `set_quantizer_attributes_partial()` (line 384): merges a
partial `dict` of attributes into existing quantizer state. Does NOT
change quantizer structure. Used for enable-only entries.
- Added `set_quantizer_by_cfg_context()` context manager (line 447) that
temporarily applies a `quant_cfg` list and restores original quantizer
state on exit.
- Deprecated `set_quantizer_attribute()` (line 525) with a
`DeprecationWarning` pointing to the new functions.

- **`tensor_quantizer.py`**:
- `TensorQuantizer.set_from_attribute_config()`: narrowed type hint from
`dict` to `dict[str, Any]`.
- Added `_axis_setter` and `_block_sizes_setter` custom setters so that
`axis` and `block_sizes` changes properly propagate to the calibrator
and maintain mutual exclusivity.
- `SequentialQuantizer.set_from_attribute_config()`: narrowed signature
to `list[QuantizerAttributeConfig] | list[dict[str, Any]]` (removed the
old union with single values).

- **`algorithms.py`**:
- Updated `_match_quantizer_cfg()` to iterate the list and return
`(matched_cfg, matched_enable)` tuple with last-match-wins.
- Updated `_cfg_to_dict()`, `estimate_quant_compression()`, and
`QuantRecipe` to work with the list-based format.
- Updated `get_auto_quantize_config()` to emit list-format `quant_cfg`.

- **`model_quant.py`**: `disable_quantizer()` / `enable_quantizer()` now
call `set_quantizer_attributes_partial()` directly instead of the
deprecated `set_quantizer_attribute()`. Updated docstrings and code
examples to show the list format.

- **`utils/core_utils.py`**: `disable_lora_quantizers_in_config()` and
`update_quant_cfg_with_kv_cache_quant()` updated to append
`QuantizerCfgEntry` dicts to the list.

- **Other**: minor updates to `backends/fp8_per_tensor_gemm.py`,
`backends/nvfp4_gemm.py`, `compress.py`, `model_calib.py`,
`export/unified_export_hf.py`, and
`sparsity/attention_sparsity/conversion.py` to use the list format.

- **`onnx/llm_export_utils/quantization_utils.py`**: Updated
quantization config construction to use list format.

**YAML recipes (`modelopt_recipes/`)**

- Converted all 5 general PTQ recipes to the new list format:
  - `general/ptq/fp8_default-fp8_kv.yml`
  - `general/ptq/nvfp4_default-fp8_kv.yml`
  - `general/ptq/nvfp4_experts_only-fp8_kv.yml`
  - `general/ptq/nvfp4_mlp_only-fp8_kv.yml`
  - `general/ptq/nvfp4_omlp_only-fp8_kv.yml`
- Converted model-specific recipe:
`models/Step3.5-Flash/nvfp4-mlp-only.yaml`

**Documentation (`docs/`)**

- New guide: `docs/source/guides/_quant_cfg.rst` — comprehensive
reference covering entry format, ordering semantics, entry atomicity,
`enable` vs `cfg` independence, `parent_class` filtering, and common
patterns (deny-all-then-enable, customizing a built-in config, building
from scratch).
- Updated `_pytorch_quantization.rst` code examples to show the list
format with `copy.deepcopy` and `.append()`.
- Added `_quant_cfg.rst` to the quantization guide table of contents.

**Examples**

- Updated all quantization examples to use the list format:
`deepseek/ptq.py`, `diffusers/quantization/config.py`,
`llm_ptq/hf_ptq.py`, `llm_qat/main.py`, `vllm_serve/vllm_ptq_utils.py`,
`llm_autodeploy/run_auto_quantize.py`, `llm_eval/quantization_utils.py`,
`llm_ptq/example_utils.py`,
`windows/torch_onnx/diffusers/qad_example/sample_example_qad_diffusers.py`,
and 2 notebooks.

**Tests**

- New test file:
`tests/unit/torch/quantization/test_config_validation.py` — unit tests
for `need_calibration()`, `normalize_quant_cfg_list()` (new format,
legacy format conversions, error cases),
`find_quant_cfg_entry_by_path()`, `_match_quantizer_cfg()`, and
`QuantizeConfig` Pydantic validators.
- Extended `tests/unit/torch/quantization/test_quantize_cpu.py` with
tests for `set_quantizer_attributes_full()` (atomicity, parent_class
filtering, SequentialQuantizer creation), list ordering, enable-only
entry behavior, and end-to-end legacy dict format.
- Updated 20+ existing test files across `tests/unit/`, `tests/gpu/`,
`tests/gpu_megatron/`, and `tests/_test_utils/` to use the list format.

##### Backward compatibility

`normalize_quant_cfg_list()` is called automatically by the
`QuantizeConfig` Pydantic `mode="before"` validator, so existing code
passing the old dict-based format (flat dict like `{"*weight_quantizer":
{"num_bits": 8}}`, single-key dict lists, or `nn.*`-scoped dicts with
`parent_class` semantics) continues to work without modification. The
legacy `"default"` key is converted to `quantizer_name: "*"`.

`set_quantizer_attribute()` is preserved as a deprecated wrapper around
`set_quantizer_attributes_partial()`.

#### Test coverage

- **Unit tests**: new `test_config_validation.py` with tests for
normalization, validation, path lookup, and cfg matching. Extended
`test_quantize_cpu.py` with tests for full/partial attribute setting,
ordering, atomicity, and legacy backward compatibility.
- **System testing**:

```
python examples/llm_ptq/hf_ptq.py \
      --model Qwen/Qwen3-8B  \
      --recipe general/ptq/fp8_default-fp8_kv \
      --export_path=build/fp8_default-fp8_kv42  \
      --calib_size=16 \
      --batch_size=0 \
      --trust_remote_code \
      --export_fmt=hf
```

### Additional Information
<!-- E.g. related issue. -->

---------

Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
2026-04-06 15:38:44 -07:00

490 lines
18 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 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.
"""High-level tests for quantization."""
import copy
import pytest
import torch
from _test_utils.torch.quantization.models import SimpleConv, SimpleConvLinear, SimpleLinear
from _test_utils.torch.quantization.quantize_common import (
INT4_AWQ_CLIP_CFG,
INT4_AWQ_FULL_CFG,
INT4_SVDQUANT_CFG,
quantize_model_and_forward,
save_restore_test,
)
from pydantic import ValidationError
import modelopt.torch.opt as mto
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.calib import MaxCalibrator
from modelopt.torch.quantization.config import QuantizerAttributeConfig
from modelopt.torch.quantization.conversion import set_quantizer_attributes_full
from modelopt.torch.quantization.nn.modules.tensor_quantizer import (
SequentialQuantizer,
TensorQuantizer,
)
# A test config with double-quant (using `SequentialQuantizers`)
WINT4INT8_CFG = {
"quant_cfg": [
{
"quantizer_name": "*weight_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": 8, "axis": 0},
],
"enable": True,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None},
"enable": True,
},
],
"algorithm": "awq_lite",
}
# Test configs for per channel MSE calibration
INT8_MSE_CFG = {
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
],
"algorithm": "mse",
}
STATIC_WEIGHT_DYNAMIC_ACTIVATION_CFG = {
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 8, "axis": 0},
}, # Per-channel quantization
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": (0, 1), "type": "dynamic"},
}, # Dynamic per-token quantization
],
"algorithm": "max",
}
class NewMaxCalibrator(MaxCalibrator):
def compute_amax(self):
return 2 * self._calib_amax
quant_cfg_custom_calib = {
"quant_cfg": [
{
"quantizer_name": "*",
"cfg": {
"num_bits": 4,
"axis": None,
"calibrator": (NewMaxCalibrator, (4, None, False)),
},
"enable": True,
}
],
"algorithm": "max",
}
@pytest.mark.parametrize("model_cls", [SimpleLinear, SimpleConv, SimpleConvLinear])
@pytest.mark.parametrize(
"config",
[
mtq.INT8_DEFAULT_CFG,
mtq.INT8_SMOOTHQUANT_CFG,
mtq.INT4_BLOCKWISE_WEIGHT_ONLY_CFG,
mtq.INT4_AWQ_CFG,
INT4_SVDQUANT_CFG,
INT4_AWQ_CLIP_CFG,
INT4_AWQ_FULL_CFG,
WINT4INT8_CFG,
INT8_MSE_CFG,
],
)
def test_quantize(model_cls, config):
"""Test quantize function can run without problems."""
model = model_cls()
calib_data = [model.get_input() for _ in range(2)]
quantize_model_and_forward(model, config, calib_data)
# For fast testing, lets just test one config
if config == mtq.INT8_DEFAULT_CFG:
mtq.print_quant_summary(model)
@pytest.mark.parametrize(
("model_cls", "quant_config"),
[
(SimpleLinear, mtq.INT8_SMOOTHQUANT_CFG),
(SimpleConvLinear, quant_cfg_custom_calib),
(SimpleConvLinear, mtq.INT8_DEFAULT_CFG),
(SimpleLinear, INT4_SVDQUANT_CFG),
],
)
def test_save_restore(model_cls, quant_config):
save_restore_test(model_cls, "cpu", quant_config)
def test_quantize_invalid_cfg():
model = SimpleLinear()
config_invalid = {
"quant_cfg": [
{"quantizer_name": "*", "cfg": {"num_bits": 4, "axis": 0, "block_sizes": {-1: 128}}}
],
"algorithm": "max",
}
with pytest.raises(ValidationError, match="axis must be None when block_sizes is not None."):
model = mtq.quantize(model, config_invalid)
def test_inplace_backward_compatibility():
model = SimpleLinear()
calib_data = [model.get_input() for _ in range(2)]
def forward_loop():
for batch in calib_data:
model(batch)
mtq.quantize(model, mtq.INT8_DEFAULT_CFG, forward_loop=forward_loop)
def test_custom_calib_config():
model_ref = SimpleLinear()
model_ref = mtq.quantize(
model_ref, quant_cfg_custom_calib, lambda model: model(model.get_input())
)
model_quant = SimpleLinear()
model_quant = mto.restore_from_modelopt_state(model_quant, mto.modelopt_state(model_ref))
model_quant.load_state_dict(model_ref.state_dict())
inputs = model_ref.get_input()
assert torch.allclose(model_ref(inputs), model_quant(inputs))
for name, module in model_quant.named_modules():
if name.endswith("quantizer"):
assert module._calibrator.__class__ == NewMaxCalibrator
def test_class_wise_config():
model = SimpleConvLinear()
config = {
"quant_cfg": [
{
"parent_class": "nn.Linear",
"quantizer_name": "*",
"cfg": {"num_bits": 4, "axis": -1},
"enable": True,
},
{
"parent_class": "nn.Conv2d",
"quantizer_name": "*",
"cfg": {"num_bits": 8},
"enable": True,
},
{"parent_class": "nn.BatchNorm2d", "quantizer_name": "*", "enable": False},
{"quantizer_name": "*output_quantizer", "cfg": {"num_bits": 8}, "enable": True},
],
"algorithm": "max",
}
model = mtq.quantize(model, config, lambda model: model(model.get_input()))
for name, module in model.named_modules():
if isinstance(module, torch.nn.Linear):
for sub_quantizer in (module.weight_quantizer, module.input_quantizer):
assert sub_quantizer.num_bits == 4
assert sub_quantizer.axis == -1
assert sub_quantizer.is_enabled
elif isinstance(module, torch.nn.Conv2d):
for sub_quantizer in (module.weight_quantizer, module.input_quantizer):
assert sub_quantizer.num_bits == 8
assert sub_quantizer.is_enabled
elif isinstance(module, torch.nn.BatchNorm2d):
assert module.input_quantizer.is_enabled is False
if name.endswith("output_quantizer"):
assert module.is_enabled
assert module.num_bits == 8
def test_static_weight_dynamic_activations():
model = SimpleLinear()
inputs = model.get_input()
model = mtq.quantize(
model, STATIC_WEIGHT_DYNAMIC_ACTIVATION_CFG, lambda model: model(model.get_input())
)
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.amax is not None
# Test that model forward works
model(inputs)
# Lets test mtq.quantize without forward_loop
model = SimpleLinear()
model = mtq.quantize(model, STATIC_WEIGHT_DYNAMIC_ACTIVATION_CFG)
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.amax is not None
def test_block_sizes_axis_model():
REF_QUANT_CFG = { # noqa: N806
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None, "type": "dynamic"},
},
],
"algorithm": "max",
}
QUANT_CFG = { # noqa: N806
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 8, "block_sizes": {1: None}},
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "block_sizes": {0: None, 1: None}, "type": "dynamic"},
},
],
"algorithm": "max",
}
model_ref = SimpleLinear()
model = copy.deepcopy(model_ref)
inputs = model_ref.get_input()
mtq.quantize(model_ref, REF_QUANT_CFG, lambda model: model(inputs))
mtq.quantize(model, QUANT_CFG, lambda model: model(inputs))
assert torch.allclose(model_ref(inputs), model(inputs))
# compare the calibrated amax of all quantizers
for (name_ref, module_ref), (name, module) in zip(
model_ref.named_modules(), model.named_modules()
):
if hasattr(module, "weight_quantizer"):
assert name_ref == name
assert torch.allclose(module_ref.weight_quantizer.amax, module.weight_quantizer.amax)
def test_quantize_twice():
"""Test that calling mtq.quantize twice on the same model works."""
model = SimpleLinear()
inputs = model.get_input()
def forward_loop(model):
return model(inputs)
model = mtq.quantize(model, mtq.INT8_DEFAULT_CFG, forward_loop=forward_loop)
out1 = model(inputs)
model = mtq.quantize(model, mtq.INT8_DEFAULT_CFG, forward_loop=forward_loop)
out2 = model(inputs)
assert torch.allclose(out1, out2), "Re-quantization with same config should be idempotent"
class TestSetQuantizerAttributesFull:
"""Tests for set_quantizer_attributes_full and its atomicity semantics."""
def _quantize(self, model):
return mtq.quantize(model, mtq.INT8_DEFAULT_CFG, lambda m: m(m.get_input()))
def test_basic_full_replacement(self):
"""set_quantizer_attributes_full replaces all attributes on matched quantizers."""
model = self._quantize(SimpleLinear())
attrs = QuantizerAttributeConfig(num_bits=4, axis=0)
set_quantizer_attributes_full(model, "*weight_quantizer", attrs)
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert isinstance(module, TensorQuantizer)
assert module.num_bits == 4
assert module.axis == 0
def test_atomicity_unset_fields_revert_to_defaults(self):
"""A full replacement reverts unspecified fields to QuantizerAttributeConfig defaults."""
model = self._quantize(SimpleLinear())
# First configure with axis=0 (non-default)
set_quantizer_attributes_full(
model, "*weight_quantizer", QuantizerAttributeConfig(num_bits=8, axis=0)
)
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.axis == 0
# Now replace with only num_bits=4; axis should revert to default (None)
set_quantizer_attributes_full(
model, "*weight_quantizer", QuantizerAttributeConfig(num_bits=4)
)
default_axis = QuantizerAttributeConfig().axis
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.num_bits == 4
assert module.axis == default_axis
def test_parent_class_filter(self):
"""parent_class restricts which quantizers are affected."""
model = self._quantize(SimpleConvLinear())
# Only set num_bits=4 for quantizers inside nn.Linear modules
set_quantizer_attributes_full(
model,
"*weight_quantizer",
QuantizerAttributeConfig(num_bits=4),
parent_class=torch.nn.Linear,
)
for name, module in model.named_modules():
if not name.endswith("weight_quantizer"):
continue
parent_name = name.rpartition(".")[0]
parent = model.get_submodule(parent_name)
if isinstance(parent, torch.nn.Linear):
assert module.num_bits == 4
else:
# Conv2d weight_quantizers should be unchanged (still 8-bit from INT8_DEFAULT_CFG)
assert module.num_bits == 8
def test_wildcard_no_match_is_noop(self):
"""A wildcard that matches nothing silently does nothing."""
model = self._quantize(SimpleLinear())
# Record state before
bits_before = {
n: m.num_bits for n, m in model.named_modules() if isinstance(m, TensorQuantizer)
}
set_quantizer_attributes_full(
model, "*nonexistent_quantizer*", QuantizerAttributeConfig(num_bits=4)
)
bits_after = {
n: m.num_bits for n, m in model.named_modules() if isinstance(m, TensorQuantizer)
}
assert bits_before == bits_after
def test_invalid_attributes_type_raises(self):
"""Passing a plain dict instead of QuantizerAttributeConfig raises ValueError."""
model = self._quantize(SimpleLinear())
with pytest.raises((ValueError, AttributeError)):
set_quantizer_attributes_full(model, "*weight_quantizer", {"num_bits": 4}) # type: ignore[arg-type]
def test_list_attributes_creates_sequential_quantizer(self):
"""A list of QuantizerAttributeConfig replaces TensorQuantizer with SequentialQuantizer."""
model = self._quantize(SimpleLinear())
attrs = [
QuantizerAttributeConfig(num_bits=4, block_sizes={-1: 128}),
QuantizerAttributeConfig(num_bits=8, axis=0),
]
set_quantizer_attributes_full(model, "*weight_quantizer", attrs)
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert isinstance(module, SequentialQuantizer)
assert len(module) == 2
def test_ordering_later_entry_overrides_earlier():
"""Later entries in quant_cfg override earlier ones for the same quantizer."""
model = SimpleLinear()
config = {
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
],
"algorithm": "max",
}
model = mtq.quantize(model, config, lambda m: m(m.get_input()))
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.num_bits == 4, "Later entry (num_bits=4) should override earlier (8)"
if name.endswith("input_quantizer"):
assert module.num_bits == 8
def test_enable_only_entry_preserves_attributes():
"""An enable-only entry toggles the quantizer without resetting its attributes."""
model = SimpleLinear()
config = {
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
# This enable-only entry should disable without resetting num_bits/axis
{"quantizer_name": "*weight_quantizer", "enable": False},
],
"algorithm": "max",
}
model = mtq.quantize(model, config, lambda m: m(m.get_input()))
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert not module.is_enabled, "weight_quantizer should be disabled"
assert module.num_bits == 4, "num_bits should be preserved by enable-only entry"
assert module.axis == 0, "axis should be preserved by enable-only entry"
def test_atomicity_later_cfg_entry_does_not_inherit_earlier():
"""When two cfg-bearing entries match the same quantizer, the second fully replaces the first."""
model = SimpleLinear()
config = {
"quant_cfg": [
# Entry 1: set axis=0
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
# Entry 2: only set num_bits=4, no axis — axis should revert to default (None), not 0
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
],
"algorithm": "max",
}
model = mtq.quantize(model, config, lambda m: m(m.get_input()))
default_axis = QuantizerAttributeConfig().axis
for name, module in model.named_modules():
if name.endswith("weight_quantizer"):
assert module.num_bits == 4
assert module.axis == default_axis, (
f"axis should revert to default ({default_axis}), not inherit 0 from earlier entry"
)
def test_legacy_dict_format_end_to_end():
"""Old dict-format quant_cfg works end-to-end through mtq.quantize via normalization."""
model = SimpleLinear()
# Old-style dict config with "default" key and wildcard keys
old_config = {
"quant_cfg": {
"default": {"enable": False},
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
},
"algorithm": "max",
}
model = mtq.quantize(model, old_config, lambda m: m(m.get_input()))
for name, module in model.named_modules():
if isinstance(module, TensorQuantizer):
if name.endswith(("weight_quantizer", "input_quantizer")):
assert module.is_enabled
assert module.num_bits == 8
elif name.endswith("output_quantizer"):
# "default" key → quantizer_name="*" with enable=False disables everything,
# but weight/input quantizers are re-enabled by subsequent entries.
# output_quantizer is NOT re-enabled so it stays disabled.
assert not module.is_enabled