mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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'},
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
),
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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}")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"] = {}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
Reference in New Issue
Block a user