mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### 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>
1107 lines
50 KiB
Python
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
|
|
}
|