Files
Model-Optimizer/modelopt/torch/quantization/model_quant.py
T
Shengliang Xu c7ed23a103 Rename modelopt_recipes/huggingface to model_type with backward-compat alias (#2328)
### What does this PR do?

Type of change: Refactor + deprecation (recipe-library restructure,
backward compatible), plus an unrelated transformers-compat test fix.

Rename the architecture-specific recipe tier
`modelopt_recipes/huggingface/` to
`modelopt_recipes/model_type/`, making explicit that it holds recipes
**shared across
every checkpoint of a Hugging Face `model_type`** — as opposed to the
checkpoint-mirror
`models/<org>/<model_id>/` tier. The old `huggingface/` path keeps
working as a
deprecated backward-compat alias (a source-tree symlink plus a loader
alias), so no
saved `--recipe` path breaks.

- **Loader alias** (`modelopt/recipe/loader.py`): generalized so saved
`--recipe huggingface/<model_type>/...` paths rewrite to
`model_type/...`, alongside
the existing `huggingface/models/... -> models/...` rewrite (checked
first as the more
specific prefix). This keeps old paths resolving for pip-installed
wheels, where the
  source-tree symlinks don't survive.
- **Internal `$import`s**: rewritten from `huggingface/... ->
model_type/...` inside the
shipped recipes so they resolve without the symlink — mandatory for
wheels, since
  `$import` resolution goes through `config_loader` (no alias there).
- **Packaging** (`pyproject.toml`, `MANIFEST.in`): extended the
symlink-exclusion globs
so the recursive `**/*.yaml` package-data glob doesn't double-ship
recipes through the
`huggingface -> model_type` and `model_type/models -> ../models`
symlinks.
- **Docs / examples / skills / tests**: migrated all internal references
to the canonical
`model_type/`; `huggingface/` remains only in the deprecated-alias tests
and explanatory
  notes.
- **Unrelated fix (2nd commit):**
`tests/unit/torch/export/test_quant_aware_conversion.py`
  failed on transformers>=5.9, which dropped `base_model_prefix` from
`WeightTransform.__slots__` (the scoped-rule tests assigned it on the
now-slotted
object). Production `_scope_prefixes` already reads it via `getattr(...,
None)` and
degrades correctly, so there is no runtime change — the tests now set it
through a
helper that suppresses `AttributeError` across the supported
transformers range.

### Usage

```bash
# New canonical path
python examples/hf_ptq/hf_ptq.py --model <ckpt> \
    --recipe model_type/qwen3_vl/ptq/fp8_vision-kv_none

# Old path still works (deprecated backward-compat alias)
python examples/hf_ptq/hf_ptq.py --model <ckpt> \
    --recipe huggingface/qwen3_vl/ptq/fp8_vision-kv_none
```

```python
from modelopt.recipe import load_recipe

load_recipe("model_type/vit/ptq/fp8")    # canonical
load_recipe("huggingface/vit/ptq/fp8")   # deprecated alias, resolves to the same recipe
```

### Testing

- `tests/unit/recipe/` — **336 passed**, including the new
`test_load_recipe_huggingface_arch_backward_compat_alias` and the
updated
  structural/doc tests (`test_recipe_docs.py`).
- `tests/unit/torch/export/test_quant_aware_conversion.py` — **16
passed** (was 4 failed
  on transformers 5.9.0).
- Built an sdist **and** a wheel and inspected both manifests: each
recipe ships exactly
once (29 `model_type/`, 13 `models/`, 2 `timm/`, 162 total) with
**zero** `huggingface/` or
  `model_type/models/` duplicates and no build error on the symlinks.
- Simulated a wheel install (symlink-free extracted tree) and confirmed
`huggingface/<arch>/...`, `model_type/...`, and `huggingface/models/...`
all resolve via
  the loader alias — including a recipe that pulls internal `$import`s.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ — old `huggingface/...` recipe
paths keep resolving via the symlink + loader alias.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅ — backward-compat alias test
added; structural/doc tests updated to the new layout.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — Deprecations entry under 0.48.0. (The transformers-compat test fix
is not changelog-worthy.)
- Did you get Claude approval on this PR?: ❌ — not yet.

### Additional Information

The `model_type/models -> ../models` symlink is kept purely as a
backward-compat alias for
old `huggingface/models/<org>/<model_id>/...` paths; `model_type/` is
otherwise
architecture-only. If we ever want it strictly architecture-only, that
symlink can be
dropped later without breaking anything, since the loader rewrites
`huggingface/models/...`
straight to the top-level `models/` tier.


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- Added post-training quantization recipes for Gemma, Gemma 4,
MiniMax-M3, Nemotron, Qwen, Step-3.7, ViT, and other architectures.
- Added vision, multimodal, mixed-precision, and experts-only
quantization options.

- **Documentation**
- Standardized architecture-specific recipes under `model_type/` and
updated examples and guidance.

- **Compatibility**
- Legacy `huggingface/` recipe paths remain supported with deprecation
warnings.
  - Local recipe files now take precedence over built-in recipes.
  - Deprecated quantization-format flags warn when explicitly provided.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
2026-09-15 12:16:12 -07:00

1107 lines
50 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""User-facing quantization API."""
import fnmatch
import inspect
import os
import warnings
from collections.abc import Callable, Iterable, Sequence
from contextlib import contextmanager
from typing import Any, cast
import torch
import torch.nn as nn
import modelopt.torch.quantization as mtq
from modelopt.torch.opt import apply_mode
from modelopt.torch.opt.searcher import ConstraintsDict, ForwardLoop
from modelopt.torch.opt.utils import forward_with_reshard
from modelopt.torch.quantization.config import QuantizeConfig
from modelopt.torch.quantization.conversion import (
preserve_quantizer_attributes_context,
set_quantizer_attributes_partial,
set_quantizer_by_cfg,
)
from modelopt.torch.utils import atomic_print
from ._auto_quantize_cost import COST_MODEL_KV_CACHE
from .algorithms import AutoQuantizeGradientSearcher, AutoQuantizeKLDivSearcher, QuantRecipe
from .algorithms import get_auto_quantize_config as _get_auto_quantize_config
from .config import QuantizeAlgoCfgType
from .kv_cache_auto_quant import AutoQuantizeKVSearcher, get_kv_cache_auto_quantize_config
from .kv_cache_auto_quant import _validate_search_inputs as _validate_kv_cache_search_inputs
from .mode import QuantizeModeRegistry, get_modelike_from_algo_cfg
from .nn import QuantModule, SequentialQuantizer, TensorQuantizer
from .utils import is_quantized
__all__ = [
"auto_quantize",
"calibrate",
"compute_quantization_mse",
"disable_quantizer",
"enable_quantizer",
"fold_weight",
"get_auto_quantize_config",
"postprocess_amax",
"print_quant_summary",
"quantize",
"temporarily_fold_weights",
]
# TODO: Descriptors for the supported algorithms
def calibrate(
model: nn.Module,
algorithm: QuantizeAlgoCfgType = "max",
forward_loop: ForwardLoop | None = None,
) -> nn.Module:
"""Adjusts weights and scaling factors based on selected algorithms.
In order to calibrate using custom user defined calibration algorithm, refer to
:ref:`custom calibration algorithm <custom_calibration_algorithm>`
Args:
model: A pytorch model with quantizer modules.
algorithm: A string or dictionary specifying the calibration algorithm to use. Supported
algorithms are ``"max"``, ``"smoothquant"``, ``"awq_lite"``, ``"awq_full"``, and
``"awq_clip"``. If a dictionary is passed, the key ``"method"`` should specify the
calibration algorithm to use. Other key-value pairs in this dictionary will be passed
as kwargs to the algorithm.
An example dictionary argument:
``{"method": "awq_clip", "max_co_batch_size": 4096}``.
If ``None``, no calibration is performed.
forward_loop: A callable which takes the model as argument and forwards calibration data
through the model. This is not required for weight-only quantization with the ``"max"``
algorithm.
Returns: The calibrated pytorch model.
"""
if forward_loop is not None:
# get the number of arguments of forward_loop
num_args = len(inspect.signature(forward_loop).parameters)
if num_args == 0:
warnings.warn(
(
"forward_loop should take model as argument, but got forward_loop without any"
" arguments. This usage will be deprecated in future versions."
),
DeprecationWarning,
)
original_forward_loop = forward_loop
def forward_loop(model):
return original_forward_loop() # type: ignore[call-arg]
# move the model to eval mode
is_training = model.training
model.eval()
with forward_with_reshard(model):
apply_mode(
model,
mode=get_modelike_from_algo_cfg(algorithm),
mode_kwargs={"forward_loop": forward_loop},
)
for name, module in model.named_modules():
if isinstance(module, TensorQuantizer):
for attr_name in ["_amax", "_pre_quant_scale"]:
module.validate_attr(attr_name=attr_name, warn_error=True, name=name)
# TODO: Re-enable when the CUDA error: unspecified launch failure is fixed.
# clear_cuda_cache()
model.train(is_training)
return model
def postprocess_amax(model: nn.Module, key: str, post_process_fn) -> nn.Module:
"""Experimental API to postprocess the amax values after calibration."""
assert isinstance(key, str), "key should be a string"
for name, module in model.named_modules():
if not isinstance(module, TensorQuantizer):
continue
if not hasattr(module, "_amax"):
continue
if not fnmatch.fnmatch(name, key):
continue
module.amax = post_process_fn(module.amax)
return model
_SKIP_WEIGHT_QUANT_CHECK_ENV = "MODELOPT_SKIP_WEIGHT_QUANT_CHECK"
def _check_weight_quantization_took_effect(model: nn.Module, config: QuantizeConfig) -> None:
"""Raise when a config asks for weight quantization but no weight quantizer is enabled.
A config whose module patterns do not match the model is not an error to
:func:`set_quantizer_by_cfg` — every pattern simply matches nothing — so the run
proceeds through calibration and export and produces a checkpoint that is silently
unquantized (``"quant_algo": null`` with an empty ``quantized_layers``). That has
bitten several MoE architectures whose module naming differs from the wildcards in
the general recipes, and it is only noticed when someone reads the exported config.
By the time this runs, :func:`set_quantizer_by_cfg` (or the ``apply_mode`` conversion
that calls it) has already applied ``config`` to ``model``, so each quantizer's
``is_enabled`` *is* the true outcome of that application — checking it directly cannot
diverge from what the config actually did. An earlier version of this check instead
re-derived "did this pattern match anything?" via a separate matcher call, which missed
the case of two *different* overlapping patterns (e.g. ``*weight_quantizer`` enabling
something a later, broader ``*`` then disables): the narrower pattern registered as
"matched" even though the quantizer it matched ended up disabled.
A config that never asks for weight quantization (activation-only or KV-cache-only)
must not raise, so the check first looks at the config's own intent — via each
pattern's *final* entry, since entries apply in order and the last one for a pattern
wins — before looking at the model at all.
"""
if os.environ.get(_SKIP_WEIGHT_QUANT_CHECK_ENV) == "1":
return
# Later entries override earlier ones, so only each pattern's final state states intent.
# A pattern naming ``weight_quantizer`` explicitly (the common case, e.g.
# ``*weight_quantizer``, ``*.experts.*weight_quantizer``) is caught by the substring
# check. A broad wildcard that never mentions "weight" -- a bare ``"*"`` catch-all, or
# ``"*_quantizer"`` -- can still match weight quantizers at runtime, so it must count
# too, or a config built only from patterns like that would never trip the guard
# regardless of what the model contains. ``fnmatch`` against the literal probe string
# ``"weight_quantizer"`` catches those (a pattern matching that bare name is, by
# construction, asking for one) without replacing the substring check: the probe alone
# would miss ``*.experts.*weight_quantizer`` (there is no ``.experts.`` in the probe
# string), which is what recognizes model-scoped patterns like the Step / MoE recipes use.
last_entry_per_pattern = {entry.quantizer_name: entry for entry in config.quant_cfg}
weight_patterns = [
pattern
for pattern, entry in last_entry_per_pattern.items()
if entry.enable
and ("weight_quantizer" in pattern or fnmatch.fnmatch("weight_quantizer", pattern))
]
if not weight_patterns:
return
# `SequentialQuantizer.is_enabled` delegates to its first member, so a list-valued `cfg`'s
# quantizers are already covered here without naming `SequentialQuantizer` explicitly:
# `named_modules()` recurses into the container and yields those children too, individually,
# named `...weight_quantizer.0` / `.1` (the substring match below still applies to them).
if any(
module.is_enabled
for name, module in model.named_modules()
if isinstance(module, TensorQuantizer) and "weight_quantizer" in name
):
return
patterns = "\n ".join(sorted(weight_patterns))
raise RuntimeError(
"The quantization config asks for weight quantization but no weight quantizer is "
f"enabled, so nothing would be quantized. These patterns asked for it:\n {patterns}\n"
"Either the patterns do not match this architecture's module names (check the "
"model-specific recipes under modelopt_recipes/model_type/<model_type>/), or the "
"modules holding the weights were never converted to quantized modules (an "
"unsupported custom module, e.g. a trust_remote_code MoE layout).\n"
"Under pipeline parallelism, a rank whose local stage genuinely has none of the "
"targeted modules (e.g. a pure-attention stage under an experts-only recipe) hits "
"this too, while other ranks proceed into calibration -- a collective hang, not "
f"just a wrong per-rank verdict. Set {_SKIP_WEIGHT_QUANT_CHECK_ENV}=1 to bypass this "
"check in that situation -- note this is process-global, so it silences the check "
"on every rank, not only the one with the legitimately empty stage."
)
def quantize(
model: nn.Module,
config: dict[str, Any | QuantizeConfig],
forward_loop: ForwardLoop | None = None,
) -> nn.Module:
"""Quantizes and calibrates the model in-place.
This method performs replacement of modules with their quantized counterparts and
performs calibration as specified by ``quant_cfg``.
``forward_loop`` is used to forward data through the model and gather statistics for calibration.
If the model is already quantized, the provided ``config`` is applied to the existing
quantizers and calibration is run.
Args:
model: A pytorch model
config: A dictionary or an instance of
: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 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>`.
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
:class:`TensorQuantizer <modelopt.torch.quantization.nn.modules.tensor_quantizer.TensorQuantizer>`.
The quantizer attributes are defined by :class:`QuantizerAttributeConfig`. See
:class:`QuantizerAttributeConfig` for details on the quantizer attributes and their values.
An example ``config`` dictionary is given below:
.. code-block::python
config = {
"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
{"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
quantization formats. See :ref:`Quantization Configs <quantization-configs>` for more details on
``config`` dictionary.
forward_loop: A callable that forwards all calibration data through the model. This is used
to gather statistics for calibration. It should take model as the argument. It does not need
to return anything.
This argument is not required for weight-only quantization with the ``"max"``
algorithm.
Here are a few examples for correct ``forward_loop`` definitions:
Example 1:
.. code-block::
def forward_loop(model) -> None:
# iterate over the data loader and forward data through the model
for batch in data_loader:
model(batch)
Example 2:
.. code-block::
def forward_loop(model) -> float:
# evaluate the model on the task
return evaluate(model, task, ....)
Example 3:
.. code-block::
def forward_loop(model) -> None:
# run evaluation pipeline
evaluator.model = model
evaluator.evaluate()
.. note::
Calibration does not require forwarding the entire dataset through the model.
Please subsample the dataset or reduce the number of batches if needed.
Returns: A pytorch model which has been quantized and calibrated.
"""
quantize_config = QuantizeConfig(**dict(config))
if not is_quantized(model):
model = apply_mode(model, mode=[("quantize", dict(config))], registry=QuantizeModeRegistry)
else:
# Already quantized, so lets apply the quant_cfg from the config
set_quantizer_by_cfg(model, quantize_config.quant_cfg)
# Fail before calibration rather than after exporting an unquantized checkpoint.
_check_weight_quantization_took_effect(model, quantize_config)
return calibrate(model, config.get("algorithm"), forward_loop=forward_loop)
# TODO: create a config interface for auto_quantize and expose setting
# quant_grouping_rules and score_module_rules as part of the config.
# This will allow users to customize the grouping and scoring rules for their models.
# This way wecan limit the granularity of quantization search. For example,
# - limit the quantization format search to decoder block level (instead of each linear layer level)
# - Same format for all self attention layers of a model etc.
_AUTO_QUANTIZE_SUPPORTED_ALGORITHMS = {
None,
"max",
"mse",
"local_hessian",
"smoothquant",
"awq_lite",
"awq_full",
"awq_clip",
}
def _process_quantization_formats(formats, custom_name_prefix):
"""Resolve search formats and preserve explicitly supplied display names."""
processed = []
for index, item in enumerate(formats):
if item is None:
continue
if isinstance(item, tuple):
if len(item) != 2:
raise ValueError("Named quantization formats must be (config, name) pairs.")
quant_cfg, name = item
if not isinstance(name, str) or not name:
raise ValueError("Quantization format names must be non-empty strings.")
else:
quant_cfg = item
name = QuantRecipe.get_auto_name_for_config(quant_cfg)
if name is None:
name = f"{custom_name_prefix}_{index}"
warnings.warn(
"Received custom quantization formats for search, auto_quantize results may "
f"not be optimal. This config will be displayed as {name}"
)
processed.append((quant_cfg, name))
return processed
def _auto_quantize_kv_cache(
model: nn.Module,
constraints: dict[str, Any],
quantization_formats: Sequence[dict[str, Any] | str | tuple[dict[str, Any], str]],
*,
data_loader: Iterable | None,
forward_step: Callable[[nn.Module, Any], Any | torch.Tensor] | None,
loss_func: Callable[[Any, Any], torch.Tensor] | None,
forward_backward_step: Callable[[nn.Module, Any], Any] | None,
disabled_layers: list[str] | str | None,
num_calib_steps: int,
num_score_steps: int,
verbose: bool,
method: str | None,
checkpoint: str | None,
module_search_spaces: list[dict[str, Any]] | None,
fixed_quantization_config: dict[str, Any] | str | None,
):
"""Run the KV-cache-specific AutoQuantize validation and search lifecycle."""
if (
torch.distributed.is_available()
and torch.distributed.is_initialized()
and torch.distributed.get_world_size() > 1
):
raise RuntimeError(
"KV-cache AutoQuantize is single-process only; distributed scoring, selection, "
"and checkpoint writes are not synchronized."
)
if is_quantized(model):
raise NotImplementedError(
"KV-cache AutoQuantize requires an unquantized model; composing it after GEMM "
"PTQ or AutoQuantize is not supported yet."
)
if method not in (None, "kl_div"):
raise ValueError("cost_model='kv_cache' requires method='kl_div'.")
if fixed_quantization_config is not None or module_search_spaces:
raise ValueError(
"KV-cache AutoQuantize does not support fixed_quantization_config or "
"module_search_spaces."
)
if loss_func is not None or forward_backward_step is not None:
raise ValueError(
"KV-cache AutoQuantize uses forward KL and does not accept loss_func or "
"forward_backward_step."
)
if not quantization_formats:
raise ValueError("cost_model='kv_cache' requires a non-empty quantization_formats list.")
if data_loader is None or forward_step is None:
raise ValueError("data_loader and forward_step must be provided for KV-cache AutoQuantize.")
processed_kv_formats: list[tuple[dict[str, Any], str | None]] = []
for candidate in quantization_formats:
if isinstance(candidate, tuple):
raw_config, name = candidate
if not isinstance(name, str) or not name:
raise ValueError("KV-cache AutoQuantize candidate names must be non-empty strings.")
elif isinstance(candidate, str):
if not hasattr(mtq, candidate):
raise ValueError(f"Unknown KV-cache quantization format: {candidate!r}.")
raw_config, name = getattr(mtq, candidate), candidate
elif isinstance(candidate, dict):
raw_config = candidate
name = QuantRecipe.get_auto_name_for_config(candidate)
else:
raise TypeError(
"KV-cache quantization formats must be config dictionaries, preset names, "
"or (config, name) tuples."
)
if not isinstance(raw_config, dict):
raise TypeError("KV-cache AutoQuantize formats must resolve to config dictionaries.")
processed_kv_formats.append((raw_config, name))
_validate_kv_cache_search_inputs(
constraints,
processed_kv_formats,
num_calib_steps,
num_score_steps,
)
model = apply_mode(model, mode="auto_quantize", registry=QuantizeModeRegistry)
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
searcher = AutoQuantizeKVSearcher()
searcher.search(
model,
cast("ConstraintsDict", constraints),
config={
"quantization_formats": processed_kv_formats,
"data_loader": data_loader,
"forward_step": forward_step,
"num_calib_steps": num_calib_steps,
"num_score_steps": num_score_steps,
"disabled_layers": disabled_layers,
"verbose": verbose,
"checkpoint": checkpoint,
},
)
return model, searcher.state_dict()
def auto_quantize(
model: nn.Module,
constraints: dict[str, Any] | None = None,
quantization_formats: Sequence[dict[str, Any] | str | tuple[dict[str, Any], str]] | None = None,
data_loader: Iterable | None = None,
forward_step: Callable[[nn.Module, Any], Any | torch.Tensor] | None = None,
loss_func: Callable[[Any, Any], torch.Tensor] | None = None,
forward_backward_step: Callable[[nn.Module, Any], Any] | None = None,
disabled_layers: list[str] | str | None = None,
num_calib_steps: int = 512,
num_score_steps: int = 128,
verbose: bool = False,
method: str | None = None,
checkpoint: str | None = None,
module_search_spaces: list[dict[str, Any]] | None = None,
fixed_quantization_config: dict[str, Any] | str | None = None,
):
r"""Perform optimal per-layer quantization by searching for the best quantization formats per-layer.
``auto_quantize`` uses sensitivity scores to rank the per-layer quantization formats and search
for the best quantization formats per-layer. The sensitivity score can be computed using gradient-based
methods (default) or KL divergence loss, controlled by the ``method`` parameter.
Internally this API runs two main phases:
#. Calibrate the quantized model exactly like :func:`quantize` would.
#. Estimate per-layer sensitivity scores to decide which format to keep.
The sensitivity scoring phase typically dominates the runtime of ``auto_quantize``, so decreasing the number of
samples used for scoring (see ``num_score_steps``) is the recommended way for improving overall auto_quantize time
with minimal accuracy impact.
Args:
model: A pytorch model with quantizer modules.
constraints: Constraints for the search. ``effective_bits`` specifies the effective number
of bits for the quantized model and defaults to 4.8. ``cost_model`` selects the metric
used for the effective-bits constraint and supports ``"weight"`` (default),
``"active_moe"``, and ``"kv_cache"``. The KV-cache cost model dispatches to isolated
forward-KL scoring over paired K/V formats; BF16/no-quant is its scoring reference but
is never solver-selectable. Additional cost-model parameters are provided through the
nested ``cost`` dict.
Here is an example for valid ``effective_bits`` argument:
.. code-block:: python
# For the default AutoQuantize effective-bits target
constraints = {"effective_bits": 4.8}
# For active-MoE accounting where 2 of 8 routed experts are active per token
constraints = {
"effective_bits": 4.8,
"cost_model": "active_moe",
"cost": {
"active_moe_expert_ratio": 0.25,
"excluded_module_name_patterns": ["*visual*", "*vision_tower*", "*mtp*"],
},
}
# For paired K/V-cache formats with exact K/V storage accounting
constraints = {"effective_bits": 5.4, "cost_model": "kv_cache"}
quantization_formats: A sequence of quantization format config dictionaries or string names to search for.
Each config dictionary should be valid as a ``config`` argument in
:meth:`quantize <modelopt.torch.quantization.model_quant.quantize>`.
The supported quantization format names are as listed by :attr:`modelopt.torch.quantization.config.choices`.
Internally we always add "do not quantize" as a choice. Therefore, it is possible that a layer is
not quantized by any of the quantization formats.
Custom quantization formats can also be defined and used as a quantization format. This is a experimental
feature and the results may not be optimal. Here is an example:
.. code-block:: python
INT8_CUSTOM_QUANT_CFG = {
"quant_cfg": [
{"quantizer_name": "*weight_quantizer", "cfg": {"num_bits": 8, "axis": 0}},
{
"quantizer_name": "*input_quantizer",
"cfg": {"num_bits": 8, "axis": None},
},
],
"algorithm": "smoothquant",
}
mtq.auto_quantize(
model,
constraints,
quantization_formats=["INT4_AWQ_CFG", INT8_CUSTOM_QUANT_CFG],
)
Internally we always add "do not quantize" as a choice. Therefore, it is possible that a layer is
not quantized by any of the quantization formats.
.. note::
The quantization formats will be applied on a per-layer match basis. The global model level name
based quantizer attribute setting will be ignored. For example, in ``FP8_DEFAULT_CFG`` quantizer
configuration the key ``"*lm_head*": {"enable": False}`` disables quantization for the ``lm_head``
layer. However in ``auto_quantize``, the quantization format for the ``lm_head`` layer will be searched.
This is because the key ``"*lm_head*"`` sets the quantizer attributes based on the global model level
name, not per-layer basis. The keys ``"*input_quantizer"``, ``"*weight_quantizer"`` etc. in
``FP8_DEFAULT_CFG`` match on a per-layer basis - hence the corresponding quantizers
will be set as specified.
Here is an example `quantization_formats` argument:
.. code-block:: python
# A valid `quantization_formats` argument
# This will search for the best per-layer quantization from FP8, W4A8_AWQ_BETA_CFG or No quantization
quantization_formats = [mtq.FP8_DEFAULT_CFG, mtq.W4A8_AWQ_BETA_CFG]
data_loader: An iterator that yields data that is to be used for calibrating quantized layers and estimating
``auto_quantize`` scores.
forward_step: A callable that takes the model and a batch of data from ``data_loader`` as input, forwards
the data through the model and returns the model output.
This is a required argument.
Here is an example for a valid ``forward_step``:
.. code-block:: python
# Takes the model and a batch of data as input and returns the model output
def forward_step(model, batch) -> torch.Tensor:
output = model(batch)
return output
loss_func: (Optional) A callable that takes the model output and the batch of data as input and computes the
loss. The model output is the output given by ``forward_step``. `.backward()` will be called on the loss.
Here is an example for a valid ``loss_func``:
.. code-block:: python
# Takes the model output and a batch of data as input and returns the loss
def loss_func(output, batch) -> torch.Tensor:
...
return loss
# loss should be a scalar tensor such that loss.backward() can be called
loss = loss_func(output, batch)
loss.backward()
If this argument is not provided, ``forward_backward_step`` should be provided.
forward_backward_step: (Optional) A callable that takes batch of data from ``data_loader``, forwards it
through the model, computes the loss and runs backward on the loss.
Here is an example for a valid ``forward_backward_step`` argument:
.. code-block:: python
# Takes the model and a batch of data as input and runs forward and backward pass
def forward_backward_step(model, batch) -> None:
output = model(batch)
loss = my_loss_func(output, batch)
run_custom_backward(loss)
If this argument is not provided, ``loss_func`` should be provided.
disabled_layers: (Optional) One or a list of wildcard strings to disable quantization for the layers. Example:
.. code-block:: python
disabled_layers = "*lm_head*"
disabled_layers = ["*lm_head*", "*mlp*"]
num_calib_steps: Number of batches to use for calibrating each candidate quantization format. Suggested value
is 512.
num_score_steps: Number of batches to use for estimating ``auto_quantize`` scores. Suggested value is 128.
A higher value could increase the time taken for performing ``auto_quantize``; reducing it speeds up the
sensitivity score estimation phase and typically affects accuracy less than lowering ``num_calib_steps``.
verbose: If True, prints the search progress/intermediate results.
method: Method to use for estimating sensitivity loss. Higher loss indicates greater sensitivity
to quantization. Options are ``"gradient"`` (default; uses gradient-based loss estimation,
linear programming search, and requires ``loss_func`` or ``forward_backward_step``) and
``"kl_div"`` (uses KL divergence between unquantized and quantized outputs, relies on
threshold-based binary search, and only requires ``forward_step`` returning logits).
checkpoint: (Optional) Path to checkpoint file for saving/restoring auto_quantize search state.
If the checkpoint file exists, the search state will be restored from it, skipping the
expensive score estimation step.
module_search_spaces: Optional module-specific candidate overrides. Each entry contains
``module_name_patterns`` (one or more glob patterns), ``quantization_formats``, and
optional ``allow_no_quant`` (default True). A matching entry replaces the global
candidate set for that runtime-grouped decision. Setting ``allow_no_quant=False``
keeps BF16/no-quant as an internal sensitivity and cost baseline but prevents the
solver from selecting it. A single candidate with ``allow_no_quant=False`` fixes the
matching module group to that format while retaining its cost in the effective-bits
constraint.
fixed_quantization_config: Optional normal PTQ config applied to modules not matched by
``module_search_spaces``. When provided, ``quantization_formats`` must be omitted and
at least one explicit module search space is required. The fixed baseline remains
active while searched modules are scored, is calibrated only with its own algorithm,
and remains part of the effective-bits numerator and denominator. This is one
integrated AutoQuantize operation, not staged PTQ followed by AutoQuantize.
Returns: A tuple (model, state_dict) where ``model`` is the searched and quantized model and
``state_dict`` contains the history and detailed stats of the search procedure.
.. note::
``auto_quantize`` groups certain layers and restricts the quantization formats for them to be same. For example,
Q, K, V linear layers belonging to the same transformer layer will have the same quantization format.
This is to ensure compatibility with TensorRT-LLM which fuses these three linear layers into a single linear
layer.
Grouping rules are defined in :attr:`quant_grouping_rules
<.algorithms.AutoQuantizeSearcher.quant_grouping_rules>`.
Each rule can be either a regex pattern or a callable function.
- **Regex patterns**: The first captured group (e.g.,
``pattern.match(name).group(1)``) determines the group key.
Layers with the same group key share the same quantization format.
- **Functions**: Should take a module name and return a group key
(or ``None`` if the rule doesn't apply).
Example regex rule: ``r"^(.*?)\.(q_proj|k_proj|v_proj)$"`` groups the
`q_proj`, `k_proj`, `v_proj` layers belonging to the same transformer layer.
You can customize the rules as needed:
.. code-block:: python
from modelopt.torch.quantization.algorithms import AutoQuantizeSearcher
# Add a regex rule to group layers in the same `mlp` module
AutoQuantizeSearcher.quant_grouping_rules.append(r"^(.*?)\.mlp")
# Or add a function rule for custom logic
AutoQuantizeSearcher.quant_grouping_rules.append(
lambda name: name.rsplit(".", 1)[0] if "expert" in name else None
)
# Perform `auto_quantize`
model, state_dict = auto_quantize(model, ...)
.. note::
The ``auto_quantize`` API and algorithm is experimental and subject to change. ``auto_quantize`` searched models
might not be readily deployable to TensorRT-LLM yet.
"""
if quantization_formats is not None:
if isinstance(quantization_formats, str) or not isinstance(quantization_formats, Sequence):
raise TypeError("`quantization_formats` must be a sequence of formats.")
quantization_formats = list(quantization_formats)
is_kv_search = constraints is not None and constraints.get("cost_model") == COST_MODEL_KV_CACHE
if is_kv_search:
assert constraints is not None
assert quantization_formats is not None
return _auto_quantize_kv_cache(
model,
constraints,
quantization_formats,
data_loader=data_loader,
forward_step=forward_step,
loss_func=loss_func,
forward_backward_step=forward_backward_step,
disabled_layers=disabled_layers,
num_calib_steps=num_calib_steps,
num_score_steps=num_score_steps,
verbose=verbose,
method=method,
checkpoint=checkpoint,
module_search_spaces=module_search_spaces,
fixed_quantization_config=fixed_quantization_config,
)
method = method or "gradient"
if fixed_quantization_config is None and quantization_formats is None:
quantization_formats = [mtq.NVFP4_AWQ_LITE_CFG, mtq.FP8_DEFAULT_CFG]
elif fixed_quantization_config is not None and quantization_formats is None:
quantization_formats = []
if fixed_quantization_config is not None and quantization_formats:
raise ValueError(
"`fixed_quantization_config` cannot be combined with global "
"`quantization_formats`; put every searched module in `module_search_spaces`."
)
if fixed_quantization_config is None and not quantization_formats:
raise ValueError("`quantization_formats` must be a non-empty list.")
processed_quantization_formats = _process_quantization_formats(quantization_formats, "CUSTOM")
if quantization_formats and not processed_quantization_formats:
raise ValueError("`quantization_formats` must contain at least one non-None format.")
processed_module_search_spaces = []
for idx, search_space in enumerate(module_search_spaces or []):
if not isinstance(search_space, dict):
raise TypeError("Each module_search_spaces entry must be a dict.")
unknown_keys = set(search_space) - {
"module_name_patterns",
"quantization_formats",
"allow_no_quant",
}
if unknown_keys:
raise ValueError(f"Unsupported module_search_spaces keys: {sorted(unknown_keys)}")
patterns = search_space.get("module_name_patterns")
if isinstance(patterns, str):
patterns = [patterns]
if (
not isinstance(patterns, list)
or not patterns
or not all(isinstance(pattern, str) for pattern in patterns)
):
raise ValueError(
"module_search_spaces.module_name_patterns must be a non-empty string list."
)
raw_formats = search_space.get("quantization_formats")
if not isinstance(raw_formats, list):
raise TypeError("module_search_spaces.quantization_formats must be a list.")
if not raw_formats:
raise ValueError("module_search_spaces.quantization_formats must be a non-empty list.")
formats = _process_quantization_formats(raw_formats, f"CUSTOM_MODULE_{idx}")
if not formats:
raise ValueError(
"module_search_spaces.quantization_formats must contain at least one non-None format."
)
allow_no_quant = search_space.get("allow_no_quant", True)
if not isinstance(allow_no_quant, bool):
raise TypeError("module_search_spaces.allow_no_quant must be a bool.")
processed_module_search_spaces.append(
{
"module_name_patterns": patterns,
"quantization_formats": formats,
"allow_no_quant": allow_no_quant,
}
)
if fixed_quantization_config is not None and not processed_module_search_spaces:
raise ValueError(
"`fixed_quantization_config` requires at least one explicit "
"`module_search_spaces` entry."
)
processed_fixed_quantization_configs = _process_quantization_formats(
[fixed_quantization_config] if fixed_quantization_config is not None else [],
"FIXED",
)
assert len(processed_fixed_quantization_configs) <= 1
processed_fixed_quantization_config = (
processed_fixed_quantization_configs[0] if processed_fixed_quantization_configs else None
)
all_processed_formats = [
*processed_quantization_formats,
*processed_fixed_quantization_configs,
*(
quant_format
for search_space in processed_module_search_spaces
for quant_format in search_space["quantization_formats"]
),
]
for quant_cfg, name in all_processed_formats:
algo = QuantRecipe(quant_cfg, name=name).config.algorithm
algo_method = algo["method"] if isinstance(algo, dict) else algo
if algo_method not in _AUTO_QUANTIZE_SUPPORTED_ALGORITHMS:
raise ValueError(
f"Algorithm '{algo_method}' in '{name}' is not supported by auto_quantize yet. "
"Please run auto_quantize with 'max' or 'mse' calibration and use "
"get_auto_quantize_config() to obtain a config for mtq.quantize()."
)
# Select the appropriate searcher based on method
if method == "gradient":
searcher = AutoQuantizeGradientSearcher()
elif method == "kl_div":
searcher = AutoQuantizeKLDivSearcher()
else:
raise ValueError(f"Invalid method: {method}. Valid options are 'gradient' or 'kl_div'.")
model = apply_mode(
model,
mode="auto_quantize",
registry=QuantizeModeRegistry,
)
search_config = {
"quantization_formats": processed_quantization_formats,
"fixed_quantization_config": processed_fixed_quantization_config,
"module_search_spaces": processed_module_search_spaces,
"data_loader": data_loader,
"forward_step": forward_step,
"loss_func": loss_func,
"forward_backward_step": forward_backward_step,
"num_calib_steps": num_calib_steps,
"num_score_steps": num_score_steps,
"disabled_layers": disabled_layers,
"verbose": verbose,
"checkpoint": checkpoint,
}
# Disable all quantizers; AutoQuantize will enable the needed ones
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
if processed_fixed_quantization_config is not None:
fixed_cfg, fixed_name = processed_fixed_quantization_config
fixed_recipe = QuantRecipe(fixed_cfg, name=fixed_name)
set_quantizer_by_cfg(model, fixed_recipe.config.quant_cfg)
search_constraints = cast("ConstraintsDict", constraints or {})
searcher.search(model, search_constraints, config=search_config)
return model, searcher.state_dict()
def get_auto_quantize_config(search_state, constraints=None, verbose=False):
"""Build a flat quant config from auto_quantize search_state.
Re-solves for ``constraints`` if provided, otherwise uses the stored best recipe.
Args:
search_state: The state dict returned by :func:`auto_quantize`.
constraints: Optional dict, e.g. ``{"effective_bits": 5.5}``, to re-solve for a
different target without re-running calibration or scoring.
verbose: If True, prints the per-layer recipe assignments.
Returns:
A config dict suitable for :func:`quantize`.
Example:
.. code-block:: python
model, search_state = mtq.auto_quantize(model, ...)
# Re-solve for a different effective_bits target (cheap, no GPU needed)
config = mtq.get_auto_quantize_config(search_state, {"effective_bits": 5.5})
# Or use the original result
config = mtq.get_auto_quantize_config(search_state)
# Reuse on the same model (e.g. run a longer calibration pass)
model = mtq.quantize(model, config, forward_loop=calibrate_loop)
# Or apply the same/customized config on a fresh model instance
# fresh_model = load_model(...)
# fresh_model = mtq.quantize(fresh_model, config, forward_loop=calibrate_loop)
"""
if search_state.get("cost_model") == COST_MODEL_KV_CACHE:
return get_kv_cache_auto_quantize_config(search_state, constraints, verbose=verbose)
return _get_auto_quantize_config(search_state, constraints, verbose=verbose)
def disable_quantizer(model: nn.Module, wildcard_or_filter_func: str | Callable):
"""Disable quantizer by wildcard or filter function."""
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_attributes_partial(model, wildcard_or_filter_func, {"enable": True})
@atomic_print
def print_quant_summary(model: nn.Module, output_dir: str | None = None):
"""Print summary of all quantizer modules in the model."""
lines = [
f"{name:80} {mod}"
for name, mod in model.named_modules()
if isinstance(mod, TensorQuantizer)
]
lines.append(f"{len(lines)} TensorQuantizers found in model")
if output_dir:
os.makedirs(output_dir, exist_ok=True)
path = os.path.join(output_dir, ".quant_summary.txt")
with open(path, "w", encoding="utf-8") as f:
f.write("\n".join(lines) + "\n")
print(f"\033[1mQuant summary saved to {path}\033[0m")
else:
print("\n".join(lines))
def fold_weight(model: nn.Module, keep_attrs: bool = False):
"""Fold weight quantizer for fast evaluation.
Any weight-quantizer rotation is folded into the weights and disabled so subsequent
forwards do not re-rotate the already-folded weights.
"""
for name, module in model.named_modules():
if isinstance(module, QuantModule):
module.fold_weight(keep_attrs)
@contextmanager
def temporarily_fold_weights(
model: nn.Module,
snapshot_device: torch.device | str | None = None,
):
"""Temporarily fold fake-quant weights for a frozen inference region.
Each :class:`QuantModule` performs its normal module-specific ``fold_weight`` operation.
Fake-quant weights affected by quantization, pre-quant scaling, or rotation and all fake-quant
runtime states are restored on exit, including after an exception. Weights are restored in
place so optimizer and distributed references remain valid.
This context is intended for repeated no-gradient forwards with no optimizer step, such as
log-probability recomputation over several microbatches. It retains calibration attributes
while folded; a retained weight ``pre_quant_scale`` is inactive inside the context because its
value is already baked into the temporary weight. Sharing a weight or weight quantizer across
multiple :class:`QuantModule` instances and using :class:`SequentialQuantizer` are not
supported.
Example::
with mtq.temporarily_fold_weights(model, snapshot_device="cpu"):
outputs = model(inputs)
Args:
model: Quantized model whose weights will be temporarily folded.
snapshot_device: Device used to store parameter snapshots. ``None`` keeps each snapshot
on the parameter's device; ``"cpu"`` avoids the additional accelerator memory.
"""
fold_pairs = []
for module in model.modules():
if not isinstance(module, QuantModule):
continue
for weight, quantizer in module.iter_weights_for_calibration():
if not isinstance(weight, torch.Tensor):
continue
if isinstance(quantizer, SequentialQuantizer):
raise NotImplementedError(
"temporarily_fold_weights does not support SequentialQuantizer"
)
if not isinstance(quantizer, TensorQuantizer) or not quantizer.fake_quant:
continue
local_weight = weight.to_local() if hasattr(weight, "to_local") else weight
fold_pairs.append((local_weight, quantizer))
weight_snapshots = {}
for weight, quantizer in fold_pairs:
if (
quantizer.is_enabled
or quantizer.pre_quant_scale is not None
or quantizer.rotate_is_enabled
):
weight_id = id(weight)
if weight_id not in weight_snapshots:
weight_snapshots[weight_id] = (
weight,
weight.detach().clone()
if snapshot_device is None
else weight.detach().to(snapshot_device, copy=True),
)
with preserve_quantizer_attributes_context(model):
try:
fold_weight(model, keep_attrs=True)
yield
finally:
with torch.no_grad():
for weight, snapshot in weight_snapshots.values():
weight.copy_(snapshot)
@torch.no_grad()
def compute_quantization_mse(
model: nn.Module,
forward_loop: ForwardLoop,
wildcards: str | Callable | list[str | Callable] = "*",
) -> dict[str, float]:
"""Compute the mean-squared quantization error for selected quantizers.
Runs ``forward_loop`` through the model while recording, for every matching
:class:`TensorQuantizer`, the MSE between the original float tensor and
its fake-quantized (Q→DQ) counterpart. Values are averaged over all
calibration batches.
Args:
model: A quantized model (output of :func:`quantize`).
forward_loop: Callable that takes ``model`` and runs data through it.
wildcards: One or more fnmatch glob patterns (or callable filters)
matched against :class:`TensorQuantizer` module names in
``model.named_modules()``. Follows the same convention as
``quant_cfg`` wildcard keys. Defaults to ``"*"`` (all quantizers).
Returns:
A dict mapping each matched quantizer's fully-qualified name to its
mean MSE (float). Quantizers that are disabled or not in fake-quant
mode are skipped and absent from the output.
Example::
mse = mtq.compute_quantization_mse(
model,
forward_loop,
wildcards=["*k_bmm_quantizer", "*v_bmm_quantizer"],
)
for name, err in sorted(mse.items()):
print(f"{name}: {err:.4e}")
"""
if not isinstance(wildcards, list):
wildcards = [wildcards]
def _matches(name: str) -> bool:
return any(fnmatch.fnmatch(name, w) if isinstance(w, str) else w(name) for w in wildcards)
accumulators: dict[str, dict] = {} # name -> {"sum": float, "count": int}
hooks = []
for name, module in model.named_modules():
if not isinstance(module, TensorQuantizer):
continue
if not _matches(name):
continue
if not (module._if_quant and module._fake_quant) or module._disabled:
continue
accumulators[name] = {"sum": 0.0, "count": 0}
def _make_hook(acc):
def hook(mod, inp, out):
original = inp[0].detach().float()
quantized = out.detach().float()
acc["sum"] += torch.mean((original - quantized) ** 2).item()
acc["count"] += 1
return hook
hooks.append(module.register_forward_hook(_make_hook(accumulators[name])))
try:
forward_loop(model)
finally:
for h in hooks:
h.remove()
return {
name: acc["sum"] / acc["count"] for name, acc in accumulators.items() if acc["count"] > 0
}