[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>
This commit is contained in:
Shengliang Xu
2026-04-06 15:38:44 -07:00
committed by GitHub
parent c542c09b11
commit 1cceb950d6
62 changed files with 3436 additions and 1395 deletions
+4
View File
@@ -14,6 +14,10 @@ NVIDIA Model Optimizer Changelog
- Add support for vLLM fakequant reload using ModelOpt state for HF models. See `examples/vllm_serve/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/vllm_serve#load-qatptq-model-and-serve-in-vllm-wip>`_ for more details.
- [Early Testing] Add Claude Code PTQ skill (``.claude/skills/ptq/``) for agent-assisted post-training quantization. The skill guides the agent through environment detection, model support checking, format selection, and execution via the launcher or manual SLURM/Docker/bare GPU paths. Includes handling for unlisted models with custom module patching. This feature is in early testing — use with caution.
**Backward Breaking Changes**
- The ``quant_cfg`` field in quantization configs is now an **ordered list** of ``QuantizerCfgEntry`` dicts instead of a flat dictionary. Each entry specifies a ``quantizer_name`` wildcard, an optional ``parent_class`` filter, a ``cfg`` dict of quantizer attributes, and/or an ``enable`` flag. Entries are applied in list order with later entries overriding earlier ones. The old dict-based format is still accepted and automatically converted via ``normalize_quant_cfg_list()``, but now emits a ``DeprecationWarning``; new code should use the list format. All built-in configs (e.g. ``FP8_DEFAULT_CFG``, ``INT4_AWQ_CFG``, ``NVFP4_DEFAULT_CFG``), examples, and YAML recipes have been updated. See the :ref:`quant-cfg` documentation for the new format reference and migration guide.
**Bug Fixes**
- Fix Minitron pruning (``mcore_minitron``) for MoE models. Importance estimation hooks were incorrectly registered for MoE modules and NAS step was hanging before this.
+1
View File
@@ -19,6 +19,7 @@ Below, you can find the documentation for the quantization toolkit in ModelOpt:
./_basic_quantization.rst
./_choosing_quant_methods.rst
./_pytorch_quantization.rst
./_quant_cfg.rst
./_customized_model_quantization.rst
./_compress_quantized_models.rst
./_onnx_quantization.rst
+21 -12
View File
@@ -237,14 +237,16 @@ For debugging purposes or simple customizations, you can modify an existing conf
.. code-block:: python
# Create a copy of the default INT8 configuration
config = mtq.INT8_DEFAULT_CFG.copy()
import copy
# Disable input quantizers for all layers
config["quant_cfg"]["*input_quantizer"]["enable"] = False
# Create a deep copy of the default INT8 configuration
config = copy.deepcopy(mtq.INT8_DEFAULT_CFG)
# Disable input quantizers for all layers (appended last, so it takes precedence)
config["quant_cfg"].append({"quantizer_name": "*input_quantizer", "enable": False})
# Disable all quantizers for layers matching the pattern "layer1.*"
config["quant_cfg"]["*layer1.*"] = {"enable": False}
config["quant_cfg"].append({"quantizer_name": "*layer1.*", "enable": False})
Advanced Configuration Creation
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
@@ -253,18 +255,23 @@ For exploring new quantization recipes, you can compose a completely new configu
.. code-block:: python
from modelopt.torch.quantization.config import _default_disabled_quantizer_cfg
# Custom configuration for INT4 block-wise weights and INT8 dynamic activations
MY_CUSTOM_CONFIG = {
"quant_cfg": {
"quant_cfg": [
# Disable all quantizers by default, then enable selectively
{"quantizer_name": "*", "enable": False},
# Configure weight quantizers with 4-bit precision and 128-element blocks
"*weight_quantizer": {"num_bits": 4, "block_sizes": {-1: 128}, "enable": True},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4, "block_sizes": {-1: 128}}, "enable": True},
# Configure input quantizers with 8-bit dynamic quantization
"*input_quantizer": {"num_bits": 8, "type": "dynamic", "block_sizes": {-1: None}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "type": "dynamic", "block_sizes": {-1: None}}},
# Include default disabled quantizer configurations
**_default_disabled_quantizer_cfg,
},
*_default_disabled_quantizer_cfg,
],
"algorithm": "max",
}
@@ -394,8 +401,10 @@ You can specify ``custom_calib`` as ``algorithm`` in ``quant_cfg`` to use it. He
# create quantization configuration with "custom_calib" method
quant_cfg = {
'quant_cfg': {'*weight_quantizer': ..},
'algorithm': {"method": 'custom_calib'},
'quant_cfg': [
{"quantizer_name": "*weight_quantizer", "cfg": {...}},
],
'algorithm': {"method": 'custom_calib'},
}
+393
View File
@@ -0,0 +1,393 @@
.. _quant-cfg:
======================================
Quantization Configuration (quant_cfg)
======================================
The ``quant_cfg`` field is the primary mechanism for controlling which quantizers are active in a
model and how they are configured. This guide explains the format, ordering semantics, and common
patterns for composing quantization configurations.
.. tip::
For the list of built-in configs and supported formats, see :any:`quantization-formats`.
For how to apply a config to a model, see :any:`_pytorch_quantization`.
----------
Overview
========
A quantization config is a Python dictionary with two top-level keys:
.. code-block:: python
config = {
"quant_cfg": [...], # ordered list of QuantizerCfgEntry dicts
"algorithm": "max", # calibration algorithm
}
The ``quant_cfg`` value is an **ordered list** of :class:`QuantizerCfgEntry
<modelopt.torch.quantization.config.QuantizerCfgEntry>` dicts. Each entry targets a set of
quantizer modules in the model and specifies their configuration.
----------
Entry Format
============
Each entry in the list is a dictionary with the following fields:
.. list-table::
:header-rows: 1
:widths: 20 15 65
* - Field
- Required
- Description
* - ``quantizer_name``
- Yes
- Wildcard string matched against quantizer module names (e.g. ``"*weight_quantizer"``).
Uses :func:`fnmatch` rules.
* - ``parent_class``
- No
- Restricts matching to quantizers whose immediate parent module is of this PyTorch class
(e.g. ``"nn.Linear"``). If omitted, all modules are targeted regardless of class.
* - ``cfg``
- No
- A dict of quantizer attributes as defined by :class:`QuantizerAttributeConfig
<modelopt.torch.quantization.config.QuantizerAttributeConfig>`, or a list of such dicts
for sequential quantization (see :ref:`sequential-quantizers`).
* - ``enable``
- No
- ``True`` or ``False``. Toggles matched quantizers on or off, independently of ``cfg``.
When ``cfg`` is absent, **only** the enabled/disabled state is changed — all other
attributes remain untouched. When ``cfg`` is present, ``enable`` sets the enabled state
of the newly-configured quantizer. When ``cfg`` is present and ``enable`` is omitted,
the quantizer is implicitly enabled (``True``).
.. note::
Every entry must specify at least one of ``cfg`` or ``enable`` in addition to
``quantizer_name``. An entry with only ``quantizer_name`` and no other keys is **invalid**
and will raise a ``ValueError`` at config-processing time. This prevents subtle bugs where
a bare ``{"quantizer_name": "*"}`` would silently behave as ``enable=True`` for all
quantizers.
----------
Default Quantizer Configuration
================================
When a quantizer is enabled but has never been touched by a ``cfg`` entry — either because no
entry in the list matched it, or because it was only reached by enable-only entries — it operates
with the default attributes of
:class:`QuantizerAttributeConfig <modelopt.torch.quantization.config.QuantizerAttributeConfig>`:
.. code-block:: python
{
"num_bits": 8, # 8-bit integer quantization
"axis": None, # per-tensor scale (no per-channel axis)
"fake_quant": True, # simulate quantization in forward pass (PTQ / QAT)
"unsigned": False, # signed integer range, e.g. [-128, 127] for INT8
"narrow_range": False, # full range; True would restrict to [-127, 127] for INT8
"type": "static", # static calibration (not dynamic per-inference)
"block_sizes": None, # no block quantization; set for NF4 / MXFP formats
"bias": None, # no affine bias correction
"calibrator": "max", # use max-abs calibration to determine amax
"rotate": False, # no Hadamard rotation (QuaRot / SpinQuant)
"pass_through_bwd": True, # straight-through estimator for QAT gradients
"trt_high_precision_dtype": "Float", # cast QDQ nodes to fp32 for TRT StronglyType export
"backend": None, # use the built-in quantization backend
"backend_extra_args": None, # no extra args for custom backends
"use_constant_amax": False, # calibrate amax; True hard-codes FP8 E4M3 max (448.0)
}
In practice this means an un-configured but enabled quantizer performs **INT8 per-tensor static
fake-quantization** with a max-calibrated scale. This is rarely the intended behavior — every
quantizer you want active should be explicitly configured with a ``cfg`` entry.
----------
Ordering and Precedence
=======================
Entries are applied **in list order**. Later entries override earlier ones for any quantizer they
match. This gives a clear, composable precedence model:
- Put broad rules (e.g. deny-all) **first**.
- Put format-specific enable rules **after**.
- Put fine-grained exclusions (specific layers, classes) **last**.
The recommended pattern used by all built-in configs is:
.. code-block:: python
"quant_cfg": [
# 1. Deny all quantizers by default
{"quantizer_name": "*", "enable": False},
# 2. Enable and configure the target quantizers
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
# 3. Apply standard exclusions last (BatchNorm, LM head, MoE routers, etc.)
*mtq.config._default_disabled_quantizer_cfg,
]
.. note::
The deny-all entry ``{"quantizer_name": "*", "enable": False}`` is available as
:data:`modelopt.torch.quantization.config._base_disable_all` and is prepended to every
built-in config. This ensures quantizers not explicitly targeted remain disabled.
----------
Entry Atomicity
===============
Each ``cfg``-bearing entry in ``quant_cfg`` is a **complete, self-contained configuration unit**.
When an entry with ``cfg`` matches a quantizer, it **completely replaces** that quantizer's
configuration — it does not merge with or incrementally update settings left by earlier entries.
Concretely, if an entry specifies only a subset of quantizer attributes (e.g. only ``num_bits``),
all unspecified attributes are filled in with their default values from
:class:`QuantizerAttributeConfig <modelopt.torch.quantization.config.QuantizerAttributeConfig>`.
The resulting *complete* config is then written to the quantizer, discarding whatever any prior
matching entry had set.
This means:
- **Last cfg-entry wins, fully.** If two entries both match ``*weight_quantizer`` and both carry
a ``cfg``, the second entry does not inherit the first entry's settings — it replaces them entirely.
- **No hidden state accumulation.** The final configuration of a quantizer depends only on the
*last* ``cfg``-bearing entry in the list that matched it, making behavior easy to reason about.
- **Changing one field requires a full spec.** Because each ``cfg`` entry is a complete replacement,
to change only one attribute of a quantizer that was already configured, you must reproduce the
full desired config in the new entry. Any attribute omitted from the entry will revert to its
default, not to the value set by an earlier entry.
**Enable-only entries are the exception.** An entry with no ``cfg`` (only ``enable``) is *not* a
full replacement — it solely flips the on/off state of matched quantizers, leaving all other
attributes unchanged:
- ``{"quantizer_name": "*", "enable": False}`` disables all quantizers without touching their
configured attributes. Use this as the first step in a deny-all-then-configure pattern.
- ``{"quantizer_name": "*weight_quantizer", "enable": True}`` (no ``cfg``) re-enables weight
quantizers using whatever attributes they currently carry (or their defaults if they were never
configured by a ``cfg`` entry).
For example, given the following two entries both matching ``*weight_quantizer``:
.. code-block:: python
# Entry 1 — sets FP8 per-channel
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": 0}},
# Entry 2 — sets INT4 blockwise (axis is NOT inherited from Entry 1)
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4, "block_sizes": {-1: 128}}},
After Entry 2 is applied, the quantizer has ``num_bits=4``, ``block_sizes={-1: 128}``, and
``axis=None`` (the default). The ``axis=0`` set by Entry 1 is gone.
.. note::
The deny-all-then-configure pattern is safe and predictable precisely because
``{"quantizer_name": "*", "enable": False}`` **only** disables quantizers without resetting
their attributes. Subsequent ``cfg`` entries then configure targets from a known default state.
----------
Common Patterns
===============
Skipping Specific Layers
------------------------
Append a disable entry after the existing config to exclude layers matched by a path pattern.
Because it is appended last, it takes precedence over all earlier entries:
.. code-block:: python
import copy
import modelopt.torch.quantization as mtq
config = copy.deepcopy(mtq.FP8_DEFAULT_CFG)
# Skip the final projection layer
config["quant_cfg"].append({"quantizer_name": "*lm_head*", "enable": False})
model = mtq.quantize(model, config, forward_loop)
Skipping Layers by Module Class
--------------------------------
Use ``parent_class`` to target quantizers only within a specific type of layer, leaving the
same quantizer path in other layer types unaffected:
.. code-block:: python
config["quant_cfg"].append({
"quantizer_name": "*input_quantizer",
"parent_class": "nn.LayerNorm",
"enable": False,
})
Overriding Quantizer Precision for Specific Layers
---------------------------------------------------
A later entry with a matching ``quantizer_name`` replaces the configuration set by an earlier
entry. This allows per-layer precision overrides without restructuring the entire config:
.. code-block:: python
config = copy.deepcopy(mtq.FP8_DEFAULT_CFG)
# Quantize attention output projections in higher-precision INT8 instead of FP8
config["quant_cfg"].append({
"quantizer_name": "*o_proj*weight_quantizer",
"cfg": {"num_bits": 8, "axis": 0},
})
Building a Config from Scratch
-------------------------------
For entirely custom recipes, compose the list directly:
.. code-block:: python
from modelopt.torch.quantization.config import _base_disable_all, _default_disabled_quantizer_cfg
MY_CUSTOM_CFG = {
"quant_cfg": [
*_base_disable_all,
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4, "block_sizes": {-1: 128}}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
*_default_disabled_quantizer_cfg,
],
"algorithm": "max",
}
model = mtq.quantize(model, MY_CUSTOM_CFG, forward_loop)
----------
.. _sequential-quantizers:
Sequential Quantization
=======================
When ``cfg`` is a **list** of attribute dicts, the matched
:class:`TensorQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.TensorQuantizer>`
is replaced with a
:class:`SequentialQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.SequentialQuantizer>`
that applies each format in sequence. This is used, for example, in W4A8 quantization where weights
are quantized first in INT4 and then in FP8:
.. code-block:: python
{
"quantizer_name": "*weight_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": (4, 3)}, # FP8
],
}
----------
.. _migrating-from-dict-format:
Migrating from Dict Format
===========================
Earlier versions of ModelOpt used a flat dictionary for ``quant_cfg``. The new list format is
preferred because it provides explicit ordering and unambiguous precedence. Existing dict-based
configs continue to work — the normalization layer converts them automatically — but new code
should use the list format.
The table below shows common patterns and their list equivalents:
.. list-table::
:header-rows: 1
:widths: 50 50
* - Legacy dict format
- New list format
* - .. code-block:: python
"quant_cfg": {
"*weight_quantizer": {
"num_bits": 8,
"axis": 0,
},
"*input_quantizer": {
"num_bits": 8,
"axis": None,
},
"default": {"enable": False},
}
- .. code-block:: python
"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}},
]
* - .. code-block:: python
# Disable by key assignment
config["quant_cfg"]["*lm_head*"] = {
"enable": False,
}
- .. code-block:: python
# Append to the end (last entry wins)
config["quant_cfg"].append(
{"quantizer_name": "*lm_head*",
"enable": False}
)
* - .. code-block:: python
# Class-scoped entry
"quant_cfg": {
"nn.Linear": {
"*input_quantizer": {
"enable": False,
},
},
}
- .. code-block:: python
"quant_cfg": [
{"quantizer_name": "*input_quantizer",
"parent_class": "nn.Linear",
"enable": False},
]
Key differences to keep in mind:
- The ``"default"`` key becomes ``{"quantizer_name": "*", "enable": False}`` placed at the
**start** of the list (deny-all-then-configure pattern).
- Dict key assignment (``config["quant_cfg"]["*lm_head*"] = ...``) becomes ``list.append()``.
Because later entries override earlier ones, appending achieves the same override effect.
- ``nn.*``-scoped dict keys become entries with a ``parent_class`` field.
----------
Reference
=========
- :class:`QuantizerCfgEntry <modelopt.torch.quantization.config.QuantizerCfgEntry>`
- :class:`QuantizerAttributeConfig <modelopt.torch.quantization.config.QuantizerAttributeConfig>`
- :class:`QuantizeConfig <modelopt.torch.quantization.config.QuantizeConfig>`
- :func:`set_quantizer_by_cfg <modelopt.torch.quantization.conversion.set_quantizer_by_cfg>`
+55 -16
View File
@@ -306,41 +306,80 @@ def ptq(
dist.barrier()
## quant config
mtq_cfg = getattr(mtq, quant_cfg)
import copy
mtq_cfg = copy.deepcopy(getattr(mtq, quant_cfg))
# disable head that corresponds to lm_head (for the huggingface checkpoint)
mtq_cfg["quant_cfg"]["*head*"] = {"enable": False}
mtq_cfg["quant_cfg"].append({"quantizer_name": "*head*", "enable": False})
allowed_mla_quant = [None, "per_tensor_fp8", "nvfp4"]
assert mla_quant in allowed_mla_quant, f"mla_quant must be {allowed_mla_quant}"
if not mla_quant:
mtq_cfg["quant_cfg"]["*attn*"] = {"enable": False}
mtq_cfg["quant_cfg"].append({"quantizer_name": "*attn*", "enable": False})
elif mla_quant == "per_tensor_fp8":
mtq_cfg["quant_cfg"]["*attn*weight_quantizer"] = {"num_bits": (4, 3), "axis": None}
mtq_cfg["quant_cfg"]["*attn*input_quantizer"] = {"num_bits": (4, 3), "axis": None}
mtq_cfg["quant_cfg"].extend(
[
{
"quantizer_name": "*attn*weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
},
{
"quantizer_name": "*attn*input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
},
]
)
elif mla_quant == "nvfp4": # for DeepSeek-R1-0528-NVFP4-Turbo
mla_linear_layers = ["*wq_a*", "*wq_b*", "*wkv_a*", "*wkv_b*", "*wo*"]
mla_nvfp4_linear_layers = ["*wq_a*", "*wkv_a*", "*wq_b*", "*wo*"]
for layer in mla_linear_layers:
if layer in mla_nvfp4_linear_layers:
# wq_a, wkv_a, wq_b, wo use NVFP4 quantization
mtq_cfg["quant_cfg"][layer + "_quantizer"] = {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"enable": True,
}
mtq_cfg["quant_cfg"].append(
{
"quantizer_name": layer + "_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
}
)
else:
mtq_cfg["quant_cfg"][layer + "_quantizer"] = {"enable": False}
mtq_cfg["quant_cfg"].append(
{"quantizer_name": layer + "_quantizer", "enable": False}
)
# Disable BMM quantizers
mtq_cfg["quant_cfg"]["*attn.kv_bmm_quantizer*"] = {"enable": False}
mtq_cfg["quant_cfg"]["*attn.pe_bmm_quantizer*"] = {"enable": False}
mtq_cfg["quant_cfg"].extend(
[
{"quantizer_name": "*attn.kv_bmm_quantizer*", "enable": False},
{"quantizer_name": "*attn.pe_bmm_quantizer*", "enable": False},
]
)
if not args.disable_wo_quant and "FP4" in quant_cfg:
mtq_cfg["quant_cfg"]["*wo*weight_quantizer"] = mtq_cfg["quant_cfg"]["*input_quantizer"]
mtq_cfg["quant_cfg"]["*wo*input_quantizer"] = mtq_cfg["quant_cfg"]["*weight_quantizer"]
# Find the default input/weight quantizer cfgs to swap for wo layers.
# cfg may be a list (SequentialQuantizer); use the first element in that case.
input_cfg = mtq.find_quant_cfg_entry_by_path(mtq_cfg["quant_cfg"], "*input_quantizer")[
"cfg"
]
weight_cfg = mtq.find_quant_cfg_entry_by_path(mtq_cfg["quant_cfg"], "*weight_quantizer")[
"cfg"
]
if isinstance(input_cfg, list):
input_cfg = input_cfg[0]
if isinstance(weight_cfg, list):
weight_cfg = weight_cfg[0]
mtq_cfg["quant_cfg"].extend(
[
{"quantizer_name": "*wo*weight_quantizer", "cfg": input_cfg},
{"quantizer_name": "*wo*input_quantizer", "cfg": weight_cfg},
]
)
## ptq
transformer = mtq.quantize(transformer, mtq_cfg, calibrate_loop)
+73 -70
View File
@@ -17,82 +17,79 @@ import torch.nn as nn
from calib.plugin_calib import PercentileCalibrator
FP8_DEFAULT_CONFIG = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*input_quantizer": {"num_bits": (4, 3), "axis": None},
"*output_quantizer": {"enable": False},
"*softmax_quantizer": {
"num_bits": (4, 3),
"axis": None,
},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*output_quantizer", "enable": False},
{"quantizer_name": "*softmax_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
],
"algorithm": "max",
}
INT8_DEFAULT_CONFIG = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
"*output_quantizer": {"enable": False},
"default": {"enable": False},
},
"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}},
{"quantizer_name": "*output_quantizer", "enable": False},
],
"algorithm": "max",
}
NVFP4_DEFAULT_CONFIG = {
"quant_cfg": {
"*weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*output_quantizer": {"enable": False},
"*softmax_quantizer": {
"num_bits": (4, 3),
"axis": None,
},
"default": {"enable": False},
},
{"quantizer_name": "*output_quantizer", "enable": False},
{"quantizer_name": "*softmax_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
],
"algorithm": "max",
}
NVFP4_FP8_MHA_CONFIG = {
"quant_cfg": {
"**weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "**weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"**input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "**input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*output_quantizer": {"enable": False},
"*[qkv]_bmm_quantizer": {
"num_bits": (4, 3),
"axis": None,
},
"*softmax_quantizer": {
"num_bits": (4, 3),
"axis": None,
},
"*bmm2_output_quantizer": {
"num_bits": (4, 3),
"axis": None,
},
"default": {"enable": False},
},
{"quantizer_name": "*output_quantizer", "enable": False},
{"quantizer_name": "*[qkv]_bmm_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*softmax_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*bmm2_output_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
],
"algorithm": {"method": "svdquant", "lowrank": 32},
}
@@ -106,8 +103,9 @@ def set_quant_config_attr(quant_config, trt_high_precision_dtype, quant_algo, **
algo_cfg["lowrank"] = kwargs["lowrank"]
quant_config["algorithm"] = algo_cfg
for p in quant_config["quant_cfg"].values():
if "num_bits" in p and "trt_high_precision_dtype" not in p:
for entry in quant_config["quant_cfg"]:
p = entry.get("cfg", {})
if isinstance(p, dict) and "num_bits" in p and "trt_high_precision_dtype" not in p:
p["trt_high_precision_dtype"] = trt_high_precision_dtype
@@ -127,18 +125,23 @@ def reset_set_int8_config(quant_config, percentile, n_steps, collect_method, bac
for name, module in backbone.named_modules():
if isinstance(module, nn.Conv2d):
aq_name = f"*{name}*input_quantizer*"
quant_config["quant_cfg"][aq_name] = {
"num_bits": 8,
"axis": None,
"calibrator": (
PercentileCalibrator,
(),
{
quant_config["quant_cfg"].append(
{
"quantizer_name": aq_name,
"cfg": {
"num_bits": 8,
"axis": None,
"percentile": percentile,
"total_step": n_steps,
"collect_method": collect_method,
"calibrator": (
PercentileCalibrator,
(),
{
"num_bits": 8,
"axis": None,
"percentile": percentile,
"total_step": n_steps,
"collect_method": collect_method,
},
),
},
),
}
}
)
+6 -1
View File
@@ -137,7 +137,12 @@ class Quantizer:
else:
raise NotImplementedError(f"Unknown format {self.config.format}")
if self.config.quantize_mha:
quant_config["quant_cfg"]["*[qkv]_bmm_quantizer"] = {"num_bits": (4, 3), "axis": None} # type: ignore[index]
quant_config["quant_cfg"].append(
{
"quantizer_name": "*[qkv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
}
)
set_quant_config_attr(
quant_config,
self.model_config.trt_high_precision_dtype.value,
+12 -2
View File
@@ -100,11 +100,21 @@ def auto_quantize(
if enable_kv_cache_quantization:
mtq.set_quantizer_by_cfg(
model,
quant_cfg={"*output_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True}},
quant_cfg=[
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
}
],
)
# Lets calibrate only the output quantizer this time. Let's disable all other quantizers.
with mtq.set_quantizer_by_cfg_context(
model, {"*": {"enable": False}, "*output_quantizer": {"enable": True}}
model,
[
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*output_quantizer", "enable": True},
],
):
mtq.calibrate(model, algorithm="max", forward_loop=calibrate_loop)
return model
+13 -5
View File
@@ -33,12 +33,20 @@ MAX_OUTPUT_LEN = 512
# Modify your custom config for debugging or research purposes.
CUSTOM_CONFIG = {
"MY_QUANT_CONFIG": {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 4, "block_sizes": {-1: 128}, "enable": True},
"*input_quantizer": {"num_bits": 8, "type": "dynamic", "block_sizes": {-1: None}},
"quant_cfg": [
*mtq.config._base_disable_all,
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 4, "block_sizes": {-1: 128}},
"enable": True,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "type": "dynamic", "block_sizes": {-1: None}},
},
# Disable sensitive layers such as `lm_head`, gate layers in MoE etc.
**mtq.config._default_disabled_quantizer_cfg,
},
*mtq.config._default_disabled_quantizer_cfg,
],
"algorithm": "max",
},
}
+10 -5
View File
@@ -205,7 +205,12 @@ def build_quant_cfg(
) -> dict[str, Any]:
quant_cfg = copy.deepcopy(quant_cfg)
if "awq" in str(quant_cfg.get("algorithm")):
weight_quantizer = quant_cfg["quant_cfg"]["*weight_quantizer"]
from modelopt.torch.quantization.config import find_quant_cfg_entry_by_path
weight_quantizer_entry = find_quant_cfg_entry_by_path(
quant_cfg["quant_cfg"], "*weight_quantizer"
)
weight_quantizer = weight_quantizer_entry.get("cfg") or {}
if isinstance(weight_quantizer, list):
weight_quantizer = weight_quantizer[0]
# If awq_block_size argument is provided, update weight_quantizer
@@ -236,10 +241,10 @@ def build_quant_cfg(
if model_type == "phi4mm":
# Only quantize the language model
quant_cfg["quant_cfg"]["*speech*"] = {"enable": False}
quant_cfg["quant_cfg"]["*audio*"] = {"enable": False}
quant_cfg["quant_cfg"]["*image*"] = {"enable": False}
quant_cfg["quant_cfg"]["*vision*"] = {"enable": False}
quant_cfg["quant_cfg"].append({"quantizer_name": "*speech*", "enable": False})
quant_cfg["quant_cfg"].append({"quantizer_name": "*audio*", "enable": False})
quant_cfg["quant_cfg"].append({"quantizer_name": "*image*", "enable": False})
quant_cfg["quant_cfg"].append({"quantizer_name": "*vision*", "enable": False})
return quant_cfg
+30 -18
View File
@@ -78,16 +78,17 @@ from modelopt.torch.utils.vlm_dataset_utils import get_vlm_dataset_dataloader
RAND_SEED = 1234
def _set_kv_cache_constant_amax(quant_cfg: dict) -> None:
def _set_kv_cache_constant_amax(quant_cfg: list) -> None:
"""Set use_constant_amax on KV cache quantizers.
Creates a new dict for the KV bmm quantizer config to avoid mutating shared references.
"""
if "*[kv]_bmm_quantizer" in quant_cfg:
quant_cfg["*[kv]_bmm_quantizer"] = {
**quant_cfg["*[kv]_bmm_quantizer"],
"use_constant_amax": True,
}
for i, entry in enumerate(quant_cfg):
if entry.get("quantizer_name") != "*[kv]_bmm_quantizer":
continue
assert isinstance(entry.get("cfg", {}), dict)
quant_cfg[i] = {**entry, "cfg": {**entry.get("cfg", {}), "use_constant_amax": True}}
break
QUANT_CFG_CHOICES: dict[str, dict[str, Any]] = {
@@ -145,7 +146,7 @@ def extract_and_prepare_language_model_from_vl(full_model):
# Apply disabled quant to all modules that are not part of language_model
# This excludes them during HF export
disabled_quant_cfg = {
"quant_cfg": {"default": {"enable": False}},
"quant_cfg": [{"quantizer_name": "*", "enable": False}],
"algorithm": "max",
}
@@ -319,7 +320,11 @@ def auto_quantize(
),
verbose=True,
# Disable all default disabled layers such as lm_head, mlp.gate, router etc.
disabled_layers=list(_default_disabled_quantizer_cfg.keys()),
disabled_layers=[
entry["quantizer_name"]
for entry in _default_disabled_quantizer_cfg
if "parent_class" not in entry
],
method=auto_quantize_method,
checkpoint=auto_quantize_checkpoint,
)
@@ -332,7 +337,9 @@ def auto_quantize(
kv_cache_quant_cfg = copy.deepcopy(
getattr(mtq, KV_QUANT_CFG_CHOICES[args.kv_cache_qformat])["quant_cfg"]
)
kv_cache_quant_cfg.pop("default", None) # keep other quantizers from auto_quantize
kv_cache_quant_cfg = [
e for e in kv_cache_quant_cfg if e["quantizer_name"] != "*"
] # keep other quantizers from auto_quantize
if args.kv_cache_qformat in _KV_CAST_FORMATS:
_set_kv_cache_constant_amax(kv_cache_quant_cfg)
@@ -341,7 +348,8 @@ def auto_quantize(
if args.kv_cache_qformat not in _KV_CAST_FORMATS:
# Calibrate only the KV cache quantizers; disable all others.
with mtq.set_quantizer_by_cfg_context(
language_model, {"*": {"enable": False}, **kv_cache_quant_cfg}
language_model,
[{"quantizer_name": "*", "enable": False}, *kv_cache_quant_cfg],
):
mtq.calibrate(language_model, algorithm="max", forward_loop=calibrate_loop)
return language_model
@@ -546,13 +554,17 @@ def mono_quantize(
# For Nemotron VL models, disable quantization of vision components
if is_nemotron_vl_model:
print("Disabling quantization for vision components in Nemotron VL model")
quant_cfg["quant_cfg"]["*vision*"] = {"enable": False}
quant_cfg["quant_cfg"]["*image*"] = {"enable": False}
quant_cfg["quant_cfg"].append({"quantizer_name": "*vision*", "enable": False})
quant_cfg["quant_cfg"].append({"quantizer_name": "*image*", "enable": False})
# Also disable radio model components specifically (for Nemotron-Parse)
quant_cfg["quant_cfg"]["*radio*"] = {"enable": False}
quant_cfg["quant_cfg"]["*visual*"] = {"enable": False}
quant_cfg["quant_cfg"]["*encoder*"] = {"enable": False} # Disable encoder
quant_cfg["quant_cfg"]["*model_encoder*"] = {"enable": False} # Nemotron-Parse specific
quant_cfg["quant_cfg"].append({"quantizer_name": "*radio*", "enable": False})
quant_cfg["quant_cfg"].append({"quantizer_name": "*visual*", "enable": False})
quant_cfg["quant_cfg"].append(
{"quantizer_name": "*encoder*", "enable": False}
) # Disable encoder
quant_cfg["quant_cfg"].append(
{"quantizer_name": "*model_encoder*", "enable": False}
) # Nemotron-Parse specific
print("Quantization will only be applied to the decoder (text generation) component")
if not model_is_already_quantized or calibration_only:
@@ -943,7 +955,7 @@ def quantize_main(
assert isinstance(recipe, ModelOptPTQRecipe), (
f"Expected PTQ recipe, but got {type(recipe).__name__} from {args.recipe}"
)
quant_cfg = recipe.ptq_cfg
quant_cfg = recipe.quantize
else:
assert len(args.qformat.split(",")) == 1, (
@@ -980,7 +992,7 @@ def quantize_main(
quant_cfg = copy.deepcopy(quant_cfg)
for prefix in mtp_layer_prefixes:
pattern = f"*{prefix}*"
quant_cfg["quant_cfg"][pattern] = {"enable": False}
quant_cfg["quant_cfg"].append({"quantizer_name": pattern, "enable": False})
print(f"Excluding MTP layer from quantization: {pattern}")
# Use constant amax for KV quantizers when a cast format is selected.
@@ -189,17 +189,7 @@
"id": "a3ce3b47-48ac-4a27-a5ed-351a10c104a9",
"metadata": {},
"outputs": [],
"source": [
"# Get default AWQ config and optionally adjust block size\n",
"quant_cfg = mtq.INT4_AWQ_CFG\n",
"weight_quantizer = quant_cfg[\"quant_cfg\"][\"*weight_quantizer\"]\n",
"if isinstance(weight_quantizer, list):\n",
" weight_quantizer = weight_quantizer[0]\n",
"weight_quantizer[\"block_sizes\"][-1] = 128 # Optional: override block size\n",
"\n",
"# Apply AWQ quantization\n",
"model = mtq.quantize(model, quant_cfg, forward_loop=forward_loop)"
]
"source": "import copy\n\nfrom modelopt.torch.quantization.config import find_quant_cfg_entry_by_path\n\n# Get default AWQ config and optionally adjust block size\nquant_cfg = copy.deepcopy(mtq.INT4_AWQ_CFG)\nweight_quantizer_entry = find_quant_cfg_entry_by_path(quant_cfg[\"quant_cfg\"], \"*weight_quantizer\")\ncfg = weight_quantizer_entry.get(\"cfg\")\nassert cfg is not None, \"Expected cfg to be set for *weight_quantizer entry\"\ncfg = copy.deepcopy(cfg)\nif isinstance(cfg, list):\n cfg = cfg[0]\ncfg[\"block_sizes\"][-1] = 128 # Optional: override block size\nweight_quantizer_entry[\"cfg\"] = cfg\n\n# Apply AWQ quantization\nmodel = mtq.quantize(model, quant_cfg, forward_loop=forward_loop)"
},
{
"cell_type": "markdown",
@@ -308,4 +298,4 @@
},
"nbformat": 4,
"nbformat_minor": 5
}
}
@@ -288,7 +288,9 @@
" mtq.set_quantizer_by_cfg(model, quant_cfg=kv_cfg)\n",
"\n",
" # Calibrate **only** those quantizers\n",
" with mtq.set_quantizer_by_cfg_context(model, {\"*\": {\"enable\": False}, **kv_cfg}):\n",
" with mtq.set_quantizer_by_cfg_context(\n",
" model, [{\"quantizer_name\": \"*\", \"enable\": False}, *kv_cfg]\n",
" ):\n",
" mtq.calibrate(model, algorithm=\"max\", forward_loop=forward_loop)\n",
"else:\n",
" print(\"KV cache left unquantized.\")"
@@ -427,4 +429,4 @@
},
"nbformat": 4,
"nbformat_minor": 5
}
}
+14 -6
View File
@@ -54,12 +54,20 @@ mto.enable_huggingface_checkpointing()
CUSTOM_QUANT_CFG = {
"INT4_WEIGHT_INT8_ACTIVATIONS": {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 4, "block_sizes": {-1: 128}, "enable": True},
"*input_quantizer": {"num_bits": 8, "axis": None, "enable": True},
"*lm_head*": {"enable": False},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 4, "block_sizes": {-1: 128}},
"enable": True,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None},
"enable": True,
},
{"quantizer_name": "*lm_head*", "enable": False},
],
"algorithm": "max",
}
}
+25 -6
View File
@@ -102,7 +102,7 @@ def calibrate_fun(calib_dataloader: DataLoader, self: Any) -> Callable[[Any], No
return calibrate_loop
def update_kv_cfg_for_mla(model: torch.nn.Module, kv_quant_cfg: dict[str, Any]) -> dict[str, Any]:
def update_kv_cfg_for_mla(model: torch.nn.Module, kv_quant_cfg: list) -> list:
"""Update KV cache quantization config for MLA models.
MLA uses `kv_c_bmm_quantizer` (compressed KV) instead of separate
@@ -117,18 +117,37 @@ def update_kv_cfg_for_mla(model: torch.nn.Module, kv_quant_cfg: dict[str, Any])
if not any(isinstance(m, MLAAttention) for m in model.modules()):
return kv_quant_cfg
if kv_config := kv_quant_cfg.get("*[kv]_bmm_quantizer"):
kv_quant_cfg["*kv_c_bmm_quantizer"] = kv_config
kv_quant_cfg["*k_pe_bmm_quantizer"] = kv_config
kv_entry = next(
(
e
for e in kv_quant_cfg
if isinstance(e, dict) and e.get("quantizer_name") == "*[kv]_bmm_quantizer"
),
None,
)
if kv_entry is not None:
kv_config = kv_entry.get("cfg", {})
kv_quant_cfg.append(
{"quantizer_name": "*kv_c_bmm_quantizer", "cfg": kv_config, "enable": True}
)
kv_quant_cfg.append(
{"quantizer_name": "*k_pe_bmm_quantizer", "cfg": kv_config, "enable": True}
)
print("MLA detected: added *kv_c_bmm_quantizer and k_pe_bmm_quantizer config")
return kv_quant_cfg
def get_quant_config(quant_config: dict[str, Any], model: Any) -> dict[str, Any]:
quant_cfg = getattr(mtq, quant_config["quant_cfg"]) if quant_config["quant_cfg"] else {}
import copy
quant_cfg = (
copy.deepcopy(getattr(mtq, quant_config["quant_cfg"])) if quant_config["quant_cfg"] else {}
)
quant_kv_cfg = (
getattr(mtq, quant_config["kv_quant_cfg"]) if quant_config["kv_quant_cfg"] else {}
copy.deepcopy(getattr(mtq, quant_config["kv_quant_cfg"]))
if quant_config["kv_quant_cfg"]
else {}
)
# Check if model has MLA and update KV config accordingly
@@ -257,26 +257,20 @@ def build_quant_config(
if exclude_blocks is None:
exclude_blocks = [0, 1, 46, 47]
quant_cfg = {
"*weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"enable": True,
},
"*input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"enable": True,
},
_nvfp4_cfg = {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
}
for pattern in SENSITIVE_LAYER_PATTERNS:
quant_cfg[pattern] = {"enable": False}
for block_idx in exclude_blocks:
quant_cfg[f"*transformer_blocks.{block_idx}.*"] = {"enable": False}
quant_cfg = [
{"quantizer_name": "*weight_quantizer", "cfg": _nvfp4_cfg, "enable": True},
{"quantizer_name": "*input_quantizer", "cfg": _nvfp4_cfg, "enable": True},
*[{"quantizer_name": pattern, "enable": False} for pattern in SENSITIVE_LAYER_PATTERNS],
*[
{"quantizer_name": f"*transformer_blocks.{i}.*", "enable": False}
for i in exclude_blocks
],
]
return {
"quant_cfg": quant_cfg,
@@ -15,6 +15,7 @@
"""Quantization utilities for LLM models."""
import copy
import time
import modelopt.torch.quantization as mtq
@@ -57,35 +58,58 @@ def _quantize_model(model, quant_config, calib_dataloader=None):
def get_quant_config(precision, lm_head_precision="fp16"):
"""Get the quantization configuration."""
if precision == "fp8":
quant_cfg = mtq.FP8_DEFAULT_CFG
quant_cfg = copy.deepcopy(mtq.FP8_DEFAULT_CFG)
elif precision == "nvfp4":
quant_cfg = mtq.NVFP4_DEFAULT_CFG
quant_cfg = copy.deepcopy(mtq.NVFP4_DEFAULT_CFG)
elif precision == "int4_awq":
quant_cfg = mtq.INT4_AWQ_CFG
quant_cfg = copy.deepcopy(mtq.INT4_AWQ_CFG) # type: ignore[arg-type]
else:
raise ValueError(f"Unsupported precision: {precision}")
config_dict = quant_cfg["quant_cfg"] # type: dict
quant_cfg_list: list = [
e for e in quant_cfg["quant_cfg"] if isinstance(e, dict) and "quantizer_name" in e
]
if lm_head_precision == "fp8":
config_dict["*lm_head.input_quantizer"] = {"num_bits": (4, 3), "axis": None}
config_dict["*lm_head.weight_quantizer"] = {"num_bits": (4, 3), "axis": None}
quant_cfg_list.append(
{
"quantizer_name": "*lm_head.input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
}
)
quant_cfg_list.append(
{
"quantizer_name": "*lm_head.weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
}
)
elif lm_head_precision == "nvfp4":
config_dict["*lm_head.input_quantizer"] = {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"enable": True,
}
config_dict["*lm_head.weight_quantizer"] = {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"enable": True,
}
quant_cfg_list.append(
{
"quantizer_name": "*lm_head.input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
}
)
quant_cfg_list.append(
{
"quantizer_name": "*lm_head.weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
}
)
quant_cfg["quant_cfg"] = quant_cfg_list
return quant_cfg
+1 -1
View File
@@ -66,7 +66,7 @@ class ModelOptRecipeBase(ModeloptBaseConfig):
class ModelOptPTQRecipe(ModelOptRecipeBase):
"""Our config class for PTQ recipes."""
ptq_cfg: dict[str, Any] = ModeloptField(
quantize: dict[str, Any] = ModeloptField(
default={},
title="PTQ config",
description="PTQ config containing quant_cfg and algorithm.",
+12 -12
View File
@@ -54,9 +54,9 @@ def load_recipe(recipe_path: str | Path | Traversable) -> ModelOptRecipeBase:
``recipe_path`` can be:
* A ``.yml`` / ``.yaml`` file with ``metadata`` and ``ptq_cfg`` sections.
* A ``.yml`` / ``.yaml`` file with ``metadata`` and ``quantize`` sections.
The suffix may be omitted and will be probed automatically.
* A directory containing ``recipe.yml`` (metadata) and ``ptq_cfg.yml``.
* A directory containing ``recipe.yml`` (metadata) and ``quantize.yml``.
The path may be relative to the built-in recipes library or an absolute /
relative filesystem path.
@@ -94,18 +94,18 @@ def _load_recipe_from_file(recipe_file: Path | Traversable) -> ModelOptRecipeBas
raise ValueError(f"Recipe file {recipe_file} must contain a 'metadata.recipe_type' field.")
if recipe_type == RecipeType.PTQ:
if "ptq_cfg" not in data:
raise ValueError(f"PTQ recipe file {recipe_file} must contain 'ptq_cfg'.")
if "quantize" not in data:
raise ValueError(f"PTQ recipe file {recipe_file} must contain 'quantize'.")
return ModelOptPTQRecipe(
recipe_type=RecipeType.PTQ,
description=metadata.get("description", "PTQ recipe."),
ptq_cfg=data["ptq_cfg"],
quantize=data["quantize"],
)
raise ValueError(f"Unsupported recipe type: {recipe_type!r}")
def _load_recipe_from_dir(recipe_dir: Path | Traversable) -> ModelOptRecipeBase:
"""Load a recipe from a directory containing ``recipe.yml`` and ``ptq_cfg.yml``."""
"""Load a recipe from a directory containing ``recipe.yml`` and ``quantize.yml``."""
recipe_file = None
for name in ("recipe.yml", "recipe.yaml"):
candidate = recipe_dir.joinpath(name)
@@ -123,19 +123,19 @@ def _load_recipe_from_dir(recipe_dir: Path | Traversable) -> ModelOptRecipeBase:
raise ValueError(f"Recipe file {recipe_file} must contain a 'metadata.recipe_type' field.")
if recipe_type == RecipeType.PTQ:
ptq_cfg_file = None
for name in ("ptq_cfg.yml", "ptq_cfg.yaml"):
quantize_file = None
for name in ("quantize.yml", "quantize.yaml"):
candidate = recipe_dir.joinpath(name)
if candidate.is_file():
ptq_cfg_file = candidate
quantize_file = candidate
break
if ptq_cfg_file is None:
if quantize_file is None:
raise ValueError(
f"Cannot find ptq_cfg in {recipe_dir}. Looked for: ptq_cfg.yml, ptq_cfg.yaml"
f"Cannot find quantize in {recipe_dir}. Looked for: quantize.yml, quantize.yaml"
)
return ModelOptPTQRecipe(
recipe_type=RecipeType.PTQ,
description=metadata.get("description", "PTQ recipe."),
ptq_cfg=load_config(ptq_cfg_file),
quantize=load_config(quantize_file),
)
raise ValueError(f"Unsupported recipe type: {recipe_type!r}")
+4 -1
View File
@@ -218,7 +218,10 @@ def _collect_shared_input_modules(
# Run dummy forward pass to collect modules sharing same input
try:
with torch.no_grad(), set_quantizer_by_cfg_context(model, {"*": {"enable": False}}):
with (
torch.no_grad(),
set_quantizer_by_cfg_context(model, [{"quantizer_name": "*", "enable": False}]),
):
dummy_forward_fn()
finally:
# Always remove hooks
+77 -23
View File
@@ -62,9 +62,22 @@ def estimate_quant_compression(quant_cfg: QuantizeConfig) -> float:
def estimate_quant_compression_for_quantizer(quantizer_attr_cfg):
if isinstance(quantizer_attr_cfg, list):
if not quantizer_attr_cfg:
return 1.0
return min(estimate_quant_compression_for_quantizer(q) for q in quantizer_attr_cfg)
if isinstance(quantizer_attr_cfg, dict):
return estimate_quant_compression_for_quantizer(list(quantizer_attr_cfg.values()))
# Handle raw quantizer cfg dicts (e.g. {"num_bits": (4, 3), "axis": None})
if not quantizer_attr_cfg.get("enable", True):
return 1.0
num_bits = quantizer_attr_cfg.get("num_bits")
if num_bits is None:
return 1.0
if isinstance(num_bits, tuple):
return (sum(num_bits) + 1) / 16
elif isinstance(num_bits, int):
return num_bits / 16
else:
raise ValueError(f"Unknown quantization config {num_bits}")
if isinstance(quantizer_attr_cfg, QuantizerAttributeConfig):
if not quantizer_attr_cfg.enable:
@@ -80,7 +93,14 @@ def estimate_quant_compression(quant_cfg: QuantizeConfig) -> float:
raise ValueError(f"Unknown type {type(quantizer_attr_cfg)}, {quantizer_attr_cfg}")
return estimate_quant_compression_for_quantizer(list(quant_cfg.quant_cfg.values()))
cfgs = []
for e in quant_cfg.quant_cfg:
if e.get("enable", True) is False:
continue
c = e.get("cfg")
if c is not None:
cfgs.append(c)
return estimate_quant_compression_for_quantizer(cfgs) if cfgs else 1.0
class QuantRecipe(CustomHPType):
@@ -97,7 +117,7 @@ class QuantRecipe(CustomHPType):
name = self.get_auto_name_for_config(quant_cfg) or name
if quant_cfg is None:
quant_cfg = {"quant_cfg": {"*": {"enable": False}}}
quant_cfg = {"quant_cfg": [{"quantizer_name": "*", "enable": False}]}
elif isinstance(quant_cfg, str):
assert hasattr(mtq_config, quant_cfg), f"Unknown quantization format {quant_cfg}"
quant_cfg = getattr(mtq_config, quant_cfg)
@@ -109,9 +129,7 @@ class QuantRecipe(CustomHPType):
# Disable KV Cache quantization
# Currently KV Cache quantization is enabled for some quantization formats and disabled for others
# This breaks the monotonicity of the quantization formats in terms of weight compression Vs accuracy
self.config.quant_cfg["*output_quantizer"] = mtq_config.QuantizerAttributeConfig(
enable=False
)
self.config.quant_cfg.append({"quantizer_name": "*output_quantizer", "enable": False})
self.compression = estimate_quant_compression(self.config)
@@ -1300,21 +1318,9 @@ def get_auto_quantize_config(search_state, constraints=None, verbose=False):
else:
best_recipe = search_state["best"]["recipe"]
quant_cfg: dict[str, Any] = {"*": {"enable": False}}
for hparam_name, recipe in best_recipe.items():
if recipe == QuantRecipe(quant_cfg=None):
continue
module_names = search_state["candidate_stats"][hparam_name]["module_names"]
for module_name in module_names:
for quantizer_attr in ("input_quantizer", "weight_quantizer"):
matched_cfg = _match_quantizer_cfg(recipe.config.quant_cfg, quantizer_attr)
if matched_cfg is not None:
quant_cfg[f"{module_name}.{quantizer_attr}"] = matched_cfg
def _cfg_to_dict(v):
if isinstance(v, mtq_config.QuantizerAttributeConfig):
return {
"enable": v.enable,
"num_bits": v.num_bits,
**v.model_dump(exclude_defaults=True),
}
@@ -1322,7 +1328,45 @@ def get_auto_quantize_config(search_state, constraints=None, verbose=False):
return [_cfg_to_dict(c) for c in v]
return v
quant_cfg = {k: _cfg_to_dict(v) for k, v in quant_cfg.items()}
quant_cfg: list[dict] = [{"quantizer_name": "*", "enable": False}]
_per_module_attrs = ("input_quantizer", "weight_quantizer", "output_quantizer")
# Track global (non per-module) recipe entries. Last recipe wins for each pattern.
global_entries: dict[str, dict] = {}
for hparam_name, recipe in best_recipe.items():
if recipe == QuantRecipe(quant_cfg=None):
continue
module_names = search_state["candidate_stats"][hparam_name]["module_names"]
for module_name in module_names:
for quantizer_attr in _per_module_attrs:
matched_cfg, matched_enable = _match_quantizer_cfg(
recipe.config.quant_cfg, quantizer_attr
)
if matched_enable is not None:
entry: dict[str, Any] = {
"quantizer_name": f"{module_name}.{quantizer_attr}",
"enable": matched_enable,
}
if matched_cfg is not None:
entry["cfg"] = _cfg_to_dict(matched_cfg)
quant_cfg.append(entry)
# Collect non-per-module entries (e.g. *[kv]_bmm_quantizer) from winning recipes.
for recipe_entry in recipe.config.quant_cfg:
pattern = recipe_entry["quantizer_name"]
if pattern == "*" or any(
fnmatch.fnmatch(attr, pattern) or pattern.endswith(attr)
for attr in _per_module_attrs
):
continue
cfg = recipe_entry.get("cfg")
enable = recipe_entry.get("enable", True)
ge: dict[str, Any] = {"quantizer_name": pattern, "enable": enable}
if cfg is not None:
ge["cfg"] = _cfg_to_dict(cfg)
global_entries[pattern] = ge
quant_cfg.extend(global_entries.values())
warnings.warn(
"get_auto_quantize_config: returned config uses algorithm='max'. "
"Per-recipe calibration algorithms (e.g. smoothquant, awq) are not preserved. "
@@ -1362,9 +1406,19 @@ def _resolve_best_recipe(search_state, constraints, verbose=False):
def _match_quantizer_cfg(quant_cfg, quantizer_attr):
# Last-match-wins to mirror set_quantizer_by_cfg behavior
# Last-match-wins to mirror set_quantizer_by_cfg behavior.
# Patterns may be path-scoped (e.g. "*mlp*weight_quantizer") while quantizer_attr
# is a bare name like "weight_quantizer". We match if the bare name matches directly
# OR if the pattern ends with the bare quantizer_attr (path-scoped match).
matched = None
for pattern, cfg in quant_cfg.items():
if fnmatch.fnmatch(quantizer_attr, pattern):
matched_enable = None
for entry in quant_cfg:
pattern = entry["quantizer_name"]
cfg = entry.get("cfg")
enable = entry.get("enable", True)
# Direct match: the bare quantizer_attr matches the whole pattern (e.g. "*weight_quantizer")
if fnmatch.fnmatch(quantizer_attr, pattern) or pattern.endswith(quantizer_attr):
matched = cfg
return matched
matched_enable = enable
return matched, matched_enable
@@ -15,13 +15,11 @@
"""This module provides a GEMM function for fp8 per tensor quantization."""
from typing import Any
import torch
from torch.autograd import Function
from modelopt.torch.quantization.backends.gemm_registry import gemm_registry
from modelopt.torch.quantization.config import FP8_DEFAULT_CFG
from modelopt.torch.quantization.config import FP8_DEFAULT_CFG, find_quant_cfg_entry_by_path
from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear
from modelopt.torch.quantization.qtensor import FP8QTensor, QTensorWrapper
from modelopt.torch.quantization.utils import reduce_amax
@@ -99,9 +97,16 @@ def fp8_per_tensor_gemm(quant_module, input, bias=None):
def _fp8_availability_check(module, input, args, kwargs):
"""Comprehensive check for FP8 GEMM availability."""
# Quantizer configs
quant_cfg: dict[str, Any] = FP8_DEFAULT_CFG["quant_cfg"]
input_cfg = quant_cfg["*input_quantizer"]
weight_cfg = quant_cfg["*weight_quantizer"]
quant_cfg_list: list = FP8_DEFAULT_CFG["quant_cfg"]
input_cfg = find_quant_cfg_entry_by_path(quant_cfg_list, "*input_quantizer").get("cfg", {})
weight_cfg = find_quant_cfg_entry_by_path(quant_cfg_list, "*weight_quantizer").get("cfg", {})
# cfg may be a list (SequentialQuantizer); fall back to the first element.
if isinstance(input_cfg, list):
input_cfg = input_cfg[0]
if isinstance(weight_cfg, list):
weight_cfg = weight_cfg[0]
if not isinstance(input_cfg, dict) or not isinstance(weight_cfg, dict):
return False
# Check hardware support
if not torch.cuda.is_available() or not fp8_compatible():
@@ -15,8 +15,6 @@
"""This module provides a GEMM function for nvfp4 quantization."""
from typing import Any
import torch
from torch.autograd import Function
@@ -213,10 +211,21 @@ def _nvfp4_availability_check(module, input, args, kwargs):
if not hasattr(module, "input_quantizer") or not hasattr(module, "weight_quantizer"):
return False
quant_cfg: dict[str, Any] = mtq.NVFP4_DEFAULT_CFG["quant_cfg"]
quant_cfg_list: list = mtq.NVFP4_DEFAULT_CFG["quant_cfg"]
# Quantizer configs
input_cfg = quant_cfg["*input_quantizer"]
weight_cfg = quant_cfg["*weight_quantizer"]
input_cfg = mtq.config.find_quant_cfg_entry_by_path(quant_cfg_list, "*input_quantizer").get(
"cfg", {}
)
weight_cfg = mtq.config.find_quant_cfg_entry_by_path(quant_cfg_list, "*weight_quantizer").get(
"cfg", {}
)
# cfg may be a list (SequentialQuantizer); fall back to the first element.
if isinstance(input_cfg, list):
input_cfg = input_cfg[0]
if isinstance(weight_cfg, list):
weight_cfg = weight_cfg[0]
if not isinstance(input_cfg, dict) or not isinstance(weight_cfg, dict):
return False
# Check input quantizer config
for key, value in input_cfg.items():
+3 -3
View File
@@ -30,7 +30,7 @@ from modelopt.torch.opt.mode import ConvertReturnType, MetadataDict
from .backends.gemm_registry import disable_real_quant_gemm, enable_real_quant_gemm
from .config import CompressCfgType, CompressConfig
from .conversion import _replace_quant_module, set_quantizer_attribute
from .conversion import _replace_quant_module, set_quantizer_attributes_partial
from .nn.modules.quant_linear import RealQuantLinear
from .qtensor import QTensorWrapper, pack_real_quantize_weight
from .utils import is_quantized_linear
@@ -87,7 +87,7 @@ def compress_convert(
compress_cfg = config.compress
if "default" in compress_cfg and isinstance(compress_cfg["default"], bool):
set_quantizer_attribute(
set_quantizer_attributes_partial(
model, "*weight_quantizer*", {"fake_quant": not compress_cfg["default"]}
)
@@ -99,7 +99,7 @@ def compress_convert(
def filter_func(name):
return fnmatch.fnmatch(name, pattern) and "weight_quantizer" in name
set_quantizer_attribute(model, filter_func, {"fake_quant": not to_compress})
set_quantizer_attributes_partial(model, filter_func, {"fake_quant": not to_compress})
else:
raise ValueError(
f"Invalid compression configuration: {to_compress}, expected a boolean as value."
File diff suppressed because it is too large Load Diff
+316 -110
View File
@@ -19,7 +19,7 @@ import fnmatch
import warnings
from collections.abc import Callable
from contextlib import contextmanager
from typing import Any
from typing import Any, cast
import torch.nn as nn
@@ -33,6 +33,7 @@ from .config import (
QuantizeQuantCfgType,
QuantizerAttributeConfig,
_QuantizeExportConfig,
normalize_quant_cfg_list,
)
from .nn import (
NVFP4StaticQuantizer,
@@ -48,6 +49,8 @@ __all__ = [
"register",
"replace_quant_module",
"set_quantizer_attribute",
"set_quantizer_attributes_full",
"set_quantizer_attributes_partial",
"set_quantizer_by_cfg",
"set_quantizer_by_cfg_context",
"unregister",
@@ -60,7 +63,7 @@ def convert_to_quantized_model(model: ModelLikeModule, config: QuantizeConfig) -
model = model.init_modellike() if isinstance(model, ModelLikeModule) else model
replace_quant_module(model, version=ModeloptStateManager(model).state_version)
set_quantizer_by_cfg(model, config.get("quant_cfg", {}))
set_quantizer_by_cfg(model, config.get("quant_cfg", []))
metadata = {}
update_quantize_metadata(model, config, metadata)
@@ -76,7 +79,7 @@ def convert_to_quantized_model_svdquant(
model = model.init_modellike() if isinstance(model, ModelLikeModule) else model
create_and_replace_svdquant_linear_on_the_fly(model)
set_quantizer_by_cfg(model, config.get("quant_cfg", {}))
set_quantizer_by_cfg(model, config.get("quant_cfg", []))
metadata = {}
update_quantize_metadata(model, config, metadata)
@@ -211,127 +214,330 @@ def _replace_quant_module(model: nn.Module, version=None, registry=QuantModuleRe
_replace_quant_module(getattr(model, name), version=version, registry=registry)
def set_quantizer_by_cfg(quant_model: nn.Module, quant_cfg: QuantizeQuantCfgType | dict):
"""Update the quantizer attributes based on the specified `quant_cfg`.
def set_quantizer_by_cfg(quant_model: nn.Module, quant_cfg: QuantizeQuantCfgType):
"""Apply a quantization config list to the quantizers in ``quant_model``.
`quant_cfg` is a dictionary mapping wildcards or filter functions
to its quantizer attributes which are defined in
:class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>`.
The wildcards or filter functions are matched against the quantizer module names.
The specified quantizer attributes of the matched quantizer modules are set accordingly.
The key ``"default"`` is a special key that sets the quantizer attributes of all the quantizers for which
no other wildcard or filter functions match the quantizer module name.
``quant_cfg`` is an **ordered list** of :class:`QuantizerCfgEntry <.config.QuantizerCfgEntry>`
dicts. Each entry has the following fields:
In addition, the dictionary entries could also be pytorch module class names mapping the class specific
quantization configuration. The pytorch modules should have a quantized equivalent.
- ``quantizer_name`` *(required)*: wildcard matched against quantizer module names via
:func:`fnmatch`.
- ``cfg`` *(optional)*: a dict of :class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>`
fields, or a list of such dicts for sequential quantization.
- ``enable`` *(optional)*: ``True`` or ``False`` to toggle matched quantizers on or off.
When omitted but ``cfg`` is present, defaults to ``True``. Every entry must specify at
least one of ``cfg`` or ``enable`` — an entry with only ``quantizer_name`` is invalid.
- ``parent_class`` *(optional)*: restricts matching to quantizers whose immediate parent
module is of this PyTorch class name.
See :meth:`set_quantizer_attribute <modelopt.torch.quantization.conversion.set_quantizer_attribute>`
for more details.
**Ordering and atomicity:** entries are applied in list order; later entries override earlier
ones for any quantizer they match. Each entry with a ``cfg`` is a **complete replacement** —
unspecified attributes revert to their defaults rather than inheriting from a prior entry.
The typical pattern is to deny all first (``{"quantizer_name": "*", "enable": False}``), then
selectively enable and configure target quantizers in subsequent entries.
**``enable`` and ``cfg`` are independent:**
- An entry with ``cfg`` (and optionally ``enable``) fully replaces the matched quantizer's
attributes. If ``enable`` is omitted, the quantizer is implicitly enabled.
- ``{"enable": False}`` without ``cfg`` **only** toggles the matched quantizers off, leaving
all other attributes unchanged.
- ``{"enable": True}`` without ``cfg`` **only** toggles the matched quantizers on, using
whatever attributes they currently have (or their defaults if never configured).
See :ref:`quant-cfg` for the full format reference and common patterns.
"""
quant_cfg = quant_cfg.copy()
if "default" in quant_cfg:
set_quantizer_attribute(quant_model, "*", quant_cfg["default"])
quant_cfg.pop("default")
quant_cfg = normalize_quant_cfg_list(quant_cfg)
for pattern, cfg in quant_cfg.items():
if str(pattern) in QuantModuleRegistry:
parent_class = QuantModuleRegistry[str(pattern)]
assert isinstance(cfg, dict), (
f"Expected a dictionary for quantizer configuration for child tensor quantizers of {parent_class}."
for entry in quant_cfg:
quantizer_name: str = entry["quantizer_name"]
cfg = entry["cfg"] # None, dict, or list — always explicit after normalization
enable: bool = entry["enable"] # always explicit after normalization
parent_class_name = entry.get("parent_class")
if parent_class_name:
try:
parent_class = QuantModuleRegistry[parent_class_name]
except KeyError:
raise ValueError(
f"parent_class {parent_class_name!r} not found in QuantModuleRegistry. "
"Make sure the class has a registered quantized equivalent."
) from None
else:
parent_class = None
if cfg is None:
# No cfg: only toggle the enable state, leave all other attributes unchanged.
set_quantizer_attributes_partial(
quant_model, quantizer_name, {"enable": enable}, parent_class
)
for sub_pattern, sub_cfg in cfg.items():
set_quantizer_attribute(quant_model, sub_pattern, sub_cfg, parent_class)
else:
# Has cfg: apply full replacement with the explicit enable value.
if isinstance(cfg, QuantizerAttributeConfig):
attributes = cfg.model_copy(update={"enable": enable})
elif isinstance(cfg, dict):
attributes = QuantizerAttributeConfig(**cfg, enable=enable)
else:
attributes = [
c.model_copy(update={"enable": enable})
if isinstance(c, QuantizerAttributeConfig)
else QuantizerAttributeConfig(**c, enable=enable)
for c in cfg
]
set_quantizer_attributes_full(quant_model, quantizer_name, attributes, parent_class)
def _match_quantizer(
wildcard_or_filter_func: str | Callable,
name: str,
module: nn.Module,
parent_class: type[nn.Module] | None,
full_model: nn.Module,
):
if not isinstance(module, (TensorQuantizer, SequentialQuantizer)):
return False
if isinstance(wildcard_or_filter_func, str):
if not fnmatch.fnmatch(name, wildcard_or_filter_func):
return False
elif callable(wildcard_or_filter_func):
if not wildcard_or_filter_func(name):
return False
else:
raise NotImplementedError(f"Unsupported type {type(wildcard_or_filter_func)}")
# Get the parent module of this quantizer. When name has no dots (root-level quantizer),
# ".".join([]) == "" and get_submodule("") returns the model itself (PyTorch convention).
return parent_class is None or isinstance(
full_model.get_submodule(".".join(name.split(".")[:-1])), parent_class
)
def set_quantizer_attributes_full(
quant_model: nn.Module,
wildcard_or_filter_func: str | Callable,
attributes: QuantizerAttributeConfig | list[QuantizerAttributeConfig],
parent_class: type[nn.Module] | None = None,
):
"""Set quantizer attributes by wildcard or filter function, fully overwriting existing attributes.
Unlike :func:`set_quantizer_attributes_partial`, this function requires a complete
:class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>` and **replaces** the
matched quantizer's attributes entirely rather than merging with existing ones.
Args:
quant_model: A pytorch model.
wildcard_or_filter_func: A wildcard string or a filter function. The wildcard string is
matched against the quantizer module names. The quantizer modules are instances of
:class:`TensorQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.TensorQuantizer>`.
The filter function takes a quantizer module name as input and returns ``True`` if the
quantizer should be adjusted and ``False`` otherwise.
attributes: A :class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>` (or a
list of them) that **fully replaces** the matched quantizer's current attributes. All
fields of the config are applied — unspecified fields revert to their defaults.
If ``attributes`` is a list, the matched
:class:`TensorQuantizer <nn.modules.tensor_quantizer.TensorQuantizer>`
modules will be replaced with
:class:`SequentialQuantizer <nn.modules.tensor_quantizer.SequentialQuantizer>`
modules having one quantizer per attribute instance in the list.
See
:meth:`set_from_attribute_config() <nn.modules.tensor_quantizer.TensorQuantizer.set_from_attribute_config>`
for details on supported attributes and their types.
parent_class: (Optional) Restrict matching to quantizers whose immediate parent module is
an instance of this class. If ``None``, all quantizers matching
``wildcard_or_filter_func`` are adjusted.
"""
if not isinstance(attributes, (QuantizerAttributeConfig, list)):
raise ValueError(
f"Invalid type for attributes: {type(attributes)}, "
"expected QuantizerAttributeConfig or list of QuantizerAttributeConfig."
)
if isinstance(attributes, list) and not all(
isinstance(attr, QuantizerAttributeConfig) for attr in attributes
):
raise ValueError(
"All elements in attributes list must be of type QuantizerAttributeConfig."
)
for name, module in quant_model.named_modules():
if _match_quantizer(wildcard_or_filter_func, name, module, parent_class, quant_model):
if isinstance(attributes, list):
if not isinstance(module, SequentialQuantizer):
parent_module = quant_model.get_submodule(name.rpartition(".")[0])
module = SequentialQuantizer(
*(TensorQuantizer() for _ in range(len(attributes)))
)
setattr(parent_module, name.split(".")[-1], module)
elif len(attributes) != len(module):
warnings.warn(
f"The number of attributes ({len(attributes)}) does not match the number of "
f"quantizers of {module} leading to partial assignment.",
)
module.set_from_attribute_config(attributes)
else:
if isinstance(module, SequentialQuantizer):
# Downgrade SequentialQuantizer back to TensorQuantizer when the
# new entry provides a single (non-list) config.
parent_module = quant_model.get_submodule(name.rpartition(".")[0])
module = TensorQuantizer()
setattr(parent_module, name.split(".")[-1], module)
cast("TensorQuantizer", module).set_from_attribute_config(attributes)
def set_quantizer_attributes_partial(
quant_model: nn.Module,
wildcard_or_filter_func: str | Callable,
partial_attributes: dict[str, Any] | list[dict[str, Any]],
parent_class: type[nn.Module] | None = None,
):
"""Update a subset of quantizer attributes by wildcard or filter function, merging with existing attributes.
Unlike :func:`set_quantizer_attributes_full`, this function accepts an arbitrary subset of
quantizer attributes as a plain ``dict`` and **merges** them into the matched quantizer's
current attributes, leaving unspecified attributes unchanged.
Args:
quant_model: A pytorch model.
wildcard_or_filter_func: A wildcard string or a filter function. The wildcard string is
matched against the quantizer module names. The quantizer modules are instances of
:class:`TensorQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.TensorQuantizer>`.
The filter function takes a quantizer module name as input and returns ``True`` if the
quantizer should be adjusted and ``False`` otherwise.
partial_attributes: A ``dict`` (or a list of ``dict``) containing only the attributes to
update. Keys must be valid fields of
:class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>`. Only the
specified keys are written; all other attributes on the quantizer remain unchanged.
When a ``dict`` is passed and the matched module is a
:class:`SequentialQuantizer <nn.modules.tensor_quantizer.SequentialQuantizer>`,
the dict is broadcast to every sub-quantizer.
When a ``list`` is passed, the matched module must already be a
:class:`SequentialQuantizer <nn.modules.tensor_quantizer.SequentialQuantizer>` —
unlike :func:`set_quantizer_attributes_full`, this function will **not** replace a
:class:`TensorQuantizer <nn.modules.tensor_quantizer.TensorQuantizer>` with a
``SequentialQuantizer``.
See
:meth:`set_from_attribute_config() <nn.modules.tensor_quantizer.TensorQuantizer.set_from_attribute_config>`
for details on supported attributes and their types.
parent_class: (Optional) Restrict matching to quantizers whose immediate parent module is
an instance of this class. If ``None``, all quantizers matching
``wildcard_or_filter_func`` are adjusted.
"""
if not isinstance(partial_attributes, (dict, list)):
raise ValueError(
f"Invalid type for attributes: {type(partial_attributes)}, expected dictionary or list of dict."
)
if isinstance(partial_attributes, list) and not all(
isinstance(attr, dict) for attr in partial_attributes
):
raise ValueError("All elements in attributes list must be of type dict.")
for name, module in quant_model.named_modules():
if _match_quantizer(wildcard_or_filter_func, name, module, parent_class, quant_model):
module = cast("TensorQuantizer | SequentialQuantizer", module) # for type checker
if isinstance(partial_attributes, list):
if not isinstance(module, SequentialQuantizer):
raise ValueError(
f"Attributes is a list but {module} is not a SequentialQuantizer."
)
module.set_from_attribute_config(partial_attributes)
elif isinstance(module, SequentialQuantizer):
# Broadcast the dict to all sub-quantizers.
module.set_from_attribute_config([partial_attributes] * len(module))
else:
module.set_from_attribute_config(partial_attributes)
@contextmanager
def set_quantizer_by_cfg_context(quant_model: nn.Module, quant_cfg: QuantizeQuantCfgType):
"""Context manager that temporarily applies a quantization config and restores the original state on exit.
Calls :func:`set_quantizer_by_cfg` on entry and reverts every
:class:`TensorQuantizer <nn.modules.tensor_quantizer.TensorQuantizer>` in
``quant_model`` to its original attributes on exit.
.. caution::
Changing stateful attributes such as ``calibrator`` inside this context may produce
unexpected behavior because those objects are not deep-copied during save/restore.
Args:
quant_model: A quantized PyTorch model whose quantizers will be temporarily reconfigured.
quant_cfg: A quantization config (or list of
:class:`QuantizerCfgEntry <.config.QuantizerCfgEntry>` dicts) passed directly to
:func:`set_quantizer_by_cfg`. Sequential ``cfg`` lists are not allowed.
Yields:
None — the context body runs with the new quantizer attributes active.
"""
quant_cfg = normalize_quant_cfg_list(quant_cfg)
for entry in quant_cfg:
if isinstance(entry.get("cfg"), list):
raise ValueError(
"Sequential cfg lists are not allowed in set_quantizer_by_cfg_context. "
"Use only single-dict cfg entries."
)
original_attributes: dict[str, dict] = {}
original_types: dict[str, type] = {}
for name, module in quant_model.named_modules():
if isinstance(module, SequentialQuantizer):
# SequentialQuantizer.get_modelopt_state does not support properties_only;
# save per-sub-quantizer state so we can fully reconstruct on restore.
original_attributes[name] = {
"is_sequential_quantizer": True,
"sub_states": [tq.get_modelopt_state(properties_only=True) for tq in module],
}
original_types[name] = SequentialQuantizer
elif isinstance(module, TensorQuantizer):
original_attributes[name] = module.get_modelopt_state(properties_only=True)
original_types[name] = TensorQuantizer
set_quantizer_by_cfg(quant_model, quant_cfg)
yield
# Restore original quantizer types and attributes. If set_quantizer_by_cfg downgraded a
# SequentialQuantizer to a TensorQuantizer (or vice-versa), we need to re-create the
# original module type before restoring attributes.
for name, module in list(quant_model.named_modules()):
if name not in original_attributes:
continue
set_quantizer_attribute(quant_model, pattern, cfg)
orig_type = original_types[name]
if orig_type is SequentialQuantizer and not isinstance(module, SequentialQuantizer):
# Restore the SequentialQuantizer that was downgraded
saved = original_attributes[name]
parent_name, _, attr_name = name.rpartition(".")
parent_module = quant_model.get_submodule(parent_name) if parent_name else quant_model
module = SequentialQuantizer(*(TensorQuantizer() for _ in saved["sub_states"]))
setattr(parent_module, attr_name, module)
for tq, sub_state in zip(module, saved["sub_states"]):
tq.set_from_modelopt_state(sub_state, properties_only=True)
elif orig_type is TensorQuantizer and not isinstance(module, TensorQuantizer):
parent_name, _, attr_name = name.rpartition(".")
parent_module = quant_model.get_submodule(parent_name) if parent_name else quant_model
module = TensorQuantizer()
setattr(parent_module, attr_name, module)
module.set_from_modelopt_state(original_attributes[name], properties_only=True)
elif orig_type is TensorQuantizer:
module.set_from_modelopt_state(original_attributes[name], properties_only=True)
elif orig_type is SequentialQuantizer:
saved = original_attributes[name]
for tq, sub_state in zip(module, saved["sub_states"]):
tq.set_from_modelopt_state(sub_state, properties_only=True)
def set_quantizer_attribute(
quant_model: nn.Module,
wildcard_or_filter_func: str | Callable,
attribute: QuantizerAttributeConfig
| list[QuantizerAttributeConfig]
| dict[
str | Callable,
QuantizerAttributeConfig | list[QuantizerAttributeConfig],
]
| dict
| list[dict],
parent_class: type | None = None,
attribute: Any,
parent_class: type[nn.Module] | None = None,
):
"""Finegrained adjustment of quantizer attribute by wildcard or filter function.
Args:
quant_model: A pytorch model
wildcard_or_filter_func: a wildcard string or a filter function. The wildcard string is matched
against the quantizer module names. The quantizer modules are
instances of
:class:`TensorQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.TensorQuantizer>`.
The filter function takes a quantized module name as input and returns ``True`` if the
quantizer should be adjusted and ``False`` otherwise.
attribute: An instance of :class:`QuantizerAttributeConfig <.config.QuantizerAttributeConfig>` or an equivalent
dictionary or a list of these two types.
If ``attribute`` is a list, the matched
:class:`TensorQuantizer <nn.modules.tensor_quantizer.TensorQuantizer>`
modules will be replaced with :class:`SequentialQuantizer <nn.modules.tensor_quantizer.SequentialQuantizer>`
modules having one quantizer for each attribute instance from the list.
See
:meth:`set_from_attribute_config() <nn.modules.tensor_quantizer.TensorQuantizer.set_from_attribute_config>`
for more details on the supported attributes and their types.
parent_class: (Optional) The parent class of the quantizer modules matching ``wildcard_or_filter_func`` which
should be adjusted. If ``None``, all the matching quantizer modules will be adjusted.
"""
for name, module in quant_model.named_modules():
if isinstance(module, (TensorQuantizer, SequentialQuantizer)):
if isinstance(wildcard_or_filter_func, str):
if not fnmatch.fnmatch(name, wildcard_or_filter_func):
continue
elif callable(wildcard_or_filter_func):
if not wildcard_or_filter_func(name):
continue
else:
raise NotImplementedError(f"Unsupported type {type(wildcard_or_filter_func)}")
if parent_class is not None and not isinstance(
quant_model.get_submodule(".".join(name.split(".")[:-1])), parent_class
):
continue
if isinstance(attribute, list) and not isinstance(module, SequentialQuantizer):
parent_module = quant_model.get_submodule(name.rpartition(".")[0])
module = SequentialQuantizer(*(TensorQuantizer() for _ in range(len(attribute))))
setattr(parent_module, name.split(".")[-1], module)
elif isinstance(attribute, list) and len(attribute) != len(module):
warnings.warn(
f"The number of attributes ({len(attribute)}) does not match the number of "
f"quantizers of {module} leading to partial assignment.",
)
module.set_from_attribute_config(attribute)
@contextmanager
def set_quantizer_by_cfg_context(quant_model: nn.Module, quant_cfg: QuantizeQuantCfgType | dict):
"""Context manager for setting quantizer attributes using `quant_cfg`.
The set attributes will be reset to the original attributes after exiting the context manager.
See :meth:`set_quantizer_by_cfg` for more details.
Use this context manager with caution. Changing certain attributes of the quantizer such as
`calibrator` can lead to unexpected behavior.
"""
assert not any(cfg for cfg in quant_cfg.values() if isinstance(cfg, (list, tuple))), (
"list of config not support."
"""Deprecated: use :func:`set_quantizer_attributes_partial` instead."""
warnings.warn(
"set_quantizer_attribute is deprecated, use set_quantizer_attributes_partial "
"or set_quantizer_attributes_full instead.",
DeprecationWarning,
stacklevel=2,
)
return set_quantizer_attributes_partial(
quant_model, wildcard_or_filter_func, attribute, parent_class
)
original_attributes = {}
for name, module in quant_model.named_modules():
if isinstance(module, TensorQuantizer):
original_attributes[name] = module.get_modelopt_state(properties_only=True)
set_quantizer_by_cfg(quant_model, quant_cfg)
yield
for name, module in quant_model.named_modules():
if isinstance(module, TensorQuantizer):
module.set_from_modelopt_state(original_attributes[name], properties_only=True)
def register(original_cls: nn.Module, quantized_cls: nn.Module):
+3 -1
View File
@@ -1100,7 +1100,9 @@ def awq_lite(
self.awq_lite.num_cache_steps += 1
self.awq_lite.num_tokens += input.numel() / input.shape[-1]
if self.awq_lite.is_input_quantized:
with set_quantizer_by_cfg_context(self.input_quantizer, {"*": {"enable": True}}):
with set_quantizer_by_cfg_context(
self.input_quantizer, [{"quantizer_name": "*", "enable": True}]
):
max_calibrate(self.input_quantizer, lambda quantizer: quantizer(input), False)
return out_actual
+27 -22
View File
@@ -30,13 +30,15 @@ from modelopt.torch.opt import apply_mode
from modelopt.torch.opt.searcher import ForwardLoop
from modelopt.torch.opt.utils import forward_with_reshard
from modelopt.torch.quantization.config import QuantizeConfig
from modelopt.torch.quantization.conversion import set_quantizer_by_cfg
from modelopt.torch.quantization.conversion import (
set_quantizer_attributes_partial,
set_quantizer_by_cfg,
)
from modelopt.torch.utils import atomic_print
from .algorithms import AutoQuantizeGradientSearcher, AutoQuantizeKLDivSearcher, QuantRecipe
from .algorithms import get_auto_quantize_config as _get_auto_quantize_config
from .config import QuantizeAlgoCfgType
from .conversion import set_quantizer_attribute
from .mode import QuantizeModeRegistry, get_modelike_from_algo_cfg
from .nn import QuantModule, TensorQuantizer
from .utils import is_quantized
@@ -159,13 +161,15 @@ def quantize(
:class:`QuantizeConfig <modelopt.torch.quantization.config.QuantizeConfig>` specifying the
values for keys ``"quant_cfg"`` and ``"algorithm"``.
It is basically a dictionary specifying the values for keys ``"quant_cfg"`` and ``"algorithm"``.
The ``"quant_cfg"`` key specifies the quantization configurations.
The ``"quant_cfg"`` key specifies the quantization configurations as an ordered list of
:class:`QuantizerCfgEntry <modelopt.torch.quantization.config.QuantizerCfgEntry>` dicts.
The ``"algorithm"`` key specifies the ``algorithm`` argument to
:meth:`calibrate <modelopt.torch.quantization.model_quant.calibrate>`.
Quantization configurations is a dictionary mapping wildcards or filter functions
to its quantizer attributes. The wildcards or filter functions are matched
against the quantizer module names. The quantizer modules have names ending with
Each entry in the ``"quant_cfg"`` list has a ``"quantizer_name"`` wildcard matched
against quantizer module names, an optional ``"cfg"`` dict of quantizer attributes,
and an optional ``"enable"`` toggle. Entries are applied in list order; later entries
override earlier ones. The quantizer modules have names ending with
``weight_quantizer`` and ``input_quantizer`` and they perform weight quantization and
input quantization (or activation quantization) respectively. The quantizer modules
are instances of
@@ -178,17 +182,15 @@ def quantize(
.. code-block::python
config = {
"quant_cfg": {
"quant_cfg": [
# Disable all quantizers by default
{"quantizer_name": "*", "enable": False},
# "num_bits" specifies the number of bits for quantization
# "axis" specifies the axis for quantization
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": -1},
# Default quantization settings
"default": {"num_bits": 8, "axis": None},
}
"algorithm": "max"
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": -1}},
],
"algorithm": "max",
}
See :ref:`Quantization Formats <quantization-formats>` to learn more about the supported
@@ -323,10 +325,13 @@ def auto_quantize(
.. code-block:: python
INT8_CUSTOM_QUANT_CFG = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
},
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None},
},
],
"algorithm": "smoothquant",
}
@@ -527,7 +532,7 @@ def auto_quantize(
"checkpoint": checkpoint,
}
# Disable all quantizers; AutoQuantize will enable the needed ones
set_quantizer_by_cfg(model, {"*": {"enable": False}})
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
searcher.search(model, constraints, config=search_config) # type: ignore[arg-type]
return model, searcher.state_dict()
@@ -574,12 +579,12 @@ def get_auto_quantize_config(search_state, constraints=None, verbose=False):
def disable_quantizer(model: nn.Module, wildcard_or_filter_func: str | Callable):
"""Disable quantizer by wildcard or filter function."""
set_quantizer_attribute(model, wildcard_or_filter_func, {"enable": False})
set_quantizer_attributes_partial(model, wildcard_or_filter_func, {"enable": False})
def enable_quantizer(model: nn.Module, wildcard_or_filter_func: str | Callable):
"""Enable quantizer by wildcard or filter function."""
set_quantizer_attribute(model, wildcard_or_filter_func, {"enable": True})
set_quantizer_attributes_partial(model, wildcard_or_filter_func, {"enable": True})
@atomic_print
@@ -203,8 +203,8 @@ class TensorQuantizer(nn.Module):
# Optional quantizer cache for caching quantizer related encoding or tensors.
self._quantizer_cache = None
def set_from_attribute_config(self, attribute_cfg: QuantizerAttributeConfig | dict):
"""Set quantizer attributes from attribute_dict.
def set_from_attribute_config(self, attribute_cfg: QuantizerAttributeConfig | dict[str, Any]):
"""Set quantizer attributes from attribute_cfg.
The attributes are defined in
:class:`QuantizerAttributeConfig <modelopt.torch.quantization.config.QuantizerAttributeConfig>`.
@@ -218,12 +218,27 @@ class TensorQuantizer(nn.Module):
calib_cls, args, kwargs = standardize_constructor_args(val)
return calib_cls(*args, **kwargs)
def _axis_setter(val):
if getattr(self, "_calibrator", None) is not None:
self._calibrator._axis = val
return val
def _block_sizes_setter(val):
if val is not None:
# block_sizes and axis are mutually exclusive; clear axis when block_sizes is set
setattr(self, "_axis", None)
if getattr(self, "_calibrator", None) is not None:
self._calibrator._axis = None
return val
# Some attributes need custom handling.
# By default, attributes from config are mapped to a name ``f"_{attribute}"``
_custom_setters: dict[str, tuple[str, Callable]] = {
"enable": ("_disabled", lambda val: val is False),
"type": ("_dynamic", lambda val: val == "dynamic"),
"calibrator": ("_calibrator", _calibrator_setter),
"axis": ("_axis", _axis_setter),
"block_sizes": ("_block_sizes", _block_sizes_setter),
"backend": ("backend", lambda val: val),
"backend_extra_args": ("backend_extra_args", lambda val: val or {}),
"use_constant_amax": ("_use_constant_amax", lambda val: val),
@@ -1408,10 +1423,7 @@ class SequentialQuantizer(nn.Sequential):
return {"num_quantizers": len(self), "is_sequential_quantizer": True}
def set_from_attribute_config(
self,
attributes: list[dict[str, Any] | QuantizerAttributeConfig]
| dict[str, Any]
| QuantizerAttributeConfig,
self, attributes: list[QuantizerAttributeConfig] | list[dict[str, Any]]
):
"""Set the attributes of contained quantizers from a list of attribute_dicts."""
if not isinstance(attributes, (list, tuple)):
@@ -27,6 +27,7 @@ from torch.distributed.fsdp import FSDPModule, MixedPrecisionPolicy, fully_shard
from torch.distributed.fsdp._fully_shard._fsdp_param import FSDPParam
from torch.distributed.tensor import Replicate
from modelopt.torch.quantization.config import QuantizerCfgEntry
from modelopt.torch.utils import get_unwrapped_name, print_rank_0
if TYPE_CHECKING:
@@ -310,11 +311,15 @@ def calibrate_with_adapters(model, args):
def disable_lora_quantizers_in_config(config, layers):
"""Turns off input, weight, and output quantizers for LoRA weights and LoRALinear layers in config."""
config["quant_cfg"]["*lora*"] = {"enable": False}
config["quant_cfg"].append({"quantizer_name": "*lora*", "enable": False})
for layer in layers:
config["quant_cfg"][f"*{layer}.input_quantizer"] = {"enable": False}
config["quant_cfg"][f"*{layer}.weight_quantizer"] = {"enable": False}
config["quant_cfg"][f"*{layer}.output_quantizer"] = {"enable": False}
config["quant_cfg"].append({"quantizer_name": f"*{layer}.input_quantizer", "enable": False})
config["quant_cfg"].append(
{"quantizer_name": f"*{layer}.weight_quantizer", "enable": False}
)
config["quant_cfg"].append(
{"quantizer_name": f"*{layer}.output_quantizer", "enable": False}
)
return config
@@ -823,13 +828,25 @@ def fsdp2_aware_weight_update(root_model, modules_to_update, reshard=True):
def update_quant_cfg_with_kv_cache_quant(
quant_cfg: dict[str, Any], kv_cache_quant_cfg: dict[str, Any]
quant_cfg: dict[str, Any], kv_cache_quant_cfg: list[QuantizerCfgEntry]
) -> dict[str, Any]:
"""Update the quant_cfg with the kv cache quant_cfg."""
"""Update the quant_cfg with the kv cache quant_cfg.
Args:
quant_cfg: The outer quantization config dict (with ``"quant_cfg"`` and ``"algorithm"`` keys).
kv_cache_quant_cfg: A list of :class:`QuantizerCfgEntry
<modelopt.torch.quantization.config.QuantizerCfgEntry>` dicts for KV cache quantization,
typically ``some_kv_cfg["quant_cfg"]``.
Returns:
A deep copy of ``quant_cfg`` with the KV cache entries appended to ``quant_cfg["quant_cfg"]``.
"""
# If quant_cfg["quant_cfg"] is None, it corresponds to only kv cache quantization case
quant_cfg = copy.deepcopy(quant_cfg)
quant_cfg["quant_cfg"] = quant_cfg.get("quant_cfg") or {"default": {"enable": False}}
quant_cfg["quant_cfg"].update(kv_cache_quant_cfg)
inner: list[QuantizerCfgEntry] = quant_cfg.get("quant_cfg") or [
{"quantizer_name": "*", "enable": False}
]
quant_cfg["quant_cfg"] = inner + list(kv_cache_quant_cfg)
# Set default algorithm for kv cache quantization if not provided.
if not quant_cfg.get("algorithm"):
@@ -208,7 +208,7 @@ def set_sparse_attention_attribute(
):
"""Set sparse attention attributes for modules matching pattern.
Similar to quantization's set_quantizer_attribute.
Similar to quantization's set_quantizer_attributes_partial.
Args:
model: Model to configure
@@ -16,49 +16,52 @@
metadata:
recipe_type: ptq
description: FP8 per-tensor weight and activation (W8A8), FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*input_quantizer':
num_bits: e4m3
axis:
'*weight_quantizer':
num_bits: e4m3
axis:
default:
- quantizer_name: '*'
enable: false
'*block_sparse_moe.gate*':
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*mlp.gate.*':
enable: false
'*mlp.shared_expert_gate.*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
- quantizer_name: '*input_quantizer'
cfg:
num_bits: e4m3
axis:
- quantizer_name: '*weight_quantizer'
cfg:
num_bits: e4m3
axis:
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*block_sparse_moe.gate*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*mlp.gate.*'
enable: false
- quantizer_name: '*mlp.shared_expert_gate.*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
@@ -16,57 +16,60 @@
metadata:
recipe_type: ptq
description: NVFP4 MLP/MoE weight only (W4A16), FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*'
enable: false
- quantizer_name: '*weight_quantizer'
enable: true
'*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*input_quantizer'
enable: true
default:
enable: false
'*block_sparse_moe.gate*':
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*mlp.gate.*':
enable: false
'*mlp.shared_expert_gate.*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*block_sparse_moe.gate*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*mlp.gate.*'
enable: false
- quantizer_name: '*mlp.shared_expert_gate.*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# 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");
@@ -16,71 +16,76 @@
metadata:
recipe_type: ptq
description: NVFP4 static weight and dynamic activation for expert layers only (W4A4), FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*mlp.experts*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*'
enable: false
- quantizer_name: '*mlp.experts*weight_quantizer'
enable: true
'*mlp.experts*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*mlp.experts*input_quantizer'
enable: true
'*block_sparse_moe*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*weight_quantizer'
enable: true
'*block_sparse_moe*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*input_quantizer'
enable: true
default:
enable: false
'*block_sparse_moe.gate*':
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*mlp.gate.*':
enable: false
'*mlp.shared_expert_gate.*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*block_sparse_moe.gate*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*mlp.gate.*'
enable: false
- quantizer_name: '*mlp.shared_expert_gate.*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
@@ -16,71 +16,76 @@
metadata:
recipe_type: ptq
description: NVFP4 static weight and dynamic activation for all linear layers (W4A4), FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*mlp*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*'
enable: false
- quantizer_name: '*mlp*weight_quantizer'
enable: true
'*mlp*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*mlp*input_quantizer'
enable: true
'*block_sparse_moe*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*weight_quantizer'
enable: true
'*block_sparse_moe*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*input_quantizer'
enable: true
default:
enable: false
'*block_sparse_moe.gate*':
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*mlp.gate.*':
enable: false
'*mlp.shared_expert_gate.*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*block_sparse_moe.gate*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*mlp.gate.*'
enable: false
- quantizer_name: '*mlp.shared_expert_gate.*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
@@ -16,85 +16,92 @@
metadata:
recipe_type: ptq
description: NVFP4 static weight and dynamic activation for all linear layers including output projections, FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*mlp*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*'
enable: false
- quantizer_name: '*mlp*weight_quantizer'
enable: true
'*mlp*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*mlp*input_quantizer'
enable: true
'*block_sparse_moe*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*weight_quantizer'
enable: true
'*block_sparse_moe*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*block_sparse_moe*input_quantizer'
enable: true
'*o_proj*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*o_proj*weight_quantizer'
enable: true
'*o_proj*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*o_proj*input_quantizer'
enable: true
default:
enable: false
'*block_sparse_moe.gate*':
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*mlp.gate.*':
enable: false
'*mlp.shared_expert_gate.*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*block_sparse_moe.gate*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*mlp.gate.*'
enable: false
- quantizer_name: '*mlp.shared_expert_gate.*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
@@ -16,69 +16,74 @@
metadata:
recipe_type: ptq
description: NVFP4 static weight and dynamic activation for MoE/MLP projections (W4A4), FP8 KV cache, max calibration.
ptq_cfg:
quantize:
algorithm: max
quant_cfg:
'*moe*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*'
enable: false
- quantizer_name: '*moe*weight_quantizer'
enable: true
'*moe*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*moe*input_quantizer'
enable: true
'*mlp*weight_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*mlp*weight_quantizer'
enable: true
'*mlp*input_quantizer':
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*mlp*input_quantizer'
enable: true
'*share_expert*':
enable: false
'*moe.gate.*':
enable: false
default:
enable: false
'*linear_attn.conv1d*':
enable: false
'*lm_head*':
enable: false
'*mixer.conv1d*':
enable: false
'*output_layer*':
enable: false
'*proj_out.*':
enable: false
'*router*':
enable: false
output.*:
enable: false
nn.BatchNorm1d:
'*':
enable: false
nn.BatchNorm2d:
'*':
enable: false
nn.BatchNorm3d:
'*':
enable: false
nn.LeakyReLU:
'*':
enable: false
'*[kv]_bmm_quantizer':
num_bits: e4m3
cfg:
block_sizes:
-1: 16
type: dynamic
scale_bits: e4m3
num_bits: e2m1
- quantizer_name: '*[kv]_bmm_quantizer'
enable: true
cfg:
num_bits: e4m3
- quantizer_name: '*share_expert*'
enable: false
- quantizer_name: '*moe.gate.*'
enable: false
- quantizer_name: '*linear_attn.conv1d*'
enable: false
- quantizer_name: '*lm_head*'
enable: false
- quantizer_name: '*mixer.conv1d*'
enable: false
- quantizer_name: '*output_layer*'
enable: false
- quantizer_name: '*proj_out.*'
enable: false
- quantizer_name: '*router*'
enable: false
- quantizer_name: 'output.*'
enable: false
- parent_class: 'nn.BatchNorm1d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm2d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.BatchNorm3d'
quantizer_name: '*'
enable: false
- parent_class: 'nn.LeakyReLU'
quantizer_name: '*'
enable: false
+176 -97
View File
@@ -85,162 +85,241 @@ class SmallQKVModel(torch.nn.Module):
# Quantization configs
partial_fp8_config = {
"quant_cfg": {
"*.1.weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*.1.input_quantizer": {"num_bits": (4, 3), "axis": None},
"default": {"num_bits": 8, "enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*.1.weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.1.input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
],
"algorithm": "max",
}
partial_w4a8_config = {
"quant_cfg": {
"*.2.weight_quantizer": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}, "enable": True},
{"num_bits": (4, 3), "axis": None, "enable": True},
],
"*.2.input_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True},
"default": {"num_bits": 8, "enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.2.weight_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": (4, 3), "axis": None},
],
"enable": True,
},
{
"quantizer_name": "*.2.input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
},
],
"algorithm": "awq_lite",
}
partial_nvfp4_config = {
"quant_cfg": {
"*.1.weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.1.weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*.1.input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.1.input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*.2.weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.2.weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*.2.input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.2.input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"default": {"enable": False},
},
],
"algorithm": "max",
}
partial_nvfp4_awq_config = {
"quant_cfg": {
"*.2.weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.2.weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*.2.input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.2.input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*.1.weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.1.weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": False,
},
"*.1.input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*.1.input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": False,
},
"default": {"enable": False},
},
],
"algorithm": "awq_lite",
}
partial_int4_awq_config = {
"quant_cfg": {
"*.2.weight_quantizer": {
"num_bits": 4,
"block_sizes": {-1: 128, "type": "static"},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.2.weight_quantizer",
"cfg": {"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
"enable": True,
},
"*.2.input_quantizer": {"enable": False},
"default": {"enable": False},
},
{"quantizer_name": "*.2.input_quantizer", "enable": False},
],
"algorithm": {"method": "awq_lite", "alpha_step": 0.1},
# "algorithm": {"method": "awq_full", "alpha_step": 0.1, "max_co_batch_size": 1024},
# "algorithm": {"method": "awq_clip", "max_co_batch_size": 2048},
}
partial_fp8_kv_cache_config = {
"quant_cfg": {
"*.1.weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*.1.input_quantizer": {"num_bits": (4, 3), "axis": None},
"*output_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*.1.weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.1.input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
},
],
"algorithm": "max",
}
partial_int8_kv_cache_config = {
"quant_cfg": {
"*.1.weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*.1.input_quantizer": {"num_bits": (4, 3), "axis": None},
"*output_quantizer": {"num_bits": 8, "axis": None, "enable": True},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*.1.weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.1.input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": 8, "axis": None},
"enable": True,
},
],
"algorithm": "max",
}
partial_nvfp4_kv_cache_config = {
"quant_cfg": {
"*.1.weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*.1.input_quantizer": {"num_bits": (4, 3), "axis": None},
"*[kv]_bmm_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*.1.weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.1.input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"default": {"enable": False},
},
],
"algorithm": "max",
}
only_weight_quantizer_fp8_config = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True},
"*input_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"*output_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
],
"algorithm": "max",
}
only_input_quantizer_fp8_config = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"*input_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True},
"*output_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
},
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
],
"algorithm": "max",
}
only_output_quantizer_fp8_config = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"*input_quantizer": {"num_bits": (4, 3), "axis": None, "enable": False},
"*output_quantizer": {"num_bits": (4, 3), "axis": None, "enable": True},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": False,
},
{
"quantizer_name": "*output_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"enable": True,
},
],
"algorithm": "max",
}
@@ -29,11 +29,11 @@ TEST_MODELS = {SimpleLinear, SimpleConv, SimpleConvLinear}
def onnx_export_tester(model, device, num_bits, per_channel_quantization, constant_folding, dtype):
axis = 0 if per_channel_quantization else None
config = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": num_bits, "axis": axis},
"*input_quantizer": {"num_bits": num_bits},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": num_bits, "axis": axis}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": num_bits}},
],
"algorithm": "max",
}
@@ -29,25 +29,28 @@ from modelopt.torch.quantization.nn.modules.tensor_quantizer import SequentialQu
from modelopt.torch.quantization.utils import is_quantized_linear
from modelopt.torch.utils import torch_to
INT4_AWQ_FULL_CFG = mtq.INT4_AWQ_CFG.copy()
INT4_AWQ_FULL_CFG = copy.deepcopy(mtq.INT4_AWQ_CFG)
INT4_AWQ_FULL_CFG["algorithm"] = "awq_full"
INT4_AWQ_CLIP_CFG = mtq.INT4_AWQ_CFG.copy()
INT4_AWQ_CLIP_CFG = copy.deepcopy(mtq.INT4_AWQ_CFG)
INT4_AWQ_CLIP_CFG["algorithm"] = "awq_clip"
# SVDQuant test cfg
INT4_SVDQUANT_CFG = mtq.INT4_AWQ_CFG.copy()
INT4_SVDQUANT_CFG = copy.deepcopy(mtq.INT4_AWQ_CFG)
INT4_SVDQUANT_CFG["algorithm"] = {"method": "svdquant", "lowrank": 8}
# SVDQuant test cfg
FP4_SVDQUANT_CFG = mtq.NVFP4_AWQ_LITE_CFG.copy()
FP4_SVDQUANT_CFG = copy.deepcopy(mtq.NVFP4_AWQ_LITE_CFG)
FP4_SVDQUANT_CFG["algorithm"] = {"method": "svdquant", "lowrank": 8}
def get_awq_config(algorithm="awq_lite", block_size=8):
config = copy.deepcopy(mtq.INT4_AWQ_CFG)
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {-1: block_size}
for entry in config["quant_cfg"]:
if entry["quantizer_name"] == "*weight_quantizer":
entry.setdefault("cfg", {})["block_sizes"] = {-1: block_size}
break
if "algorithm" not in config or not isinstance(config["algorithm"], dict):
config["algorithm"] = {}
+7 -11
View File
@@ -77,7 +77,7 @@ def test_kv_rotate(rotate_fp32):
model = nn.Sequential(SDPAAttention())
mtq.replace_quant_module(model)
set_quantizer_by_cfg(model, {"*": {"enable": False}})
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
dummy_input = SDPAAttention.get_input(device="cuda")
output_ref = model(dummy_input)
if rotate_fp32:
@@ -86,11 +86,9 @@ def test_kv_rotate(rotate_fp32):
rotate = True
with set_quantizer_by_cfg_context(
model,
{
"*[qk]_bmm_quantizer": {
"rotate": rotate,
},
},
[
{"quantizer_name": "*[qk]_bmm_quantizer", "cfg": {"rotate": rotate}},
],
):
output_test = model(dummy_input)
assert torch.allclose(output_ref, output_test, atol=0.05)
@@ -98,11 +96,9 @@ def test_kv_rotate(rotate_fp32):
# Test the rotation is actually applied by turning on only one of the query, key quantizers
with set_quantizer_by_cfg_context(
model,
{
"*k_bmm_quantizer": {
"rotate": rotate,
},
},
[
{"quantizer_name": "*k_bmm_quantizer", "cfg": {"rotate": rotate}},
],
):
output_test1 = model(dummy_input)
assert not torch.allclose(output_ref, output_test1, atol=0.05)
@@ -21,7 +21,7 @@ import pytest
import torch
import torch.nn as nn
from modelopt.torch.quantization import set_quantizer_attribute
from modelopt.torch.quantization.conversion import set_quantizer_attributes_partial
from modelopt.torch.quantization.nn import QuantModuleRegistry
@@ -44,7 +44,7 @@ def test_no_quant_proj(original_cls, bidirectional, bias):
rnn_object_original = copy.deepcopy(rnn_object)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"enable": False})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"enable": False})
test_input = torch.randn((3, 2, 8), device="cuda")
@@ -15,6 +15,8 @@
"""High-level tests for quantization."""
import copy
import pytest
from _test_utils.torch.quantization.models import SimpleConv, SimpleConvLinear, SimpleLinear
from _test_utils.torch.quantization.quantize_common import (
@@ -29,20 +31,26 @@ import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.extensions import get_cuda_ext_mx
NVFP4_WEIGHT_ACT_MSE_CFG = {
"quant_cfg": {
"*weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "static", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{
"quantizer_name": "*weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "static", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
},
],
"algorithm": {
"method": "mse",
"step_size": 0.25,
@@ -52,17 +60,18 @@ NVFP4_WEIGHT_ACT_MSE_CFG = {
}
NVFP4_WEIGHT_MSE_FP8_SWEEP_CFG = {
"quant_cfg": {
"*weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "static", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{
"quantizer_name": "*weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "static", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*input_quantizer": {
"enable": False,
},
},
{"quantizer_name": "*input_quantizer", "enable": False},
],
"algorithm": {
"method": "mse",
"fp8_scale_sweep": True,
@@ -123,7 +132,10 @@ def test_quantize(model_cls, config):
if config == mtq.FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG:
# reduce block sizes for simple testing models
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {-1: 8, -2: 8}
config = copy.deepcopy(config)
for entry in config["quant_cfg"]:
if entry.get("quantizer_name") == "*weight_quantizer":
entry.setdefault("cfg", {})["block_sizes"] = {-1: 8, -2: 8}
model = model_cls().cuda()
calib_data = [model.get_input().cuda() for _ in range(8)]
quantize_model_and_forward(model, config, calib_data)
@@ -15,6 +15,7 @@
"""High-level tests for real weight-only quantization."""
import copy
import fnmatch
import pytest
@@ -47,10 +48,14 @@ def test_real_quantize(model_cls, config):
# update config to fit test cases
if config == mtq.INT4_AWQ_CFG:
# reduce block sizes for simple testing models
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {
-1: 16,
"scale_bits": 8,
}
config = copy.deepcopy(config)
for entry in config["quant_cfg"]:
if entry.get("quantizer_name") == "*weight_quantizer":
entry.setdefault("cfg", {})["block_sizes"] = {
-1: 16,
"scale_bits": 8,
}
break
if model_cls is SimpleConv or model_cls is SimpleConvLinear:
pytest.skip(
"INT4_AWQ_CFG requires even number of elements on last dimension for weights."
@@ -101,10 +106,14 @@ def test_save_restore(model_cls, config):
# update config to fit test cases
if config == mtq.INT4_AWQ_CFG:
# reduce block sizes for simple testing models
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {
-1: 16,
"scale_bits": 8,
}
config = copy.deepcopy(config)
for entry in config["quant_cfg"]:
if entry.get("quantizer_name") == "*weight_quantizer":
entry.setdefault("cfg", {})["block_sizes"] = {
-1: 16,
"scale_bits": 8,
}
break
if model_cls is SimpleConv or model_cls is SimpleConvLinear:
pytest.skip(
"INT4_AWQ_CFG requires even number of elements on last dimension for weights."
@@ -33,23 +33,32 @@ from modelopt.torch.peft.lora.layer import LoRAModule
from modelopt.torch.utils.plugins import megatron_prefill
NVFP4_DEFAULT_CONFIG = {
"quant_cfg": {
"*weight_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*input_quantizer": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
{
"quantizer_name": "*input_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"axis": None,
},
"enable": True,
},
"*output_quantizer": {"enable": False},
"*output_layer*": {"enable": False}, # Note: only output_layer is disabled.
"default": {"enable": False},
},
{"quantizer_name": "*output_quantizer", "enable": False},
{
"quantizer_name": "*output_layer*",
"enable": False,
}, # Note: only output_layer is disabled.
],
"algorithm": "max",
}
@@ -84,15 +84,15 @@ def test_convert_apex_parallel_linear(distributed_setup_size_1):
assert hasattr(module, "weight_quantizer")
assert hasattr(module, "output_quantizer")
mtq.set_quantizer_attribute(model_test, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_test, "*", {"enable": False})
x = model_ref.get_dummy_input().cuda()
out_1 = model_ref(x)
out_2 = model_test(x)
assert torch.allclose(out_1, out_2)
mtq.set_quantizer_attribute(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attribute(model_test, "*weight_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*weight_quantizer", {"enable": True})
model_ref = RegularQuantModelForTP().cuda()
model_ref.load_state_dict(model_test.state_dict())
@@ -82,15 +82,15 @@ def test_convert_megatron_parallel_linear(distributed_setup_size_1):
assert hasattr(module, "weight_quantizer")
assert hasattr(module, "output_quantizer")
mtq.set_quantizer_attribute(model_test, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_test, "*", {"enable": False})
x = model_ref.get_dummy_input().cuda()
out_1 = model_ref(x)
out_2 = model_test(x)
assert torch.allclose(out_1, out_2)
mtq.set_quantizer_attribute(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attribute(model_test, "*weight_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*weight_quantizer", {"enable": True})
model_ref = RegularQuantModelForTP().cuda()
model_ref.load_state_dict(model_test.state_dict(), strict=False)
@@ -304,7 +304,7 @@ def _test_sharded_state_dict(
):
# Must disable output_layer quantization since output_layer amax cannot be restore via
# sharded_state_dict. All output_layer quantizers state are removed.
config["quant_cfg"]["*output_layer*"] = {"enable": False}
config["quant_cfg"].append({"quantizer_name": "*output_layer*", "enable": False})
if modelopt_version is not None:
mto.conversion.__version__ = modelopt_version
@@ -383,36 +383,44 @@ def _test_sharded_state_dict(
mixed_precision_config = copy.deepcopy(mtq.W4A8_AWQ_BETA_CFG)
mixed_precision_config["quant_cfg"].update(
{
"*.1.*": {"enable": False},
"*.2.*weight_quantizer": {"num_bits": (4, 3), "axis": None},
"*.2.*input_quantizer": {"num_bits": (4, 3), "axis": None},
"*.3.*weight_quantizer.0": {"num_bits": 8, "axis": 0},
"*.3.*weight_quantizer.1": {"enable": False},
"*.3.*input_quantizer": {"num_bits": 8, "axis": None},
}
mixed_precision_config["quant_cfg"].extend(
[
{"quantizer_name": "*.1.*", "enable": False},
{"quantizer_name": "*.2.*weight_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.2.*input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{"quantizer_name": "*.3.*weight_quantizer.0", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*.3.*weight_quantizer.1", "enable": False},
{"quantizer_name": "*.3.*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
]
)
mixed_block_size_config = copy.deepcopy(mtq.INT4_BLOCKWISE_WEIGHT_ONLY_CFG)
mixed_block_size_config["quant_cfg"].update(
{
"*.1.*": {"enable": False},
"*.2.*weight_quantizer": {"num_bits": 4, "block_sizes": {-1: 64}, "enable": True},
"*.2.*input_quantizer": {"num_bits": (4, 3), "axis": None},
"*.3.*weight_quantizer": {"num_bits": 4, "block_sizes": {-1: 128, -2: 64}, "enable": True},
"*.3.*input_quantizer": {"num_bits": 8, "axis": None},
}
mixed_block_size_config["quant_cfg"].extend(
[
{"quantizer_name": "*.1.*", "enable": False},
{
"quantizer_name": "*.2.*weight_quantizer",
"cfg": {"num_bits": 4, "block_sizes": {-1: 64}},
"enable": True,
},
{"quantizer_name": "*.2.*input_quantizer", "cfg": {"num_bits": (4, 3), "axis": None}},
{
"quantizer_name": "*.3.*weight_quantizer",
"cfg": {"num_bits": 4, "block_sizes": {-1: 128, -2: 64}},
"enable": True,
},
{"quantizer_name": "*.3.*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
]
)
# Combined NVFP4 GEMM + KV cache quantization config
NVFP4_GEMM_KV_CFG = copy.deepcopy(mtq.NVFP4_DEFAULT_CFG)
NVFP4_GEMM_KV_CFG["quant_cfg"].update(mtq.NVFP4_KV_CFG["quant_cfg"])
NVFP4_GEMM_KV_CFG["quant_cfg"].extend(mtq.NVFP4_KV_CFG["quant_cfg"])
# Combined FP8 GEMM + KV cache quantization config
FP8_GEMM_KV_CFG = copy.deepcopy(mtq.FP8_DEFAULT_CFG)
FP8_GEMM_KV_CFG["quant_cfg"].update(mtq.FP8_KV_CFG["quant_cfg"])
FP8_GEMM_KV_CFG["quant_cfg"].extend(mtq.FP8_KV_CFG["quant_cfg"])
@pytest.mark.parametrize(
@@ -13,6 +13,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
import pytest
import torch
import torch.nn as nn
@@ -73,7 +75,11 @@ def test_quantize(model_cls, config):
if config == mtq.FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG:
# reduce block sizes for simple testing models
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {-1: 8, -2: 8}
config = copy.deepcopy(config)
for entry in config["quant_cfg"]:
if entry.get("quantizer_name") == "*weight_quantizer":
entry["cfg"]["block_sizes"] = {-1: 8, -2: 8}
break
model = model_cls().cuda()
calib_data = [model.get_input().cuda() for _ in range(1)]
quantize_model_and_forward(model, config, calib_data)
+58 -20
View File
@@ -15,6 +15,8 @@
"""Unit tests for modelopt.recipe.loader and modelopt.recipe.loader.load_config."""
import re
import pytest
from modelopt.recipe.config import ModelOptPTQRecipe, RecipeType
@@ -36,10 +38,10 @@ key: val
CFG_RECIPE_MISSING_TYPE = """\
metadata:
description: Missing recipe_type.
ptq_cfg: {}
quantize: {}
"""
CFG_RECIPE_MISSING_PTQ_CFG = """\
CFG_RECIPE_MISSING_quantize = """\
metadata:
recipe_type: ptq
"""
@@ -86,7 +88,7 @@ def test_load_recipe_builtin_with_suffix():
recipe = load_recipe("general/ptq/fp8_default-fp8_kv.yml")
assert recipe.recipe_type == RecipeType.PTQ
assert isinstance(recipe, ModelOptPTQRecipe)
assert recipe.ptq_cfg
assert recipe.quantize
def test_load_recipe_builtin_without_suffix():
@@ -112,11 +114,11 @@ _BUILTIN_PTQ_RECIPES = [
@pytest.mark.parametrize("recipe_path", _BUILTIN_PTQ_RECIPES)
def test_load_recipe_all_builtins(recipe_path):
"""Smoke-test: every built-in PTQ recipe loads without error and has ptq_cfg."""
"""Smoke-test: every built-in PTQ recipe loads without error and has quantize."""
recipe = load_recipe(recipe_path)
assert recipe.recipe_type == RecipeType.PTQ
assert isinstance(recipe, ModelOptPTQRecipe)
assert recipe.ptq_cfg
assert recipe.quantize
# ---------------------------------------------------------------------------
@@ -138,11 +140,11 @@ def test_load_recipe_missing_recipe_type_raises(tmp_path):
load_recipe(bad)
def test_load_recipe_missing_ptq_cfg_raises(tmp_path):
"""load_recipe raises ValueError when ptq_cfg is absent for a PTQ recipe."""
def test_load_recipe_missing_quantize_raises(tmp_path):
"""load_recipe raises ValueError when quantize is absent for a PTQ recipe."""
bad = tmp_path / "bad.yml"
bad.write_text(CFG_RECIPE_MISSING_PTQ_CFG)
with pytest.raises(ValueError, match="ptq_cfg"):
bad.write_text(CFG_RECIPE_MISSING_quantize)
with pytest.raises(ValueError, match="quantize"):
load_recipe(bad)
@@ -160,28 +162,28 @@ def test_load_recipe_unsupported_type_raises(tmp_path):
def test_load_recipe_dir(tmp_path):
"""load_recipe loads a recipe from a directory with recipe.yml + ptq_cfg.yml."""
"""load_recipe loads a recipe from a directory with recipe.yml + quantize.yml."""
(tmp_path / "recipe.yml").write_text(
"metadata:\n recipe_type: ptq\n description: Dir test.\n"
)
(tmp_path / "ptq_cfg.yml").write_text("algorithm: max\nquant_cfg: {}\n")
(tmp_path / "quantize.yml").write_text("algorithm: max\nquant_cfg: []\n")
recipe = load_recipe(tmp_path)
assert recipe.recipe_type == RecipeType.PTQ
assert recipe.description == "Dir test."
assert recipe.ptq_cfg == {"algorithm": "max", "quant_cfg": {}}
assert recipe.quantize == {"algorithm": "max", "quant_cfg": []}
def test_load_recipe_dir_missing_recipe_raises(tmp_path):
"""load_recipe raises ValueError when recipe.yml is absent from the directory."""
(tmp_path / "ptq_cfg.yml").write_text("algorithm: max\nquant_cfg: {}\n")
(tmp_path / "quantize.yml").write_text("algorithm: max\nquant_cfg: {}\n")
with pytest.raises(ValueError, match="recipe descriptor"):
load_recipe(tmp_path)
def test_load_recipe_dir_missing_ptq_cfg_raises(tmp_path):
"""load_recipe raises ValueError when ptq_cfg.yml is absent from the directory."""
def test_load_recipe_dir_missing_quantize_raises(tmp_path):
"""load_recipe raises ValueError when quantize.yml is absent from the directory."""
(tmp_path / "recipe.yml").write_text("metadata:\n recipe_type: ptq\n")
with pytest.raises(ValueError, match="ptq_cfg"):
with pytest.raises(ValueError, match="quantize"):
load_recipe(tmp_path)
@@ -200,13 +202,49 @@ def test_load_recipe_dir_missing_ptq_cfg_raises(tmp_path):
],
)
def test_general_ptq_yaml_matches_config_dicts(yaml_path, model_cfg_name, kv_cfg_name):
"""Each general/ptq YAML's merged quant_cfg matches the corresponding config.py dicts."""
"""Each general/ptq YAML's quant_cfg list matches the merged Python config dicts."""
import json
import modelopt.torch.quantization.config as qcfg
from modelopt.torch.quantization.config import normalize_quant_cfg_list
model_cfg = getattr(qcfg, model_cfg_name)
kv_cfg = getattr(qcfg, kv_cfg_name)
yaml_data = load_config(yaml_path)
ptq = yaml_data["ptq_cfg"]
assert {**model_cfg["quant_cfg"], **kv_cfg["quant_cfg"]} == ptq["quant_cfg"]
assert model_cfg["algorithm"] == ptq["algorithm"]
def _normalize_fpx(val):
"""Normalize FPx representations to a canonical ``[E, M]`` list.
Python configs may use tuple form ``(E, M)`` or string alias ``"eEmM"``;
YAML always uses the string form. Both are converted to ``[E, M]`` so the
comparison is representation-agnostic.
"""
if isinstance(val, str):
m = re.fullmatch(r"e(\d+)m(\d+)", val)
if m:
return [int(m.group(1)), int(m.group(2))]
if isinstance(val, tuple) and len(val) == 2 and all(isinstance(x, int) for x in val):
return list(val)
if isinstance(val, dict):
return {str(k): _normalize_fpx(v) for k, v in val.items()}
return val
def _normalize_entries(raw_entries):
"""Normalize a raw quant_cfg list to a canonical, JSON-serialisable form."""
entries = normalize_quant_cfg_list(list(raw_entries))
result = []
for entry in entries:
e = {k: v for k, v in entry.items() if v is not None}
if "cfg" in e and e["cfg"] is not None:
e["cfg"] = _normalize_fpx(e["cfg"])
result.append(e)
return result
def _sort_key(entry):
return json.dumps(entry, sort_keys=True, default=str)
python_entries = _normalize_entries(model_cfg["quant_cfg"] + kv_cfg["quant_cfg"])
yaml_entries = _normalize_entries(yaml_data["quantize"]["quant_cfg"])
assert sorted(python_entries, key=_sort_key) == sorted(yaml_entries, key=_sort_key)
assert model_cfg["algorithm"] == yaml_data["quantize"]["algorithm"]
@@ -61,10 +61,10 @@ class SDPAAttention(nn.Module):
kv_cache_config = {
"quant_cfg": {
"*[kv]_bmm_quantizer": {"num_bits": 4, "enable": True},
"*softmax_quantizer": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*[kv]_bmm_quantizer", "cfg": {"num_bits": 4}, "enable": True},
{"quantizer_name": "*softmax_quantizer", "enable": False},
],
"algorithm": "max",
}
@@ -87,7 +87,7 @@ def test_convert_conv1d():
assert hasattr(module, "weight_quantizer")
assert hasattr(module, "output_quantizer")
mtq.set_quantizer_attribute(model_test, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_test, "*", {"enable": False})
x = torch.randn(2, 3)
out_1 = model_ref(x)
@@ -95,8 +95,8 @@ def test_convert_conv1d():
assert torch.allclose(out_1, out_2)
mtq.set_quantizer_attribute(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attribute(model_test, "*weight_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*input_quantizer", {"enable": True})
mtq.set_quantizer_attributes_partial(model_test, "*weight_quantizer", {"enable": True})
model_ref = PytorchModel()
model_ref.load_state_dict(model_test.state_dict())
@@ -136,7 +136,7 @@ def test_dbrx():
expertglu_ref.w1,
)
mtq.set_quantizer_attribute(model_test, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_test, "*", {"enable": False})
x = torch.randn(1, 4, 32)
out_1 = model_ref(x)
@@ -193,7 +193,13 @@ def test_quantized_transformers_save_restore(tmp_path, model_cls, quant_config):
tiny_llama_dir = create_tiny_llama_dir(tmp_path)
# update config to fit test cases
if quant_config == mtq.INT4_AWQ_CFG:
quant_config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {-1: 16}
import copy
quant_config = copy.deepcopy(quant_config)
for entry in quant_config["quant_cfg"]:
if entry["quantizer_name"] == "*weight_quantizer":
entry.setdefault("cfg", {})["block_sizes"] = {-1: 16}
break
else:
raise ValueError(f"Unsupported quant_config: {quant_config}")
@@ -48,7 +48,7 @@ def test_convert_loralinear():
assert hasattr(module, "weight_quantizer")
assert hasattr(module, "output_quantizer")
mtq.set_quantizer_attribute(model_test, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_test, "*", {"enable": False})
tf_output_tester(model_ref, model_test)
+28 -15
View File
@@ -28,7 +28,7 @@ from modelopt.torch.quantization.algorithms import (
QuantRecipeHparam,
estimate_quant_compression,
)
from modelopt.torch.quantization.config import _default_disabled_quantizer_cfg
from modelopt.torch.quantization.config import _base_disable_all, _default_disabled_quantizer_cfg
from modelopt.torch.utils.distributed import DistributedProcessGroup
@@ -110,11 +110,12 @@ def test_quant_recipe_hparam():
# use this config to test custom quantization config
INT8_CUSTOM_QUANT_TEST_CFG = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
**_default_disabled_quantizer_cfg,
},
"quant_cfg": [
*_base_disable_all,
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
*_default_disabled_quantizer_cfg,
],
"algorithm": "smoothquant",
}
@@ -230,14 +231,22 @@ def test_auto_quantize_disabled_layers_no_poison():
INT4INT8_AWQ_CFG = {
"quant_cfg": {
"*weight_quantizer": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}, "enable": True},
{"num_bits": 8, "axis": None, "enable": True},
],
"*input_quantizer": {"num_bits": 8, "axis": None, "enable": True},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": 8, "axis": None},
],
"enable": True,
},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None},
"enable": True,
},
],
"algorithm": "awq_lite",
}
@@ -480,7 +489,11 @@ def test_get_auto_quantize_config(method):
# Use stored best recipe
config = mtq.get_auto_quantize_config(search_state)
assert "quant_cfg" in config
assert config["quant_cfg"]["*"] == {"enable": False}
assert isinstance(config["quant_cfg"], list)
assert any(
entry["quantizer_name"] == "*" and entry.get("enable") is False
for entry in config["quant_cfg"]
)
assert config["algorithm"] == "max"
# Re-solve with different constraints
@@ -22,10 +22,10 @@ import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.nn import TensorQuantizer
INT8_CFG = {
"quant_cfg": {
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
},
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 8, "axis": None}},
],
"algorithm": "max",
}
@@ -15,6 +15,9 @@
"""Test of quantization config validations."""
import pytest
from pydantic import ValidationError
from modelopt.torch.quantization.config import (
FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG,
FP8_DEFAULT_CFG,
@@ -22,7 +25,10 @@ from modelopt.torch.quantization.config import (
INT4_AWQ_CFG,
NVFP4_DEFAULT_CFG,
W4A8_AWQ_BETA_CFG,
QuantizeConfig,
find_quant_cfg_entry_by_path,
need_calibration,
normalize_quant_cfg_list,
)
@@ -33,3 +39,435 @@ def test_need_calibration():
assert need_calibration(INT4_AWQ_CFG)
assert need_calibration(W4A8_AWQ_BETA_CFG)
assert need_calibration(NVFP4_DEFAULT_CFG)
def test_need_calibration_with_list_cfg():
"""need_calibration must handle sequential (list) cfg entries without crashing."""
# Static list-cfg on a non-weight quantizer → needs calibration
cfg_static = {
"quant_cfg": [
{
"quantizer_name": "*input_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": (4, 3)},
],
"enable": True,
},
],
"algorithm": "max",
}
assert need_calibration(cfg_static)
# Dynamic list-cfg on a non-weight quantizer → no calibration needed
cfg_dynamic = {
"quant_cfg": [
{
"quantizer_name": "*input_quantizer",
"cfg": [{"num_bits": (4, 3), "type": "dynamic"}],
"enable": True,
},
],
"algorithm": "max",
}
assert not need_calibration(cfg_dynamic)
class TestNormalizeQuantCfgList:
def test_new_format_passthrough(self):
"""New-format entries are returned unchanged (only canonical defaults added)."""
raw = [{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}}]
result = normalize_quant_cfg_list(raw)
assert len(result) == 1
assert result[0]["quantizer_name"] == "*weight_quantizer"
assert result[0]["cfg"] == {"num_bits": 8, "axis": 0}
assert result[0]["enable"] is True # defaulted
def test_new_format_enable_false(self):
"""Explicit enable=False is preserved."""
raw = [{"quantizer_name": "*", "enable": False}]
result = normalize_quant_cfg_list(raw)
assert result[0]["enable"] is False
assert result[0]["cfg"] is None # defaulted
def test_new_format_explicit_enable_true_no_cfg(self):
"""Explicit enable=True with no cfg is valid and cfg defaults to None."""
raw = [{"quantizer_name": "*", "enable": True}]
result = normalize_quant_cfg_list(raw)
assert result[0]["enable"] is True
assert result[0]["cfg"] is None
def test_legacy_single_key_dict(self):
"""Legacy {'*path': {attrs}} is converted to new format."""
raw = [{"*weight_quantizer": {"num_bits": 8, "axis": 0}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["quantizer_name"] == "*weight_quantizer"
assert result[0]["cfg"] == {"num_bits": 8, "axis": 0}
assert result[0]["enable"] is True # defaulted
def test_legacy_single_key_dict_with_enable(self):
"""Legacy {'*path': {'enable': False}} splits enable out from cfg."""
raw = [{"*input_quantizer": {"enable": False}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["quantizer_name"] == "*input_quantizer"
assert result[0]["enable"] is False
assert result[0]["cfg"] is None
def test_legacy_nn_class_scoped(self):
"""Legacy {'nn.Linear': {'*': {attrs}}} is converted with parent_class."""
raw = [{"nn.Linear": {"*": {"enable": False}}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["parent_class"] == "nn.Linear"
assert result[0]["quantizer_name"] == "*"
assert result[0]["enable"] is False
def test_normalization_cfg_defaults_to_none(self):
"""Entries without cfg get cfg=None after normalization."""
raw = [{"quantizer_name": "*lm_head*", "enable": False}]
result = normalize_quant_cfg_list(raw)
assert "cfg" in result[0]
assert result[0]["cfg"] is None
def test_normalization_enable_defaults_to_true(self):
"""Entries with cfg but no enable get enable=True after normalization."""
raw = [{"quantizer_name": "*", "cfg": {"num_bits": 4}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["enable"] is True
def test_empty_list(self):
"""Empty list is returned unchanged."""
assert normalize_quant_cfg_list([]) == []
def test_multiple_entries_order_preserved(self):
"""The order of entries is preserved."""
raw = [
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}},
]
result = normalize_quant_cfg_list(raw)
assert result[0]["quantizer_name"] == "*"
assert result[1]["quantizer_name"] == "*weight_quantizer"
def test_error_on_quantizer_name_only(self):
"""Entry with only quantizer_name and no cfg or enable is rejected."""
with pytest.raises(ValueError, match="must specify 'cfg', 'enable'"):
normalize_quant_cfg_list([{"quantizer_name": "*"}])
def test_error_on_empty_dict(self):
"""An empty dict entry is rejected."""
with pytest.raises(ValueError):
normalize_quant_cfg_list([{}])
def test_error_on_multi_key_legacy_dict(self):
"""A multi-key legacy dict (no quantizer_name, no nn.* keys) is rejected."""
with pytest.raises(ValueError):
normalize_quant_cfg_list([{"*weight_quantizer": {}, "*input_quantizer": {}}])
def test_new_format_with_list_cfg(self):
"""cfg can be a list of dicts for SequentialQuantizer."""
raw = [
{
"quantizer_name": "*weight_quantizer",
"cfg": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": (4, 3)},
],
}
]
result = normalize_quant_cfg_list(raw)
assert len(result) == 1
assert result[0]["cfg"] == raw[0]["cfg"]
assert result[0]["enable"] is True
def test_legacy_flat_dict_conversion(self):
"""Legacy flat dict {'*': {...}, '*weight_quantizer': {...}} is converted to list."""
raw = {"*": {"enable": False}, "*weight_quantizer": {"num_bits": 8, "axis": 0}}
result = normalize_quant_cfg_list(raw)
assert len(result) == 2
assert result[0]["quantizer_name"] == "*"
assert result[0]["enable"] is False
assert result[0]["cfg"] is None
assert result[1]["quantizer_name"] == "*weight_quantizer"
assert result[1]["cfg"] == {"num_bits": 8, "axis": 0}
assert result[1]["enable"] is True
def test_legacy_enable_only_produces_cfg_none(self):
"""Legacy {'*': {'enable': False}} should produce cfg=None, not cfg={}."""
raw = [{"*": {"enable": False}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["cfg"] is None
assert result[0]["enable"] is False
def test_legacy_nn_class_enable_only_produces_cfg_none(self):
"""Legacy nn.* scoped format with only enable produces cfg=None."""
raw = [{"nn.Linear": {"*": {"enable": False}}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["cfg"] is None
assert result[0]["enable"] is False
assert result[0]["parent_class"] == "nn.Linear"
def test_legacy_default_key(self):
"""Legacy 'default' key is converted to quantizer_name='*'."""
raw = [{"default": {"enable": False}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["quantizer_name"] == "*"
assert result[0]["enable"] is False
assert result[0]["cfg"] is None
def test_legacy_default_key_with_cfg(self):
"""Legacy 'default' key with cfg attributes maps to '*'."""
raw = [{"default": {"num_bits": 8, "axis": None}}]
result = normalize_quant_cfg_list(raw)
assert result[0]["quantizer_name"] == "*"
assert result[0]["cfg"] == {"num_bits": 8, "axis": None}
assert result[0]["enable"] is True
def test_legacy_flat_dict_with_default_key(self):
"""Legacy flat dict containing 'default' key converts it to '*'."""
raw = {"default": {"enable": False}, "*weight_quantizer": {"num_bits": 8}}
result = normalize_quant_cfg_list(raw)
default_entries = [e for e in result if e["quantizer_name"] == "*"]
assert len(default_entries) == 1
assert default_entries[0]["enable"] is False
def test_legacy_nn_class_multi_key(self):
"""Legacy nn.* scoped format with multiple sub-keys produces multiple entries."""
raw = [
{
"nn.Linear": {
"*input_quantizer": {"enable": False},
"*weight_quantizer": {"num_bits": 4},
}
}
]
result = normalize_quant_cfg_list(raw)
assert len(result) == 2
paths = {e["quantizer_name"] for e in result}
assert paths == {"*input_quantizer", "*weight_quantizer"}
for e in result:
assert e["parent_class"] == "nn.Linear"
def test_legacy_nn_class_with_cfg(self):
"""Legacy nn.* scoped format with actual quantizer attributes (not just enable)."""
raw = [{"nn.Linear": {"*weight_quantizer": {"num_bits": 4, "axis": 0}}}]
result = normalize_quant_cfg_list(raw)
assert len(result) == 1
assert result[0]["parent_class"] == "nn.Linear"
assert result[0]["quantizer_name"] == "*weight_quantizer"
assert result[0]["cfg"] == {"num_bits": 4, "axis": 0}
assert result[0]["enable"] is True
def test_legacy_list_valued_cfg(self):
"""Legacy dict format with list-valued cfg (SequentialQuantizer) normalizes correctly."""
raw = [
{
"*weight_quantizer": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
{"num_bits": 8, "axis": 0},
]
}
]
result = normalize_quant_cfg_list(raw)
assert len(result) == 1
assert result[0]["quantizer_name"] == "*weight_quantizer"
assert isinstance(result[0]["cfg"], list)
assert len(result[0]["cfg"]) == 2
assert result[0]["cfg"][0]["num_bits"] == 4
assert result[0]["cfg"][1]["num_bits"] == 8
assert result[0]["enable"] is True
class TestFindQuantCfgEntry:
def test_finds_last_match(self):
"""When multiple entries share the same quantizer_name, returns the last one."""
entries = normalize_quant_cfg_list(
[
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}},
{"quantizer_name": "*input_quantizer", "cfg": {"num_bits": 4}},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4}},
]
)
result = find_quant_cfg_entry_by_path(entries, "*weight_quantizer")
assert result["cfg"] == {"num_bits": 4}
def test_exact_match_only(self):
"""Does not do fnmatch — only exact string equality on quantizer_name."""
entries = normalize_quant_cfg_list(
[{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}}]
)
with pytest.raises(KeyError):
find_quant_cfg_entry_by_path(entries, "model.layer.weight_quantizer")
def test_raises_on_missing(self):
"""Raises KeyError when no entry matches."""
entries = normalize_quant_cfg_list(
[{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}}]
)
with pytest.raises(KeyError):
find_quant_cfg_entry_by_path(entries, "*input_quantizer")
def test_single_entry(self):
entries = normalize_quant_cfg_list([{"quantizer_name": "*", "enable": False}])
result = find_quant_cfg_entry_by_path(entries, "*")
assert result["enable"] is False
def test_empty_list(self):
with pytest.raises(KeyError):
find_quant_cfg_entry_by_path([], "*")
def test_need_calibration_with_legacy_dict_format():
"""need_calibration should accept legacy dict-format quant_cfg without crashing."""
legacy_config = {
"quant_cfg": {"*input_quantizer": {"num_bits": 8, "axis": None}},
"algorithm": "max",
}
assert need_calibration(legacy_config)
def test_need_calibration_with_legacy_list_of_single_key_dicts():
"""need_calibration should accept legacy list-of-single-key-dicts format."""
legacy_config = {
"quant_cfg": [{"*input_quantizer": {"num_bits": 8, "axis": None}}],
"algorithm": "max",
}
assert need_calibration(legacy_config)
class TestMatchQuantizerCfg:
"""Tests for _match_quantizer_cfg in algorithms.py."""
def test_wildcard_matches_bare_name(self):
"""'*weight_quantizer' matches bare 'weight_quantizer'."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}}]
)
matched, enable = _match_quantizer_cfg(quant_cfg, "weight_quantizer")
assert matched == {"num_bits": 8}
assert enable is True
def test_star_matches_any_bare_name(self):
"""'*' matches any bare quantizer name."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list([{"quantizer_name": "*", "enable": False}])
matched, enable = _match_quantizer_cfg(quant_cfg, "weight_quantizer")
assert matched is None # enable-only entry has cfg=None
assert enable is False
def test_path_scoped_pattern_matches_matching_suffix(self):
"""'*mlp*weight_quantizer' matches bare 'weight_quantizer' (suffix match)."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[{"quantizer_name": "*mlp*weight_quantizer", "cfg": {"num_bits": 4}}]
)
matched, enable = _match_quantizer_cfg(quant_cfg, "weight_quantizer")
assert matched == {"num_bits": 4}
def test_path_scoped_pattern_does_not_match_different_suffix(self):
"""'*mlp*weight_quantizer' does NOT match bare 'input_quantizer'."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[{"quantizer_name": "*mlp*weight_quantizer", "cfg": {"num_bits": 4}}]
)
matched, enable = _match_quantizer_cfg(quant_cfg, "input_quantizer")
assert matched is None
assert enable is None
def test_last_match_wins(self):
"""Later entries override earlier ones."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}},
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 4}},
]
)
matched, _ = _match_quantizer_cfg(quant_cfg, "weight_quantizer")
assert matched == {"num_bits": 4}
def test_no_match_returns_none(self):
"""No matching entry returns (None, None)."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8}}]
)
matched, enable = _match_quantizer_cfg(quant_cfg, "output_quantizer")
assert matched is None
assert enable is None
def test_bracket_pattern_matches_correctly(self):
"""'*[kv]_bmm_quantizer' matches 'k_bmm_quantizer' and 'v_bmm_quantizer'."""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[{"quantizer_name": "*[kv]_bmm_quantizer", "cfg": {"num_bits": (4, 3)}}]
)
matched_k, _ = _match_quantizer_cfg(quant_cfg, "k_bmm_quantizer")
matched_v, _ = _match_quantizer_cfg(quant_cfg, "v_bmm_quantizer")
matched_w, _ = _match_quantizer_cfg(quant_cfg, "weight_quantizer")
assert matched_k is not None
assert matched_v is not None
assert matched_w is None
def test_path_scoped_does_not_overmatch(self):
"""'*mixer*weight_quantizer' should NOT match 'input_quantizer'.
Regression test: the old rsplit('*') logic would strip to 'weight_quantizer' and
overmatch any quantizer ending in 'weight_quantizer', but should not match unrelated names.
"""
from modelopt.torch.quantization.algorithms import _match_quantizer_cfg
quant_cfg = normalize_quant_cfg_list(
[
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*mixer*weight_quantizer", "cfg": {"num_bits": 4}},
]
)
# input_quantizer should only match the disable-all, not the mixer pattern
matched, enable = _match_quantizer_cfg(quant_cfg, "input_quantizer")
assert matched is None # cfg is None (enable-only entry)
assert enable is False
class TestQuantizeConfigValidators:
"""Tests for QuantizeConfig Pydantic field validators."""
def test_normalize_validator_converts_legacy_dict(self):
"""The 'before' validator auto-normalizes legacy dict format."""
cfg = QuantizeConfig(
quant_cfg={"*": {"enable": False}, "*weight_quantizer": {"num_bits": 8}},
algorithm="max",
)
assert isinstance(cfg.quant_cfg, list)
assert all("quantizer_name" in e for e in cfg.quant_cfg)
def test_validate_quant_cfg_entries_catches_invalid_cfg(self):
"""The 'after' validator surfaces QuantizerAttributeConfig errors early."""
with pytest.raises(ValidationError):
QuantizeConfig(
quant_cfg=[
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 8, "axis": 0, "block_sizes": {-1: 128}},
}
],
algorithm="max",
)
def test_validate_quant_cfg_entries_accepts_valid_cfg(self):
"""The 'after' validator passes for valid configs."""
cfg = QuantizeConfig(
quant_cfg=[
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{"quantizer_name": "*input_quantizer", "enable": False},
],
algorithm="max",
)
assert len(cfg.quant_cfg) == 2
@@ -42,16 +42,19 @@ def test_custom_backend_via_quantize():
model = torch.nn.Linear(16, 16, bias=False)
cfg = {
"quant_cfg": {
"*weight_quantizer": {
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {
"num_bits": 8,
"axis": None,
"backend": "dummy_backend",
"backend_extra_args": {"offset": 2.5},
},
"enable": True,
"num_bits": 8,
"axis": None,
"backend": "dummy_backend",
"backend_extra_args": {"offset": 2.5},
},
"default": {"enable": False},
},
],
"algorithm": "max",
}
@@ -88,10 +91,14 @@ def test_custom_backend_with_quantizer_cache():
model = torch.nn.Linear(16, 16, bias=False)
cfg = {
"quant_cfg": {
"*weight_quantizer": {"enable": True, "backend": "cached_backend"},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"backend": "cached_backend"},
"enable": True,
},
],
"algorithm": "max",
}
inputs = torch.randn(1, 16)
@@ -19,7 +19,7 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from modelopt.torch.quantization import set_quantizer_attribute, tensor_quant
from modelopt.torch.quantization import set_quantizer_attributes_partial, tensor_quant
from modelopt.torch.quantization.nn import QuantModuleRegistry
@@ -42,7 +42,7 @@ class TestQuantLeakyReLU:
negative_slope = 0.01
leaky_relu_object = nn.LeakyReLU(negative_slope=negative_slope)
quant_leaky_relu_object = QuantModuleRegistry.convert(leaky_relu_object)
set_quantizer_attribute(quant_leaky_relu_object, lambda name: True, {"axis": (1)})
set_quantizer_attributes_partial(quant_leaky_relu_object, lambda name: True, {"axis": (1)})
test_input = torch.randn(input_shape)
quant_input = tensor_quant.fake_tensor_quant(
@@ -20,7 +20,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from modelopt.torch.quantization import set_quantizer_attribute, tensor_quant
from modelopt.torch.quantization import tensor_quant
from modelopt.torch.quantization.conversion import set_quantizer_attributes_partial
from modelopt.torch.quantization.nn import QuantModuleRegistry
NUM_CHANNELS = 3
@@ -90,7 +91,7 @@ class TestQuantBatchNormND:
def test_fake_quant_per_channel(self, original_cls, input_shape):
batchnorm_object = original_cls(NUM_CHANNELS, affine=True)
quant_batchnorm_object = QuantModuleRegistry.convert(batchnorm_object)
set_quantizer_attribute(quant_batchnorm_object, lambda name: True, {"axis": (1)})
set_quantizer_attributes_partial(quant_batchnorm_object, lambda name: True, {"axis": (1)})
test_input = torch.randn(input_shape)
reduce_dims = list(range(len(test_input.shape)))
@@ -21,7 +21,8 @@ import pytest
import torch
import torch.nn as nn
from modelopt.torch.quantization import set_quantizer_attribute, tensor_quant
from modelopt.torch.quantization import tensor_quant
from modelopt.torch.quantization.conversion import set_quantizer_attributes_partial
from modelopt.torch.quantization.nn import QuantModuleRegistry
from modelopt.torch.quantization.nn.modules.quant_rnn import VFRNNForward
@@ -52,7 +53,7 @@ class TestQuantRNN:
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
rnn_object.eval()
rnn_object_original.eval()
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"enable": False})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"enable": False})
assert torch.allclose(
quant_rnn_object.weight_ih_l0, rnn_object_original.weight_ih_l0, atol=1e-6
@@ -86,7 +87,7 @@ class TestQuantRNN:
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
rnn_object.eval()
rnn_object_original.eval()
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"enable": False})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"enable": False})
assert torch.allclose(
quant_rnn_object.weight_ih_l0, rnn_object_original.weight_ih_l0, atol=1e-6
@@ -124,7 +125,7 @@ class TestQuantRNN:
rnn_object_original = copy.deepcopy(rnn_object)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"enable": False})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"enable": False})
test_input = torch.randn(INPUT_SHAPE)
@@ -150,7 +151,7 @@ class TestQuantRNN:
rnn_object_original = copy.deepcopy(rnn_object)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"enable": False})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"enable": False})
test_input = torch.randn([INPUT_SHAPE[1], INPUT_SHAPE[0], INPUT_SHAPE[2]])
@@ -176,7 +177,7 @@ class TestQuantRNN:
)
rnn_object_original = copy.deepcopy(rnn_object)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"axis": None})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"axis": None})
quant_rnn_object._disable_input_quantizers()
for name, weight in rnn_object_original.named_parameters():
@@ -205,7 +206,7 @@ class TestQuantRNN:
rnn_object = original_cls(HIDDEN_SIZE, HIDDEN_SIZE, NUM_LAYERS, bidirectional=bidirectional)
rnn_object_original = copy.deepcopy(rnn_object)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"axis": (0)})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"axis": (0)})
quant_rnn_object._disable_input_quantizers()
for name, weight in rnn_object_original.named_parameters():
@@ -234,7 +235,7 @@ class TestQuantRNN:
HIDDEN_SIZE, HIDDEN_SIZE, NUM_LAYERS, bidirectional=bidirectional, bias=True
)
quant_rnn_object = QuantModuleRegistry.convert(rnn_object)
set_quantizer_attribute(quant_rnn_object, lambda name: True, {"axis": None})
set_quantizer_attributes_partial(quant_rnn_object, lambda name: True, {"axis": None})
quant_rnn_object._disable_weight_quantizers()
num_directions = 2 if bidirectional else 1
@@ -32,41 +32,54 @@ 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": {
"*weight_quantizer": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}, "enable": True},
{"num_bits": 8, "axis": 0, "enable": True},
],
"*input_quantizer": {"num_bits": 8, "axis": None, "enable": True},
},
"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": {
"*weight_quantizer": {"num_bits": 8, "axis": 0},
"*input_quantizer": {"num_bits": 8, "axis": None},
},
"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": {
"*weight_quantizer": {
"num_bits": 8,
"axis": 0,
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 8, "axis": 0},
}, # Per-channel quantization
"*input_quantizer": {
"num_bits": 8,
"axis": (0, 1),
"type": "dynamic",
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": (0, 1), "type": "dynamic"},
}, # Dynamic per-token quantization
"default": {"enable": False},
},
],
"algorithm": "max",
}
@@ -77,14 +90,17 @@ class NewMaxCalibrator(MaxCalibrator):
quant_cfg_custom_calib = {
"quant_cfg": {
"*": {
"num_bits": 4,
"axis": None,
"quant_cfg": [
{
"quantizer_name": "*",
"cfg": {
"num_bits": 4,
"axis": None,
"calibrator": (NewMaxCalibrator, (4, None, False)),
},
"enable": True,
"calibrator": (NewMaxCalibrator, (4, None, False)),
}
},
],
"algorithm": "max",
}
@@ -131,7 +147,9 @@ def test_save_restore(model_cls, quant_config):
def test_quantize_invalid_cfg():
model = SimpleLinear()
config_invalid = {
"quant_cfg": {"*": {"num_bits": 4, "axis": 0, "block_sizes": {-1: 128}}},
"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."):
@@ -170,12 +188,22 @@ def test_custom_calib_config():
def test_class_wise_config():
model = SimpleConvLinear()
config = {
"quant_cfg": {
"nn.Linear": {"*": {"num_bits": 4, "axis": -1, "enable": True}},
"nn.Conv2d": {"*": {"num_bits": 8, "enable": True}},
"nn.BatchNorm2d": {"*": {"enable": False}},
"*output_quantizer": {"num_bits": 8, "enable": True},
},
"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",
}
@@ -222,33 +250,28 @@ def test_static_weight_dynamic_activations():
def test_block_sizes_axis_model():
REF_QUANT_CFG = { # noqa: N806
"quant_cfg": {
"*weight_quantizer": {
"num_bits": 8,
"axis": 0,
"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"},
},
"*input_quantizer": {
"num_bits": 8,
"axis": None,
"type": "dynamic",
},
"default": {"enable": False},
},
],
"algorithm": "max",
}
QUANT_CFG = { # noqa: N806
"quant_cfg": {
"*weight_quantizer": {
"num_bits": 8,
"block_sizes": {1: None},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*weight_quantizer",
"cfg": {"num_bits": 8, "block_sizes": {1: None}},
},
"*input_quantizer": {
"num_bits": 8,
"block_sizes": {0: None, 1: None},
"type": "dynamic",
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "block_sizes": {0: None, 1: None}, "type": "dynamic"},
},
"default": {"enable": False},
},
],
"algorithm": "max",
}
model_ref = SimpleLinear()
@@ -283,3 +306,184 @@ def test_quantize_twice():
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
@@ -47,7 +47,7 @@ def test_quantize_replace(model_cls):
assert not isinstance(module, nn.Conv2d) or _is_quantized_linear_conv(module)
assert not isinstance(module, nn.Linear) or _is_quantized_linear_conv(module)
mtq.set_quantizer_attribute(model_atq, "*", {"enable": False})
mtq.set_quantizer_attributes_partial(model_atq, "*", {"enable": False})
out_ref = model_ref(dummy_input)
out_atq = model_atq(dummy_input)
@@ -89,14 +89,18 @@ class TestQuantizerAttributeConfig:
WINT4INT8_CFG = {
"quant_cfg": {
"*weight_quantizer": [
{"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}, "enable": True},
{"num_bits": 8, "axis": 0, "enable": True},
],
"*input_quantizer": {"num_bits": 8, "enable": True},
"default": {"enable": False},
},
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"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}, "enable": True},
],
"algorithm": "awq_full",
}
@@ -109,10 +113,14 @@ def test_set_quantizer_cxt():
state_dict = model.state_dict()
output_ref = model(inputs)
mtq.set_quantizer_by_cfg(model, {"*output_quantizer": {"enable": True}})
mtq.set_quantizer_by_cfg(model, [{"quantizer_name": "*output_quantizer", "enable": True}])
with mtq.set_quantizer_by_cfg_context(
model, {"*": {"enable": False}, "*output_quantizer": {"enable": True}}
model,
[
{"quantizer_name": "*", "enable": False},
{"quantizer_name": "*output_quantizer", "enable": True},
],
):
for name, module in model.named_modules():
if not isinstance(module, TensorQuantizer):
@@ -123,7 +131,7 @@ def test_set_quantizer_cxt():
assert not module.is_enabled
mtq.calibrate(model, "max", lambda model: model(inputs * 10))
mtq.set_quantizer_by_cfg(model, {"*output_quantizer": {"enable": False}})
mtq.set_quantizer_by_cfg(model, [{"quantizer_name": "*output_quantizer", "enable": False}])
output_test = model(inputs)
assert torch.allclose(output_ref, output_test)