Files
Model-Optimizer/modelopt/torch/quantization/algorithms.py
T
JoshuaandCursor 63c4b660bd Add Aumann-Shapley sensitivity scoring method to auto_quantize (#2183)
[Paper](https://arxiv.org/abs/2607.12266) ·
[Overview](https://x.com/waterloo_intern/status/2076460984475263401) ·
[Implementation
thread](https://x.com/the_joshua_hill/status/2076427869388255635)

Depends on #2231, which provides the shared AutoQuantize
backward-scoring infrastructure. Until that PR merges, the focused PR B
diff is available
[here](https://github.com/joshua-hill/Model-Optimizer/compare/fix/autoquant-scoring-infrastructure...feat/aumann-shapley-autoquant).

### What does this PR do?

Type of change: new feature

This PR adds `method="aumann_shapley"` to `mtq.auto_quantize`. It is a
label-free scoring method that measures how each candidate quantization
format affects the model across the path from full precision to
quantized.

The new method is opt-in; existing `gradient` and `kl_div` behavior is
unchanged.

For each calibration batch, the method:

1. Runs the baseline model and saves its next-token distribution.
2. Measures each candidate format at a configurable number of points
along the quantization path.
3. Uses the KL-divergence gradients at those points to assign a damage
contribution to every runtime group and candidate format.
4. Measures the most aggressive candidate configuration once and uses
that value to calibrate the per-group damage model.

The resulting scores use the existing AutoQuantize linear-program
solver. The search can either:

- choose the least damaging configuration that meets an `effective_bits`
target; or
- choose the smallest configuration whose predicted damage stays below
`max_predicted_damage`.

The selected recipe records `predicted_damage` in mean per-token KL
units together with its validity and fit diagnostics.

### Public API

`auto_quantize` gains an optional `method_options` dictionary. For
`method="aumann_shapley"`, it accepts:

| Option | Default | Meaning |
|---|---:|---|
| `num_path_nodes` | `2` | Number of points used to average gradients
along the quantization path. |
| `damage_link` | `"coverage"` | How per-group scores combine.
`"coverage"` uses `damage = c * (1 - exp(-sum(b)))`; `"additive"` sums
the path contributions. |
| `max_predicted_damage` | `None` | Replaces the bit target with a
maximum predicted mean per-token KL. |

Method options are validated before the model is modified. Unknown
options and incompatible targets fail early.

### Usage

Select a configuration for a target effective bit width:

```python
import modelopt.torch.quantization as mtq

model, search_state = mtq.auto_quantize(
    model,
    constraints={"effective_bits": 4.8},
    quantization_formats=["NVFP4_DEFAULT_CFG", "FP8_DEFAULT_CFG"],
    data_loader=calib_loader,
    forward_step=lambda model, batch: model(**batch),
    method="aumann_shapley",
)

print(search_state["best"]["predicted_damage"])
print(search_state["best"]["predicted_damage_valid"])
```

Or let the search choose the bit width for a predicted-damage target:

```python
model, search_state = mtq.auto_quantize(
    model,
    constraints={},
    quantization_formats=["NVFP4_DEFAULT_CFG", "FP8_DEFAULT_CFG"],
    data_loader=calib_loader,
    forward_step=forward_step,
    method="aumann_shapley",
    method_options={"max_predicted_damage": 0.05},
)
```

### Implementation

- Reuses the candidate-replay and backward-scoring lifecycle introduced
in #2231.
- Reuses the existing AutoQuantize linear-program solver for both search
directions.
- Numerically integrates the coverage path when converting measured
contributions into per-group damage costs.
- Preserves deterministic runtime-group and candidate ordering.
- Keeps raw measurements, solver scores, and damage-model diagnostics
distinct in the search state.
- Rejects incompatible checkpoint resumes while allowing the same scores
to be re-solved for a new bit budget.
- Retains the shared MoE score-module rules so routed experts are scored
at their enclosing block.

Recipe integration will follow separately.

### Testing

Focused tests:

```text
pytest -q \
  tests/unit/torch/quantization/test_autoquant.py::test_backward_scoring_session_restores_partial_setup \
  tests/unit/torch/quantization/test_autoquant_shapley.py
```

Result: **51 passed**.

The tests cover:

- end-to-end scoring and configuration generation;
- agreement between path contributions and measured quantization damage;
- exact allocation checks against exhaustive search;
- effective-bits and predicted-damage search modes;
- checkpoint resume and offline re-solving;
- custom formats and heterogeneous candidate ladders;
- distributed reductions and nested MoE score modules;
- reused score modules and model-specific backward support;
- non-finite measurements and invalid-fit reporting; and
- input validation before model conversion.

End-to-end checks with NVFP4 and FP8 candidates at a 6.0-bit target:

- `Qwen/Qwen2.5-0.5B-Instruct` reaches 5.998 effective bits, with the
summed path contributions reproducing 98% of the directly measured
lowest-precision KL.
- `Qwen/Qwen3-30B-A3B` reaches 6.000 effective bits and reproduces 99%,
with all 48 MoE layers scored once at the sparse-MoE block rather than
per expert.

### Production use

We use this method in production for NVFP4 checkpoints of Kimi-K3,
MiniMax-M3, and GLM-5.2.

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

- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow the contributing guidance?: ✅ No copied code
and no new dependencies.
- Did you write the necessary tests?: ✅
- Did you update `CHANGELOG.rst`?: ✅
- Are the commits signed and signed off?: ✅


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

- **New Features**
  - Added label-free Aumann–Shapley scoring for automatic quantization.
- Added configurable path sampling, damage modeling, effective-bit
targets, and predicted-damage bounds.
- Added temporary weight-folding support with automatic state
restoration.
- Added method-specific search options, checkpoint resumption, and
distributed scoring.

- **Bug Fixes**
- Improved cleanup and restoration of quantizer state, gradients, hooks,
and forward behavior after scoring or failures.
- Added validation and clearer handling for unsupported configurations
and invalid measurements.

- **Tests**
- Expanded coverage for scoring, solver behavior, distributed execution,
checkpointing, and custom quantization formats.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Joshua Hill <joshua.hill@baseten.co>
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-23 20:51:55 -07:00

2303 lines
98 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.
"""Module for advanced quantization algorithms."""
import copy
import fnmatch
import functools
import gc
import types
import warnings
from abc import ABC, abstractmethod
from collections import defaultdict
from collections.abc import Callable, Sequence
from contextlib import ExitStack, nullcontext
from typing import Any
import regex as re
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
from modelopt.torch.opt.conversion import ModeloptStateManager
from modelopt.torch.opt.hparam import CustomHPType, Hparam, HPType
from modelopt.torch.opt.searcher import LPS, BaseSearcher, SearchConfig, SearchStateDict
from modelopt.torch.opt.utils import get_hparam, named_hparams
from modelopt.torch.utils import create_param_grad_clear_hook, print_rank_0, report_memory
from modelopt.torch.utils.distributed import DistributedProcessGroup, ParallelState, is_master
from . import config as mtq_config
from . import model_calib
from ._auto_quantize_cost import (
ACTIVE_MOE_EXPERT_RATIO_KEY,
AUTO_QUANTIZE_CONSTRAINT_KEYS,
COST_MODEL_ACTIVE_MOE,
COST_MODEL_WEIGHT,
_get_module_weight_numel,
get_auto_quantize_cost_model,
normalize_auto_quantize_constraints,
)
from .config import QuantizeConfig, QuantizerAttributeConfig, QuantizerCfgEntry
from .conversion import set_quantizer_by_cfg
from .nn import QuantLinearConvBase, QuantModule, SequentialQuantizer, TensorQuantizer
from .utils import is_quantized_linear
def _is_hf_quant_fused_experts_module(module: nn.Module) -> bool:
"""Return True for a converted HF fused-MoE-experts quantization wrapper."""
# Late import avoids a circular import: the HF plugin registers AutoQuantize
# support from this module at import time.
try:
from .plugins.huggingface import _is_quant_fused_experts_module
except ImportError:
return False
return _is_quant_fused_experts_module(module)
# Quantizer attribute names that participate in AutoQuantize snapshot/restore.
_STD_QUANTIZER_ATTRS = ("input_quantizer", "weight_quantizer", "output_quantizer")
_FUSED_EXPERTS_QUANTIZER_ATTRS = (
"gate_up_proj_input_quantizer",
"gate_up_proj_weight_quantizers",
"down_proj_input_quantizer",
"down_proj_weight_quantizers",
)
_FUSED_EXPERTS_REPLAY_QUANTIZER_ATTRS = (
"gate_up_proj_input_quantizer",
"gate_up_proj_weight_quantizer",
"down_proj_input_quantizer",
"down_proj_weight_quantizer",
)
_NON_GATED_FUSED_EXPERTS_REPLAY_QUANTIZER_ATTRS = (
"up_proj_input_quantizer",
"up_proj_weight_quantizer",
"down_proj_input_quantizer",
"down_proj_weight_quantizer",
)
def _get_replay_quantizer_attr(attr_name: str) -> str:
"""Return the quantizer name used by config matching/replay."""
if attr_name.endswith("_quantizers"):
return attr_name.removesuffix("s")
return attr_name
def _get_quantizer_attrs(module: nn.Module) -> tuple[str, ...]:
"""Return the quantizer attribute names that AutoQuantize must snapshot/restore.
For fused MoE experts, this returns the four plural quantizer attrs (two
shared input quantizers + two ``ModuleList`` of per-expert weight quantizers).
For standard Linear-derived QuantModules, returns the canonical trio.
"""
if _is_hf_quant_fused_experts_module(module):
try:
from .plugins.huggingface import _get_fused_experts_quantizer_attr_names
except ImportError:
return _FUSED_EXPERTS_QUANTIZER_ATTRS
return _get_fused_experts_quantizer_attr_names(module)
return _STD_QUANTIZER_ATTRS
def _make_fresh_quantizer_for_attr(module: nn.Module, attr_name: str) -> nn.Module:
"""Return a fresh, default quantizer object suitable to overwrite ``module.<attr_name>``.
For ModuleList attrs (per-expert quantizers on fused-experts modules), the
returned ModuleList preserves the original list length so per-expert
enumeration stays consistent across recipes.
"""
current = getattr(module, attr_name, None)
if isinstance(current, nn.ModuleList):
return nn.ModuleList(TensorQuantizer() for _ in range(len(current)))
return TensorQuantizer()
def _iter_tensor_quantizers(module: nn.Module):
if isinstance(module, TensorQuantizer):
yield module
elif isinstance(module, nn.ModuleList | SequentialQuantizer):
for child in module:
yield from _iter_tensor_quantizers(child)
def _tensor_quantizer_format_signature(quantizer: TensorQuantizer) -> tuple:
"""Return the numerical-format fields relevant to runtime fusion compatibility."""
return (
quantizer.is_enabled,
quantizer.num_bits,
getattr(quantizer, "_effective_bits", None),
quantizer.axis,
repr(quantizer.block_sizes),
getattr(quantizer, "_dynamic", False),
quantizer.fake_quant,
quantizer.backend,
repr(quantizer.backend_extra_args),
)
def _fixed_module_format_signature(module: nn.Module) -> tuple:
return tuple(
(
_get_replay_quantizer_attr(attr_name),
tuple(
_tensor_quantizer_format_signature(quantizer)
for quantizer in _iter_tensor_quantizers(getattr(module, attr_name))
),
)
for attr_name in _get_quantizer_attrs(module)
)
def _fixed_module_weight_compression(
module: nn.Module, effective_bits_override: float | None = None
) -> float:
weight_quantizers = []
for attr_name in _get_quantizer_attrs(module):
if "weight_quantizer" not in attr_name:
continue
weight_quantizers.extend(_iter_tensor_quantizers(getattr(module, attr_name)))
if not weight_quantizers or all(not quantizer.is_enabled for quantizer in weight_quantizers):
return 1.0
if any(not quantizer.is_enabled for quantizer in weight_quantizers):
raise ValueError(
"The fixed quantize baseline enables only some weight quantizers within one "
"quantizable module. Move that module into an explicit AutoQuantize "
"module_search_spaces entry."
)
if effective_bits_override is not None:
return effective_bits_override / 16
compressions = []
for quantizer in weight_quantizers:
effective_bits = getattr(quantizer, "_effective_bits", None)
num_bits = quantizer.num_bits
if effective_bits is not None:
compressions.append(effective_bits / 16)
elif isinstance(num_bits, tuple):
compressions.append((sum(num_bits) + 1) / 16)
elif isinstance(num_bits, int):
compressions.append(num_bits / 16)
else:
raise ValueError(f"Cannot infer AutoQuantize cost from num_bits={num_bits!r}.")
if any(abs(value - compressions[0]) > 1e-12 for value in compressions[1:]):
raise ValueError(
"The fixed quantize baseline assigns different weight formats within one quantizable "
"module. Move that module into an explicit AutoQuantize module_search_spaces entry."
)
return compressions[0]
def estimate_quant_compression(quant_cfg: QuantizeConfig) -> float:
"""Estimate the compression ratio of a quantization configuration.
Effective bits per element resolve in priority order: (1) recipe-level
``quant_cfg.effective_bits``; (2) per-entry ``cfg.effective_bits`` (library default,
e.g. NVFP4 = 4.5); (3) the ``num_bits`` heuristic (``num_bits / 16`` for ints,
``(E + M + 1) / 16`` for FP tuples). Per-entry values are aggregated via ``min``, which
still under-counts activation cost for mixed weight+activation formats.
Args:
quant_cfg: The quantization configuration to estimate compression for.
Returns:
float: The estimated compression ratio (0.0 to 1.0).
"""
if quant_cfg.effective_bits is not None:
return quant_cfg.effective_bits / 16.0
def estimate_quant_compression_for_quantizer(quantizer_attr_cfg):
if isinstance(quantizer_attr_cfg, list):
if not quantizer_attr_cfg:
return 1.0
return min(estimate_quant_compression_for_quantizer(q) for q in quantizer_attr_cfg)
if isinstance(quantizer_attr_cfg, dict):
# Handle raw quantizer cfg dicts (e.g. {"num_bits": (4, 3), "axis": None})
if not quantizer_attr_cfg.get("enable", True):
return 1.0
effective_bits = quantizer_attr_cfg.get("effective_bits")
if effective_bits is not None:
return effective_bits / 16
num_bits = quantizer_attr_cfg.get("num_bits")
if num_bits is None:
return 1.0
if isinstance(num_bits, tuple):
return (sum(num_bits) + 1) / 16
elif isinstance(num_bits, int):
return num_bits / 16
else:
raise ValueError(f"Unknown quantization config {num_bits}")
if isinstance(quantizer_attr_cfg, QuantizerAttributeConfig):
if not quantizer_attr_cfg.enable:
return 1.0
if quantizer_attr_cfg.effective_bits is not None:
return quantizer_attr_cfg.effective_bits / 16
if not hasattr(quantizer_attr_cfg, "num_bits"):
return 1.0
if isinstance(quantizer_attr_cfg.num_bits, tuple):
return (sum(quantizer_attr_cfg.num_bits) + 1) / 16
elif isinstance(quantizer_attr_cfg.num_bits, int):
return quantizer_attr_cfg.num_bits / 16
else:
raise ValueError(f"Unknown quantization config {quantizer_attr_cfg.num_bits}")
raise ValueError(f"Unknown type {type(quantizer_attr_cfg)}, {quantizer_attr_cfg}")
cfgs = []
for e in quant_cfg.quant_cfg:
if e.get("enable", True) is False:
continue
c = e.get("cfg")
if c is not None:
cfgs.append(c)
return estimate_quant_compression_for_quantizer(cfgs) if cfgs else 1.0
@functools.cache
def _no_quant_signature() -> str:
"""Return the canonical signature used to identify the no-quant recipe."""
return QuantRecipe(quant_cfg=None).checkpoint_signature
class QuantRecipe(CustomHPType):
"""A subclass of QuantizeConfig enabling auto_quantize specific configurations.
Args:
quant_cfg: str or dict or None. dict is used for custom quantization formats.
name: name for custom quantization formats. Only used if quantization format is a custom
format not available in :mod:`modelopt.torch.quantization.config`.
"""
def __init__(self, quant_cfg: str | dict[str, Any] | None = None, name: str | None = None):
"""Initialize the QuantRecipe with the quantization configuration."""
name = self.get_auto_name_for_config(quant_cfg) or name
if quant_cfg is None:
quant_cfg = {"quant_cfg": [{"quantizer_name": "*", "enable": False}]}
elif isinstance(quant_cfg, str):
assert hasattr(mtq_config, quant_cfg), f"Unknown quantization format {quant_cfg}"
quant_cfg = getattr(mtq_config, quant_cfg)
else:
assert name is not None, "name must be provided for custom quantization formats"
self.config = mtq_config.QuantizeConfig(**quant_cfg) # type: ignore [arg-type]
# Disable KV Cache quantization
# Currently KV Cache quantization is enabled for some quantization formats and disabled for others
# This breaks the monotonicity of the quantization formats in terms of weight compression Vs accuracy
self.config.quant_cfg.append(
QuantizerCfgEntry(quantizer_name="*output_quantizer", enable=False)
)
self.compression = estimate_quant_compression(self.config)
self._str_repr: str = f"{name}(effective-bits: {self.compression * 16})"
self._config_signature = self.config.model_dump_json()
@property
def checkpoint_signature(self) -> str:
"""Return the canonical identity used for ordering and checkpoint validation."""
return getattr(self, "_config_signature", self.config.model_dump_json())
@property
def is_no_quant(self) -> bool:
"""Whether this recipe leaves the module unquantized."""
return self.checkpoint_signature == _no_quant_signature()
@staticmethod
def get_auto_name_for_config(quant_cfg: str | dict[str, Any] | None) -> str | None:
"""Get a name for the quantization configuration."""
if quant_cfg is None:
return "NONE"
if isinstance(quant_cfg, str):
return quant_cfg
for quant_cfg_name in mtq_config.choices:
if quant_cfg == getattr(mtq_config, quant_cfg_name):
return quant_cfg_name
return None
@property
def num_bits(self) -> int:
"""Get the number of bits for the quantization format."""
return int(self.compression * 16)
def __str__(self) -> str:
return self._str_repr
def __repr__(self) -> str:
return self._str_repr
def __lt__(self, other: "QuantRecipe"):
# Callers treat the last choice as the unquantized end of the format ladder.
return (self.compression, self.is_no_quant, self.checkpoint_signature) < (
other.compression,
other.is_no_quant,
other.checkpoint_signature,
)
def __eq__(self, other: object):
return (
isinstance(other, QuantRecipe)
and self.checkpoint_signature == other.checkpoint_signature
)
def __hash__(self) -> int:
return hash(self.checkpoint_signature)
@staticmethod
def disable_folding_pqs_to_weights():
"""Disable the folding of pre_quant_scale to weights."""
model_calib._ENABLE_FOLDING_PQS_TO_WEIGHTS = False
@staticmethod
def fold_pqs_to_weights(model):
"""Fold the pre_quant_scale in weight_quantizers to weights."""
model_calib._ENABLE_FOLDING_PQS_TO_WEIGHTS = True
for name, module in model.named_modules():
if is_quantized_linear(module):
with SequentialQuantizer.convert_to_single_quantizer(module):
if module.weight_quantizer.pre_quant_scale is not None:
weight_pqs = module.weight_quantizer.pre_quant_scale
delattr(module.weight_quantizer, "_pre_quant_scale")
model_calib._apply_weight_pre_quant_scale(module, weight_pqs)
class QuantRecipeHparam(Hparam):
"""An Hparam for quantization recipes.
See :class:`Hparam <modelopt.torch.opt.hparam.Hparam>` for more details. In addition, this Hparam also:
* Keeps a link to its ``quant_modules`` and ``score_modules`` and sets the quantizers for the
``quant_modules`` based on the active recipe.
* Provides ``get_score()`` and ``get_cost()`` methods to evaluate recipes.
* Registers itself with each ``score_module`` via the ``_hparams_for_scoring`` attribute.
"""
def __init__(
self,
choices: Sequence[QuantRecipe] | None = None,
quant_modules: list[nn.Module] | None = None,
score_modules: list[nn.Module] | None = None,
name: str | None = None,
quant_module_names: list[str] | None = None,
cost_weight: float = 1.0,
allow_no_quant: bool = True,
fixed_recipe: QuantRecipe | None = None,
) -> None:
"""Initializes Hparam with internal scoring choices and solver selectability."""
candidate_choices = sorted(set(choices or []))
if fixed_recipe is not None:
assert candidate_choices == [fixed_recipe]
assert not allow_no_quant
# A one-format rule with no-quant disallowed is genuinely fixed: keep that
# format active while other groups are scored. Multi-format rules retain an
# internal no-quant reference for sensitivity estimation, then filter it out
# before LP selection.
choices = (
candidate_choices
if not allow_no_quant and len(candidate_choices) == 1
else sorted({*candidate_choices, QuantRecipe(quant_cfg=None)})
)
super().__init__(choices, original=choices[0])
self.name = name
self.quant_module_names = quant_module_names or []
self.quant_module_replay_attrs = {
name: tuple(_get_replay_quantizer_attr(attr) for attr in _get_quantizer_attrs(module))
for module, name in zip(quant_modules or [], self.quant_module_names)
}
assert cost_weight >= 0.0, "cost_weight must be non-negative."
self.cost_weight = cost_weight
self.allow_no_quant = allow_no_quant
self.is_fixed = fixed_recipe is not None
# Module hashes depend on object identity, so sets can produce different orders per rank.
self.quant_modules = list(dict.fromkeys(quant_modules or []))
self.score_modules = list(dict.fromkeys(score_modules or self.quant_modules))
self._warned_parallel_state_fallbacks: set[nn.Module] = set()
fixed_quantizers = (
{
module: {
attr_name: getattr(module, attr_name)
for attr_name in _get_quantizer_attrs(module)
}
for module in self.quant_modules
}
if fixed_recipe is not None
else {}
)
# This is a hack; We dont want to make the input_quantizer, weight_quantizer, output_quantizer
# a dynamic attribute for backward compatibility with the model_calib.py
# TODO: Make input_quantizer, weight_quantizer, output_quantizer a dynamic attribute and get rid of this hack
# NOTE: For fused-experts modules, the relevant attrs are plural
# (``*_input_quantizer`` + ``*_weight_quantizers`` ModuleList) — see
# ``_get_quantizer_attrs``. Both layouts share the same snapshot dict
# shape so ``active.setter`` swaps the right child modules.
no_quant_recipe = QuantRecipe(quant_cfg=None)
calibration_recipes = sorted({*self.choices, no_quant_recipe})
self._all_quantizer_choices = {quant_recipe: {} for quant_recipe in calibration_recipes}
quant_recipe: QuantRecipe
for quant_recipe in calibration_recipes:
for quant_module in self.quant_modules:
attr_names = _get_quantizer_attrs(quant_module)
if quant_recipe == fixed_recipe:
self._all_quantizer_choices[quant_recipe][quant_module] = fixed_quantizers[
quant_module
]
continue
for attr_name in attr_names:
setattr(
quant_module,
attr_name,
_make_fresh_quantizer_for_attr(quant_module, attr_name),
)
set_quantizer_by_cfg(quant_module, quant_recipe.config.quant_cfg)
self._all_quantizer_choices[quant_recipe][quant_module] = {
attr_name: getattr(quant_module, attr_name) for attr_name in attr_names
}
self.active = self.original
# Importance dict is keyed by score_module (where the score is computed)
self._importance_dict = {
quant_recipe: dict.fromkeys(self.score_modules) for quant_recipe in self.choices
}
# Registration order follows the rank-stable runtime-group construction order.
for score_module in self.score_modules:
if not hasattr(score_module, "_hparams_for_scoring"):
score_module._hparams_for_scoring = []
score_module._hparams_for_scoring.append(self)
@property
def active(self) -> HPType:
"""Return the currently active value."""
return self._active
@active.setter
def active(self, val: HPType | None):
"""Set the active value with a sanity check for choices and dynamic hparams."""
val = self.original if val is None else val
assert isinstance(val, QuantRecipe)
assert val in self._choices, f"val = {val}, choices = {self.choices}"
if self.is_configurable:
self._active = val
else:
assert self._active == val
self._apply_quantizer_choice(val)
def _apply_quantizer_choice(self, recipe: QuantRecipe) -> None:
for nn_module, quantizer_choices in self._all_quantizer_choices[recipe].items():
for quantizer_attr_name, quantizer in quantizer_choices.items():
setattr(nn_module, quantizer_attr_name, quantizer)
def set_calibration_recipe(self, recipe: QuantRecipe) -> None:
"""Enable ``recipe`` for this calibration pass or isolate the group."""
calibration_recipe = recipe if recipe in self.choices else QuantRecipe(quant_cfg=None)
self._apply_quantizer_choice(calibration_recipe)
def restore_active_quantizers(self) -> None:
"""Restore quantizer objects corresponding to the solver-visible active recipe."""
active = self.active
assert isinstance(active, QuantRecipe)
self._apply_quantizer_choice(active)
@property
def solver_choices(self) -> list[QuantRecipe]:
"""Return choices exposed to the LP after removing an internal no-quant baseline."""
no_quant_recipe = QuantRecipe(quant_cfg=None)
recipes: list[QuantRecipe] = []
for recipe in self.choices:
assert isinstance(recipe, QuantRecipe)
if self.is_fixed or self.allow_no_quant or recipe != no_quant_recipe:
recipes.append(recipe)
return recipes
@property
def importance(self) -> dict:
"""Raises an error since this is not a useful abstraction for AutoQuantize."""
raise NotImplementedError
def get_score(self, recipe: QuantRecipe) -> float:
"""Get the score for a given recipe."""
total_score = 0
for score_module in self.score_modules:
importance = self._importance_dict[recipe][score_module]
if importance is None:
continue
parallel_state = getattr(score_module, "parallel_state", None)
if parallel_state is None:
# TODO: Prefer parallel_state owned by the score module; this temporary fallback
# inherits the first quantized child's state and assumes all grouped quant modules
# share the same parallel groups.
parallel_state_source = next(
(
(module, state)
for module in self.quant_modules
if (state := getattr(module, "parallel_state", None)) is not None
),
None,
)
if parallel_state_source is not None:
quant_module, parallel_state = parallel_state_source
if (
torch.distributed.is_initialized()
and score_module not in self._warned_parallel_state_fallbacks
):
warnings.warn(
"Distributed training is initialized but no parallel_state is set for "
f"score module {type(score_module)}. Using parallel_state from its first "
f"quantized child {type(quant_module)}. All grouped quant modules must "
"share the same parallel groups."
)
self._warned_parallel_state_fallbacks.add(score_module)
if parallel_state is None:
total_score += importance.cpu().item()
continue
importance = importance.cpu()
importance = DistributedProcessGroup.get_dist_syncd_obj(
importance,
[
parallel_state.tensor_parallel_group,
parallel_state.data_parallel_group,
parallel_state.expert_model_parallel_group,
],
sum,
)
total_score += importance.item()
return total_score
def get_cost(self, recipe: QuantRecipe, cost_weight: float | None = None) -> float:
"""Get the cost for a given recipe.
The cost is the total weight size of the quantizable modules multiplied by
the compression ratio of the recipe.
"""
cost_weight = self.cost_weight if cost_weight is None else cost_weight
cost = 0
for quant_module in self.quant_modules:
weight_size = (
_AutoQuantizeBaseSearcher._get_total_weight_size([quant_module]) * cost_weight
)
parallel_state = getattr(quant_module, "parallel_state", None)
if parallel_state is None:
cost += weight_size * recipe.compression
continue
weight_size = DistributedProcessGroup.get_dist_syncd_obj(
weight_size,
[
parallel_state.tensor_parallel_group,
parallel_state.expert_model_parallel_group,
],
sum,
)
# Across data parallel groups, the weight size is the same for all the ranks.
weight_size = DistributedProcessGroup.get_dist_syncd_obj(
weight_size,
[parallel_state.data_parallel_group],
lambda a: a[0],
)
cost += weight_size * recipe.compression
return cost
@property
def attrs(self) -> list[str]:
"""Return the attributes of the hparam for repr."""
return ["name", "cost_weight", "allow_no_quant", "is_fixed", *super().attrs]
_LINEAR_ATTN_QKVZ_RE = re.compile(r"^(.*?\.linear_attn)\.(?:in_proj_qkv|in_proj_z)$")
_LINEAR_ATTN_BA_RE = re.compile(r"^(.*?\.linear_attn)\.(?:in_proj_a|in_proj_b)$")
def _linear_attn_qkvz_group_key(_model, name: str) -> str | None:
m = _LINEAR_ATTN_QKVZ_RE.match(name)
return f"{m.group(1)}/qkvz" if m else None
def _linear_attn_ba_group_key(_model, name: str) -> str | None:
m = _LINEAR_ATTN_BA_RE.match(name)
return f"{m.group(1)}/ba" if m else None
def _module_search_space_signature(module_search_spaces) -> tuple:
"""Return a checkpoint-stable description of module-specific candidate spaces."""
return tuple(
(
tuple(search_space["module_name_patterns"]),
tuple(sorted(recipe.checkpoint_signature for recipe in search_space["quant_recipes"])),
search_space["allow_no_quant"],
)
for search_space in module_search_spaces
)
def _quantization_formats_signature(quant_recipes) -> tuple[str, ...]:
"""Return a checkpoint-stable description of the global candidate formats."""
return tuple(sorted(recipe.checkpoint_signature for recipe in quant_recipes))
class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
"""Base searcher for AutoQuantize algorithm."""
# This searcher finds optimal per-layer quantization by searching across quantization formats
# for each quantizable module (quant module). Optionally, quant grouping rules can restrict
# certain modules to share the same format. Sensitivity scores are computed from perturbations
# at score modules. See AutoQuantizeGradientSearcher for detailed documentation.
candidate_stats: dict[str, dict[str, Any]]
best: dict[str, Any]
quantizer_states: dict
method_name: str | None = None
method_config_keys: frozenset[str] = frozenset()
quant_grouping_rules = [
r"^(.*?)\.(q_proj|k_proj|v_proj)$", # q_proj, k_proj, v_proj for llama like models
# gate_proj, up_proj, down_proj for Qwen3 like MoE models
r"^(.*?\.mlp\.experts)\.\d+\.(gate_proj|up_proj|down_proj)$",
r"^(.*?\.mixer\.experts)\.\d+\.(up_proj|down_proj)$", # NemotronH MoE experts
# NemotronH MoE experts in MCore naming (linear_fc1=gate+up fused, linear_fc2=down)
r"^(.*?\.mlp\.experts\.local_experts)\.\d+\.(linear_fc1|linear_fc2)$",
r"^(.*?)\.(gate_proj|up_proj)$", # gate_proj, up_proj for llama like models
r"^(.*?)\.(\d+\.(w1|w2|w3))$", # mixtral experts
r"^(.*?)\.((w1_linear|w2_linear|w3_linear)\.\d+)$", # dbrx experts
# Qwen3.5/3.6 hybrid linear_attn: vLLM fuses (in_proj_qkv, in_proj_z)
# into ``in_proj_qkvz`` and (in_proj_a, in_proj_b) into ``in_proj_ba`` and
# requires fused shards to share quant_algo. Two callables (not one
# regex) so qkv+z and a+b produce DIFFERENT group keys; each pair
# stays with its own fusion partner.
_linear_attn_qkvz_group_key,
_linear_attn_ba_group_key,
]
score_module_rules = []
@property
def default_search_config(self):
"""Get the default config for the searcher."""
return {
"quantization_formats": ["NVFP4_DEFAULT_CFG", "FP8_DEFAULT_CFG"],
"fixed_quantization_config": None,
"module_search_spaces": [],
"data_loader": None,
"num_calib_steps": 512,
"num_score_steps": 128,
"deployment": None,
"disabled_layers": None,
"verbose": is_master(),
"checkpoint": None,
"cost_model": COST_MODEL_WEIGHT,
"cost": {},
"active_moe_expert_ratio": None,
}
@property
def default_state_dict(self) -> SearchStateDict:
"""Get the default state dict for AutoQuantize."""
return {
"method": self.method_name,
"cost_model": "weight",
"cost": {},
"active_moe_expert_ratio": None,
"cost_denominator": None,
"quantization_formats_signature": None,
"fixed_quantization_config_signature": None,
"fixed_quantization_config": None,
"module_search_space_signature": None,
"resolved_search_setup_signature": None,
"disabled_layers": None,
"candidate_stats": defaultdict(dict),
"quantizer_states": {},
"best": {"recipe": {}, "constraints": {}, "score": float("inf"), "is_satisfied": False},
}
def sanitize_search_config(self, config: SearchConfig | None) -> SearchConfig:
"""Sanitize the search config dict."""
config = config or {}
config = super().sanitize_search_config(config)
assert config["data_loader"] is not None, (
"`data_loader` must be provided for `auto_quantize`."
)
assert config["forward_step"] is not None, (
"`forward_step` must be provided for `auto_quantize`."
)
return config
def validate_search_input(self, constraints, config) -> None:
"""Validate method-specific inputs before quantizing the model."""
def load_search_checkpoint(self) -> bool:
return super().load_search_checkpoint(strict=False)
@staticmethod
def _is_auto_quantize_module(module):
if (is_quantized_linear(module) or isinstance(module, QuantLinearConvBase)) and isinstance(
module, QuantModule
):
return True
# Fused MoE experts: a single ``QuantModule`` that owns N per-expert
# weight quantizers in an ``nn.ModuleList`` plus shared input quantizers.
# All N experts in a layer share one search dimension (one recipe per
# fused module).
return _is_hf_quant_fused_experts_module(module) and isinstance(module, QuantModule)
@staticmethod
def _get_search_recipes(quantization_formats):
return sorted(
{
QuantRecipe(quant_cfg=q[0], name=q[1])
if isinstance(q, tuple)
else QuantRecipe(quant_cfg=q)
for q in quantization_formats
}
)
def _apply_quant_group_rule(self, name: str, rule) -> str | None:
"""Apply a single quant_group_rule to a module name.
Args:
name: Module name
rule: Either a regex pattern string or a callable that returns a unique key;
If callable, it should take the model and the name as input and return the unique key
Returns:
The group key if the rule matches, None otherwise
"""
if callable(rule):
return rule(self.model, name)
else:
# Regex pattern
pattern = re.compile(rule)
match = pattern.match(name)
if match:
return match.group(1)
return None
def _apply_score_group_rule(self, name: str, rule) -> str | None:
"""Apply a single score_group_rule to a module name.
Args:
name: Module name
rule: Either a regex pattern string or a callable that returns the score module name.
If callable, it should take the model and the name as input and return the score module name
Returns:
The score module name if the rule matches, None otherwise
"""
if callable(rule):
return rule(self.model, name)
else:
# Regex pattern - return the matched name or full match
pattern = re.compile(rule)
match = pattern.match(name)
if match:
# For score rules, return the full match or first group
return match.group(0) if match.lastindex is None else match.group(1)
return None
def _get_score_module_from_name(
self, model: nn.Module, score_module_name: str, quant_module: nn.Module
) -> nn.Module:
"""Get the actual score module object from its name.
Args:
model: The model containing all modules
score_module_name: The name of the score module to retrieve
quant_module: The quantized module for which the score is estimated
Returns:
The score module object, or the quantized module itself if the score module is not found
"""
try:
score_module = model.get_submodule(score_module_name)
return score_module
except AttributeError:
warnings.warn(
f"Score module '{score_module_name}' not found. Score will estimated from the quantized module itself."
)
return quant_module
def _normalize_module_search_spaces(self, module_search_spaces):
"""Convert processed API search spaces to QuantRecipe-based rules."""
return [
{
"module_name_patterns": tuple(search_space["module_name_patterns"]),
"quant_recipes": self._get_search_recipes(search_space["quantization_formats"]),
"allow_no_quant": search_space["allow_no_quant"],
}
for search_space in module_search_spaces
]
@staticmethod
def _match_module_search_space(quant_module_names, module_search_spaces):
"""Return the unique rule that fully covers a runtime-grouped decision."""
matched_search_spaces = []
for search_space in module_search_spaces:
matches = [
any(
fnmatch.fnmatch(module_name, pattern)
for pattern in search_space["module_name_patterns"]
)
for module_name in quant_module_names
]
if not any(matches):
continue
if not all(matches):
raise ValueError(
"A module_search_spaces rule partially matches runtime-grouped modules "
f"{quant_module_names}. Update its module_name_patterns so the rule covers "
"the entire group or none of it."
)
matched_search_spaces.append(search_space)
if len(matched_search_spaces) > 1:
raise ValueError(
"Multiple module_search_spaces rules match runtime-grouped modules "
f"{quant_module_names}. Make the module_name_patterns disjoint."
)
return matched_search_spaces[0] if matched_search_spaces else None
@staticmethod
def _resolve_fixed_group_recipe(quant_modules, quant_module_names, fixed_recipe):
"""Resolve a full-model PTQ baseline to one runtime-group-compatible fixed choice."""
format_signatures = [_fixed_module_format_signature(module) for module in quant_modules]
if any(signature != format_signatures[0] for signature in format_signatures[1:]):
raise ValueError(
"The fixed quantize baseline assigns incompatible formats to runtime-grouped "
f"modules {quant_module_names}. Move the entire group into one explicit "
"AutoQuantize module_search_spaces entry."
)
compressions = [
_fixed_module_weight_compression(module, fixed_recipe.config.effective_bits)
for module in quant_modules
]
if any(abs(value - compressions[0]) > 1e-12 for value in compressions[1:]):
raise ValueError(
"The fixed quantize baseline assigns different weight costs to runtime-grouped "
f"modules {quant_module_names}. Move the entire group into one explicit "
"AutoQuantize module_search_spaces entry."
)
compression = compressions[0]
if abs(compression - 1.0) <= 1e-12:
return QuantRecipe(quant_cfg=None)
if abs(compression - fixed_recipe.compression) > 1e-12:
raise ValueError(
"The fixed quantize baseline resolves some unmatched modules to a different "
"numerical format than its effective_bits cost. Use one uniform PTQ format as "
"the baseline and put format-specific modules in AutoQuantize "
"module_search_spaces."
)
return fixed_recipe
def insert_hparams_after_merge_rules(
self,
model,
quant_recipes,
disabled_layers=None,
module_search_spaces=None,
fixed_recipe=None,
):
"""Restrict the search space using the merge rules and insert the hparams for the model."""
# TRTLLM fuses linear layers such as q_proj, k_proj, v_proj into same layer
# Hence we need to restrict the search space so that all these layers share the same recipe
# Lets group the modules based on the rules and insert the same hparam for all the modules in the group
if disabled_layers is None:
disabled_layers = []
elif isinstance(disabled_layers, str):
disabled_layers = [disabled_layers]
# Map from group key to list of (quant_module, name, disabled, score_module)
search_map: dict[str, list[tuple[nn.Module, str, bool, nn.Module]]] = {}
for name, module in model.named_modules():
if not self._is_auto_quantize_module(module):
continue
# Skip layers that match disabled_layers patterns
disabled = False
for pattern in disabled_layers:
if fnmatch.fnmatch(name, pattern):
disabled = True
break
# Apply quant_grouping_rules to determine the group key
group_key = name # Default: each module in its own group
for rule in self.quant_grouping_rules:
result = self._apply_quant_group_rule(name, rule)
if result is not None:
group_key = result
# We support only one rule for matching per module
break
# Apply score_module_rules to determine the score module name, then get the actual module
score_module_name = name # Default: score from same module
for rule in self.score_module_rules:
result = self._apply_score_group_rule(name, rule)
if result is not None:
score_module_name = result
# We support only one rule for matching per module
break
# Get the actual score module object immediately
score_module = self._get_score_module_from_name(model, score_module_name, module)
if group_key not in search_map:
search_map[group_key] = [(module, name, disabled, score_module)]
else:
search_map[group_key].append((module, name, disabled, score_module))
for group_key, module_info_list in search_map.items():
quant_modules = [module for module, _, _, _ in module_info_list]
disabled = any(disabled for _, _, disabled, _ in module_info_list)
score_modules = [score_module for _, _, _, score_module in module_info_list]
quant_module_names = [name for _, name, _, _ in module_info_list]
cost_weight = self._cost_model.module_cost_weight(
quant_module_names, self.config["cost"]
)
search_space = self._match_module_search_space(
quant_module_names, module_search_spaces or []
)
if disabled:
_quant_recipes = None
allow_no_quant = True
resolved_fixed_recipe = None
elif search_space is not None:
_quant_recipes = search_space["quant_recipes"]
allow_no_quant = search_space["allow_no_quant"]
resolved_fixed_recipe = None
elif fixed_recipe is not None:
resolved_fixed_recipe = self._resolve_fixed_group_recipe(
quant_modules, quant_module_names, fixed_recipe
)
_quant_recipes = [resolved_fixed_recipe]
allow_no_quant = False
else:
_quant_recipes = quant_recipes
allow_no_quant = True
resolved_fixed_recipe = None
hparam = QuantRecipeHparam(
_quant_recipes,
quant_modules=quant_modules,
score_modules=score_modules,
name=str(group_key),
quant_module_names=quant_module_names,
cost_weight=cost_weight,
allow_no_quant=allow_no_quant,
fixed_recipe=resolved_fixed_recipe,
)
for module in quant_modules:
module._register_hparam("quant_recipe", hparam)
def _get_formatted_weight_compression_constraint(self):
effective_bits = self.constraints["effective_bits"]
assert effective_bits > 0 and effective_bits <= 16, (
"effective_bits should be between 0 and 16."
)
weight_compression = self.constraints["effective_bits"] / 16.0
return weight_compression
def _verify_constraint(self, search_recipes):
assert self.constraints["effective_bits"] >= search_recipes[0].num_bits, (
f"The effective_bits {self.constraints['effective_bits']} constraint cannot be lower than the "
f"num_bits of most aggressive quantization format for this search which is "
f"{search_recipes[0]} whose num_bits = {search_recipes[0].num_bits}."
)
def _resolved_search_setup_signature(self, quant_recipe_hparams) -> tuple:
"""Fingerprint the runtime groups, choices, scoring boundaries, and cost weights."""
module_names = {id(module): name for name, module in self.model.named_modules()}
signature = []
for hparam in quant_recipe_hparams:
replay_attrs = tuple(
(module_name, tuple(attrs))
for module_name, attrs in sorted(hparam.quant_module_replay_attrs.items())
)
score_module_names = tuple(
sorted(
module_names.get(id(module), type(module).__qualname__)
for module in hparam.score_modules
)
)
signature.append(
(
hparam.name,
tuple(sorted(hparam.quant_module_names)),
replay_attrs,
score_module_names,
tuple(sorted(recipe.checkpoint_signature for recipe in hparam.solver_choices)),
hparam.allow_no_quant,
hparam.is_fixed,
float(hparam.cost_weight),
)
)
return tuple(sorted(signature, key=repr))
def _verify_resolved_constraint(self, quant_recipe_hparams) -> None:
"""Fail before calibration when resolved per-group choices cannot meet the budget."""
no_quant_recipe = QuantRecipe(quant_cfg=None)
uncompressed_cost = sum(hparam.get_cost(no_quant_recipe) for hparam in quant_recipe_hparams)
if uncompressed_cost <= 0:
raise ValueError(
"AutoQuantize cost denominator is zero after applying the resolved cost "
"constraints. Include at least one quantizable module in the cost model."
)
minimum_cost = sum(
min(hparam.get_cost(recipe) for recipe in hparam.solver_choices)
for hparam in quant_recipe_hparams
)
target_cost = uncompressed_cost * self._get_formatted_weight_compression_constraint()
tolerance = uncompressed_cost * 1e-12
if minimum_cost > target_cost + tolerance:
minimum_effective_bits = minimum_cost / uncompressed_cost * 16
raise ValueError(
f"The effective_bits target {self.constraints['effective_bits']} is infeasible "
"for the resolved module search spaces. The minimum achievable effective bits "
f"is {minimum_effective_bits:.4f}."
)
@abstractmethod
def estimate_sensitivity_scores(self) -> None:
"""Estimate sensitivity scores and track them with Hparam."""
def initialize_candidate_stats(self):
"""Initialize the candidate stats for the model."""
no_quant_recipe = QuantRecipe(quant_cfg=None)
for name, hparam in named_hparams(self.model, unique=True):
if not isinstance(hparam, QuantRecipeHparam):
continue
formats, raw_scores, scores, costs = [], [], [], []
prev_score = float("inf")
for recipe in hparam.solver_choices:
formats.append(recipe)
score = hparam.get_score(recipe)
cost = hparam.get_cost(recipe)
raw_scores.append(score)
score = min(score, prev_score) # TODO: Should we get rid of this?
scores.append(score)
costs.append(cost)
prev_score = score
self.candidate_stats[name]["formats"] = formats
self.candidate_stats[name]["scores"] = scores
self.candidate_stats[name]["raw_scores"] = raw_scores
self.candidate_stats[name]["costs"] = costs
self.candidate_stats[name]["module_names"] = hparam.quant_module_names
self.candidate_stats[name]["quantizer_attrs"] = hparam.quant_module_replay_attrs
self.candidate_stats[name]["cost_weight"] = hparam.cost_weight
self.candidate_stats[name]["allow_no_quant"] = hparam.allow_no_quant
self.candidate_stats[name]["is_fixed"] = hparam.is_fixed
# Keep the no-quant cost as denominator metadata even when no-quant is not
# solver-selectable for this hparam. Fixed formats must remain in the cost model.
self.candidate_stats[name]["uncompressed_cost"] = hparam.get_cost(no_quant_recipe)
def _run_func(self, func, num_iters=1, desc=""):
for i, data in tqdm(
zip(range(num_iters), self.config["data_loader"]),
desc=desc,
total=num_iters,
):
func(self.model, data)
def before_search(self):
"""Prepare the model for search by calibrating the quantizers and collecting ``AutoQuantize`` score."""
# Import here to avoid circular import
from modelopt.torch.quantization.model_quant import calibrate
from .conversion import restore_quantizer_state, update_quantize_metadata
from .utils import get_quantizer_state_dict, set_quantizer_state_dict
super().before_search()
self.constraints = normalize_auto_quantize_constraints(self.model, self.constraints)
self.config["cost_model"] = self.constraints["cost_model"]
self.config["cost"] = self.constraints.get("cost", {})
self.config["active_moe_expert_ratio"] = self.config["cost"].get(
ACTIVE_MOE_EXPERT_RATIO_KEY
)
cost_model = get_auto_quantize_cost_model(self.config["cost_model"])
restored_method = getattr(self, "method", None)
if self.candidate_stats and restored_method not in (None, self.method_name):
raise ValueError(
f"Checkpoint method '{restored_method}' does not match current method "
f"'{self.method_name}'. Use a different checkpoint path."
)
restored_cost_model = getattr(self, "cost_model", "weight")
restored_active_moe_expert_ratio = getattr(self, "active_moe_expert_ratio", None)
if self.candidate_stats and (
restored_cost_model != self.config["cost_model"]
or restored_active_moe_expert_ratio != self.config["active_moe_expert_ratio"]
):
raise ValueError(
"Checkpoint AutoQuantize cost model does not match current search config: "
f"checkpoint=({restored_cost_model}, {restored_active_moe_expert_ratio}), "
f"current=({self.config['cost_model']}, {self.config['active_moe_expert_ratio']}). "
"Use a different checkpoint path."
)
self.method = self.method_name
self.cost_model = self.config["cost_model"]
self.cost = self.config["cost"]
self.active_moe_expert_ratio = self.config["active_moe_expert_ratio"]
self.disabled_layers = self.config["disabled_layers"]
self.cost_denominator = getattr(self, "cost_denominator", None)
module_search_spaces = self._normalize_module_search_spaces(
self.config["module_search_spaces"]
)
default_search_recipes = self._get_search_recipes(self.config["quantization_formats"])
fixed_search_recipes = self._get_search_recipes(
[self.config["fixed_quantization_config"]]
if self.config["fixed_quantization_config"] is not None
else []
)
assert len(fixed_search_recipes) <= 1
fixed_recipe = fixed_search_recipes[0] if fixed_search_recipes else None
quantization_formats_signature = _quantization_formats_signature(default_search_recipes)
fixed_quantization_config_signature = (
fixed_recipe.checkpoint_signature if fixed_recipe is not None else None
)
module_search_space_signature = _module_search_space_signature(module_search_spaces)
restored_quantization_formats_signature = getattr(
self, "quantization_formats_signature", None
)
restored_fixed_quantization_config_signature = getattr(
self, "fixed_quantization_config_signature", None
)
restored_module_search_space_signature = getattr(
self, "module_search_space_signature", None
)
has_restored_calibration_or_scores = bool(self.quantizer_states or self.candidate_stats)
if has_restored_calibration_or_scores and restored_quantization_formats_signature is None:
raise ValueError(
"Checkpoint does not record its quantization_formats signature and cannot be "
"safely reused. Use a different checkpoint path."
)
if (
has_restored_calibration_or_scores
and restored_quantization_formats_signature != quantization_formats_signature
):
raise ValueError(
"Checkpoint quantization_formats do not match the current search config. "
"Use a different checkpoint path."
)
if (
has_restored_calibration_or_scores
and restored_fixed_quantization_config_signature != fixed_quantization_config_signature
):
raise ValueError(
"Checkpoint fixed_quantization_config does not match the current search config. "
"Use a different checkpoint path."
)
if has_restored_calibration_or_scores and (
(restored_module_search_space_signature is None and module_search_space_signature)
or (
restored_module_search_space_signature is not None
and restored_module_search_space_signature != module_search_space_signature
)
):
raise ValueError(
"Checkpoint module_search_spaces do not match the current search config. "
"Use a different checkpoint path."
)
self.quantization_formats_signature = quantization_formats_signature
self.fixed_quantization_config_signature = fixed_quantization_config_signature
self.fixed_quantization_config = (
fixed_recipe.config.model_dump() if fixed_recipe is not None else None
)
self.module_search_space_signature = module_search_space_signature
search_recipes = sorted(
{
*default_search_recipes,
*fixed_search_recipes,
*(
recipe
for search_space in module_search_spaces
for recipe in search_space["quant_recipes"]
),
}
)
self._verify_constraint(search_recipes)
self._cost_model = cost_model
self.insert_hparams_after_merge_rules(
self.model,
default_search_recipes,
self.config["disabled_layers"],
module_search_spaces,
fixed_recipe,
)
quant_recipe_hparams = [
hparam
for _, hparam in named_hparams(self.model, unique=True)
if isinstance(hparam, QuantRecipeHparam)
]
resolved_search_setup_signature = self._resolved_search_setup_signature(
quant_recipe_hparams
)
restored_resolved_search_setup_signature = getattr(
self, "resolved_search_setup_signature", None
)
if has_restored_calibration_or_scores and restored_resolved_search_setup_signature is None:
raise ValueError(
"Checkpoint does not record its resolved search setup and cannot be safely "
"reused. Use a different checkpoint path."
)
if (
has_restored_calibration_or_scores
and restored_resolved_search_setup_signature != resolved_search_setup_signature
):
raise ValueError(
"Checkpoint resolved search setup does not match the current runtime groups, "
"allowed choices, scoring boundaries, or cost weights. Use a different "
"checkpoint path."
)
self.resolved_search_setup_signature = resolved_search_setup_signature
self._verify_resolved_constraint(quant_recipe_hparams)
QuantRecipe.disable_folding_pqs_to_weights()
# Iterate over the search recipes and calibrate the quantizers for each recipe
calibrated_new = False
try:
for recipe in search_recipes:
if recipe == QuantRecipe(quant_cfg=None): # No-quant format
continue
for hparam in quant_recipe_hparams:
hparam.set_calibration_recipe(recipe)
if recipe in self.quantizer_states:
saved = self.quantizer_states[recipe]
# config is unused by restore_quantizer_state
restore_quantizer_state(
self.model, QuantizeConfig(), {"quantizer_state": saved["metadata"]}
)
set_quantizer_state_dict(self.model, saved["state_dict"])
if self.config["verbose"]:
print_rank_0(f"AutoQuantize: Restored calibration for {recipe}")
continue
# Lets reduce the number of calibration steps for AWQ since it takes longer
num_calib_steps = (
self.config["num_calib_steps"]
if "awq" not in str(recipe.config.algorithm)
else max(1, self.config["num_calib_steps"] // 4)
)
def forward_loop(model):
self._run_func(
self.config["forward_step"],
num_iters=num_calib_steps,
desc=f"Calibrating for {recipe}",
)
calibrate(
self.model,
algorithm=recipe.config.algorithm,
forward_loop=forward_loop,
)
# Calibrate adds a new mode to the model. Since auto_quantize mixes the quantization recipes
# across layers, lets not save this new mode in the modelopt state.
# TODO: This is a hack. We need to create a mode for auto_quantize to handle this in a clean way.
ModeloptStateManager(self.model).state_dict().pop()
metadata: dict = {}
# config is unused by update_quantize_metadata
update_quantize_metadata(self.model, QuantizeConfig(), metadata)
self.quantizer_states[recipe] = {
"metadata": metadata["quantizer_state"],
"state_dict": get_quantizer_state_dict(self.model),
}
calibrated_new = True
finally:
for hparam in quant_recipe_hparams:
hparam.restore_active_quantizers()
if calibrated_new:
self.save_search_checkpoint(verbose=self.config["verbose"])
if self.candidate_stats:
if self.config["verbose"]:
print_rank_0("AutoQuantize: Restored from checkpoint, skipping scoring")
return
self.estimate_sensitivity_scores()
self.initialize_candidate_stats()
self.save_search_checkpoint(verbose=self.config["verbose"])
@staticmethod
def _print_recipe_summary(best_recipe, total_cost, total_weight_size, prefix="AutoQuantize"):
for name, recipe in best_recipe.items():
print_rank_0(f"{prefix} best recipe for {name.replace('.quant_recipe', '')}: {recipe}")
effective_bits = (total_cost / total_weight_size) * 16
print_rank_0(f"{prefix} effective bits: {effective_bits:.2f}")
return effective_bits
@staticmethod
def _get_total_weight_size(modules):
return sum(
_get_module_weight_numel(module)
if _AutoQuantizeBaseSearcher._is_auto_quantize_module(module)
else 0
for module in modules
)
@staticmethod
def _get_total_weight_size_from_candidate_stats(candidate_stats):
no_quant_recipe = QuantRecipe(quant_cfg=None)
total_weight_size = 0
for candidate_stat in candidate_stats.values():
if "uncompressed_cost" in candidate_stat:
total_weight_size += candidate_stat["uncompressed_cost"]
continue
no_quant_idx = candidate_stat["formats"].index(no_quant_recipe)
total_weight_size += candidate_stat["costs"][no_quant_idx]
return total_weight_size
def _get_constraints_for_search(self, max_weight_size, lower_bound=None):
constraints = {
"weight_size_after_compression": (
lower_bound * max_weight_size if lower_bound else lower_bound,
max_weight_size,
)
}
return constraints, "weight_size_after_compression"
def _get_search_lower_bounds(self):
cost_model = getattr(self, "cost_model", getattr(self, "config", {}).get("cost_model"))
if cost_model == COST_MODEL_ACTIVE_MOE:
return [0.99, 0.90, None]
return [None, 0.99, 0.90]
def _run_linear_program_search(self, max_weight_size, verbose=False):
"""Select recipes with the standard AutoQuantize linear program."""
for lower_bound in self._get_search_lower_bounds():
constraints, constraint_name = self._get_constraints_for_search(
max_weight_size, lower_bound
)
lps = LPS(
name="AutoQuantize",
constraints=constraints,
constraints_to_candidate_costs={
constraint_name: [
candidate_stat["costs"] for candidate_stat in self.candidate_stats.values()
]
},
candidate_scores=[
candidate_stat["scores"] for candidate_stat in self.candidate_stats.values()
],
objective_type="minimize",
verbose=verbose,
)
selections, self.status = lps()
if self.status == "Optimal":
break
is_satisfied = self.status == "Optimal"
if not is_satisfied:
warnings.warn(
"AutoQuantize FAILED to find a solution! The searched model might not meet all constraints. "
)
best_recipes = {}
for name, selected_idx in zip(self.candidate_stats, selections, strict=True):
best_recipes[name] = {
"format": self.candidate_stats[name]["formats"][selected_idx],
"costs": self.candidate_stats[name]["costs"][selected_idx],
"scores": self.candidate_stats[name]["scores"][selected_idx],
}
return best_recipes, is_satisfied
@abstractmethod
def run_search_with_stats(self, max_weight_size, verbose=False):
"""Run the search with stats to get the best recipe and whether the constraints are satisfied."""
def run_search(self):
"""Search for the best per-layer quantization configuration and return the best model and configuration."""
verbose = self.config["verbose"]
assert "effective_bits" in self.constraints and (
set(self.constraints) <= AUTO_QUANTIZE_CONSTRAINT_KEYS
), (
"`constraints` must contain 'effective_bits' and may contain 'cost_model' and 'cost'. "
f"Got {self.constraints.keys()}."
)
compression = self._get_formatted_weight_compression_constraint()
assert self.candidate_stats, (
"candidate_stats must be populated by before_search() before run_search()"
)
total_weight_size = self._get_total_weight_size_from_candidate_stats(self.candidate_stats)
self.cost_denominator = total_weight_size
max_weight_size = total_weight_size * compression
if verbose:
print_rank_0(
"AutoQuantize cost model: "
f"{self.config['cost_model']}"
+ (
f" (active_moe_expert_ratio={self.config['active_moe_expert_ratio']})"
if self.config["cost_model"] == COST_MODEL_ACTIVE_MOE
else ""
)
)
# Run the search with stats to get the best recipe and whether the constraints are satisfied
best_recipe_info, is_satisfied = self.run_search_with_stats(max_weight_size, verbose)
self.best["is_satisfied"] = is_satisfied
best_recipe = {}
best_constraints, best_scores = 0, 0
for name, best_hparam_recipe_info in best_recipe_info.items():
# Solvers could give different solutions for the same layer across DP/TP/EP groups even though
# the scores and costs are the same. Lets make sure the same recipe is selected across DP/TP/EP
_ps = self.model.get_submodule(name.split(".quant_recipe")[0]).parallel_state
best_format = DistributedProcessGroup.get_dist_syncd_obj(
best_hparam_recipe_info["format"],
[
_ps.data_parallel_group,
_ps.tensor_parallel_group,
_ps.expert_model_parallel_group,
],
lambda a: a[0],
)
best_recipe[name] = best_format
get_hparam(self.model, name).active = best_format
best_constraints += best_hparam_recipe_info["costs"]
best_scores += best_hparam_recipe_info["scores"]
if verbose:
effective_bits_from_search = self._print_recipe_summary(
best_recipe, best_constraints, total_weight_size
)
else:
effective_bits_from_search = (best_constraints / total_weight_size) * 16
self.best["recipe"] = best_recipe
self.best["constraints"] = {"effective_bits": effective_bits_from_search}
self.best["score"] = best_scores
QuantRecipe.fold_pqs_to_weights(self.model)
def _get_auto_quantize_score(grad_output, output_diff):
x = grad_output.float() * output_diff.float()
return x.clamp(-1e10, 1e10).square().sum()
class _AutoQuantizeBackwardScoringSession(ABC):
"""Manage temporary model state used by activation-backward scoring."""
def __init__(
self,
model: nn.Module,
score_modules: Sequence[nn.Module],
is_param_grad_enabled: Callable,
verbose: bool = False,
) -> None:
self.model = model
self.score_modules = tuple(score_modules)
self.is_param_grad_enabled = is_param_grad_enabled
self.verbose = verbose
self._stack = ExitStack()
self._original_forwards: dict[nn.Module, Callable] = {}
self._output_grad_hook_handles: set[Any] = set()
self._grad_accumulators: list[Any] = []
def __enter__(self):
"""Install scoring hooks and parameter settings."""
try:
hparams = list(
dict.fromkeys(
hparam
for module in self.score_modules
for hparam in module._hparams_for_scoring
)
)
for hparam in hparams:
self._stack.callback(setattr, hparam, "active", hparam.active)
def patched_forward(module, *args, **kwargs):
return self.forward(module, *args, **kwargs)
for module in self.score_modules:
original_forward = module.forward
self._original_forwards[module] = original_forward
had_instance_forward = "forward" in module.__dict__
instance_forward = module.__dict__.get("forward")
module.forward = types.MethodType(patched_forward, module)
if had_instance_forward:
self._stack.callback(setattr, module, "forward", instance_forward)
else:
self._stack.callback(module.__dict__.pop, "forward", None)
for name, param in self.model.named_parameters():
requires_grad = param.requires_grad
enable_grad = self.is_param_grad_enabled(name, self.model)
param.requires_grad = enable_grad
self._stack.callback(setattr, param, "requires_grad", requires_grad)
if not enable_grad:
continue
if self.verbose:
print_rank_0(f"AutoQuantize: Enabling gradient for param {name}.")
accumulator, hook = create_param_grad_clear_hook(param)
self._grad_accumulators.append(accumulator)
self._stack.callback(hook.remove)
except Exception:
self._stack.close()
raise
return self
def __exit__(self, exc_type, exc_value, traceback) -> None:
"""Restore all model state changed for scoring."""
self._clear_output_grad_hooks()
self._stack.close()
self._original_forwards.clear()
self._grad_accumulators.clear()
def original_forward(self, module: nn.Module) -> Callable:
"""Return the forward method saved before scoring."""
return self._original_forwards[module]
def _clear_output_grad_hooks(self) -> None:
"""Remove output hooks whose backward pass has not run."""
for handle in self._output_grad_hook_handles:
handle.remove()
self._output_grad_hook_handles.clear()
def _register_output_grad_hook(self, output: torch.Tensor, hook: Callable) -> None:
"""Attach an invocation-specific output-gradient hook for this session."""
def run_once(grad):
try:
return hook(grad)
finally:
handle.remove()
self._output_grad_hook_handles.discard(handle)
handle = output.register_hook(run_once)
self._output_grad_hook_handles.add(handle)
@abstractmethod
def forward(self, module: nn.Module, *args, **kwargs):
"""Run a score module forward pass and collect method-specific state."""
class _AutoQuantizeCandidateReplayScoringSession(_AutoQuantizeBackwardScoringSession):
"""Share baseline execution, candidate replay, and score accumulation."""
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.no_quant = QuantRecipe(quant_cfg=None)
def _run_unquantized(self, module: nn.Module, *args, **kwargs):
"""Run a score module with each configurable group unquantized."""
for hparam in module._hparams_for_scoring:
if hparam.is_configurable:
hparam.active = self.no_quant
output = self.original_forward(module)(*args, **kwargs)
base = output[0] if isinstance(output, tuple) else output
return output, base
def _replay_candidates(self, module: nn.Module, base, candidate_recipes, *args, **kwargs):
"""Measure output differences for the requested candidates."""
output_diffs = {}
with torch.no_grad():
for hparam in module._hparams_for_scoring:
if not hparam.is_configurable:
continue
recipe_diffs = {}
for recipe in candidate_recipes(hparam):
if recipe == self.no_quant:
continue
hparam.active = recipe
try:
replay = self.original_forward(module)(*args, **kwargs)
finally:
hparam.active = self.no_quant
replay = replay[0] if isinstance(replay, tuple) else replay
recipe_diffs[recipe] = (replay - base).detach()
if recipe_diffs:
output_diffs[hparam] = recipe_diffs
return output_diffs
def _accumulate_candidate_scores(self, module, output_diffs, grad_output) -> None:
"""Apply the method's score functional to replayed output differences."""
if grad_output is None:
return
if not torch.isfinite(grad_output).all():
module_name = next(
name for name, child in self.model.named_modules() if child is module
)
raise RuntimeError(
f"AutoQuantize: Non-finite output gradients in module '{module_name or '<root>'}'. "
"Cannot compute reliable sensitivity scores. Check the model, data, and loss. "
"cuDNN SDPA backward on fully masked attention rows is one possible cause; "
"try torch.backends.cuda.enable_cudnn_sdp(False) before rerunning auto_quantize."
)
with torch.no_grad():
for hparam, recipe_diffs in output_diffs.items():
for recipe, output_diff in recipe_diffs.items():
contribution = self._score_contribution(grad_output, output_diff)
current = hparam._importance_dict[recipe][module]
hparam._importance_dict[recipe][module] = (
contribution if current is None else current + contribution
)
def _register_candidate_score_hook(self, module, output, output_diffs) -> None:
"""Bind replayed candidate differences to one output invocation."""
self._register_output_grad_hook(
output,
lambda grad_output: self._accumulate_candidate_scores(
module, output_diffs, grad_output
),
)
@abstractmethod
def _score_contribution(self, grad_output, output_diff):
"""Return this method's score contribution for one replayed candidate."""
class _AutoQuantizeGradientScoringSession(_AutoQuantizeCandidateReplayScoringSession):
"""Collect gradient-based scores while candidate recipes are replayed."""
def forward(self, module: nn.Module, *args, **kwargs):
"""Run the reference forward and cache each recipe's output perturbation."""
output, base = self._run_unquantized(module, *args, **kwargs)
# Checkpointed modules recompute with gradients enabled during backward.
if not torch.is_grad_enabled() or not base.requires_grad:
return output
output_diffs = self._replay_candidates(
module, base, lambda hparam: hparam.choices, *args, **kwargs
)
self._register_candidate_score_hook(module, base, output_diffs)
return output
def _score_contribution(self, grad_output, output_diff):
return _get_auto_quantize_score(grad_output, output_diff)
class _AutoQuantizeBackwardScoringSearcher(_AutoQuantizeBaseSearcher):
"""Share orchestration used by activation-backward scoring methods."""
score_module_rules = [
# Score MoE projections together at their enclosing MLP or mixer output.
r"^(.*?\.mlp)\.experts\.\d+\.(gate_proj|up_proj|down_proj)$",
r"^(.*?\.mixer)\.experts\.\d+\.(up_proj|down_proj)$",
r"^(.*?)\.(\d+\.(w1|w2|w3))$",
r"^(.*?)\.((w1_linear|w2_linear|w3_linear)\.\d+)$",
]
_custom_support: list[tuple[Callable, Callable, Callable]] = []
@classmethod
def register_custom_support(
cls,
is_supported_checker: Callable,
grad_ckpt_context: Callable,
is_param_grad_enabled: Callable,
) -> None:
"""Register optional hooks for memory-efficient backward scoring.
`is_supported_checker` selects models that use these hooks.
`grad_ckpt_context` enables their gradient-checkpointing context, and
`is_param_grad_enabled` selects the minimum parameters needed to propagate
activation gradients.
"""
cls._custom_support.append((is_supported_checker, grad_ckpt_context, is_param_grad_enabled))
def _configurable_score_modules(self) -> list[nn.Module]:
return [
module
for module in self.model.modules()
if hasattr(module, "_hparams_for_scoring")
and any(hparam.is_configurable for hparam in module._hparams_for_scoring)
]
@abstractmethod
def _estimate_auto_quantize_scores(self, is_param_grad_enabled: Callable) -> None:
"""Estimate scores while activation gradients are enabled."""
def estimate_sensitivity_scores(self) -> None:
"""Run backward scoring with the first matching model-specific support hook."""
self.model.eval()
def default_is_param_grad_enabled(_name, _model):
return True
grad_checkpointing_context = None
is_param_grad_enabled = default_is_param_grad_enabled
for is_supported, context_candidate, grad_candidate in self._custom_support:
if is_supported(self.model):
grad_checkpointing_context = context_candidate
is_param_grad_enabled = grad_candidate
break
context = (
grad_checkpointing_context(self.model)
if grad_checkpointing_context is not None
else nullcontext()
)
with context:
self._estimate_auto_quantize_scores(is_param_grad_enabled)
class AutoQuantizeGradientSearcher(_AutoQuantizeBackwardScoringSearcher):
"""A searcher for AutoQuantize algorithm that uses gradient based score estimation.
In AutoQuantize, we search for the best per-layer quantization configuration that minimizes the sum of per-layer
scores while meeting the specified constraint. AutoQuantize uses Linear Programming Solver to find the
optimal quantization configuration.
The auto_quantize score for a layer quantization configuration is an approximation of model loss change due
to quantizing the particular layer with the particular configuration.
The approximation is based on taylor expansion of the loss function wrt to the quantized output of the layer and
substitution of Fisher information for Hessian.
This approximation is mathematically correct for models where the loss
is a log likelihood loss such as BERT, GPT, etc. However, the auto_quantize score can still be used as a proxy
for other models such as ResNet.
**Quant Modules:**
This searcher operates on quantizable modules (quant modules), which are typically Linear or Conv layers
that support quantization. Optionally, grouping rules can be applied to ensure certain layers share the same
quantization format (e.g., Q, K, V projections in the same attention layer). For details on quant_grouping_rules
and customization, see the :meth:`auto_quantize <modelopt.torch.quantization.model_quant.auto_quantize>`
API documentation.
**Score Modules:**
By default, for each quant module, its sensitivity score is estimated using that module's output perturbation.
However, the sensitivity can also be estimated by looking at perturbation at a separate point in the neural
network (score module). This is helpful in some cases such as MoEs for speed and lower memory consumption.
Since all experts are already restricted to the same quant format by quant grouping rules, their sensitivity
can be estimated together at a single point (e.g., the MLP output level).
"""
method_name = "gradient"
@property
def default_search_config(self):
"""Get the default config for the searcher."""
config = super().default_search_config
config.update(
{
"forward_step": None,
"loss_func": None,
"forward_backward_step": None,
}
)
return config
def sanitize_search_config(self, config: SearchConfig | None) -> SearchConfig:
"""Sanitize the search config dict."""
config = config or {}
if "score_func" in config:
if config["score_func"] is not None:
warnings.warn("`score_func` is ignored for gradient based `auto_quantize`.")
config.pop("score_func")
config = super().sanitize_search_config(config)
if config["forward_backward_step"] is None:
assert config["loss_func"] is not None, (
"`loss_func` or `forward_backward_step` must be provided for `auto_quantize`."
)
config["forward_backward_step"] = self._get_default_forward_backward_step()
return config
def _get_default_forward_backward_step(self):
def forward_backward_step(model, data):
output = self.config["forward_step"](model, data)
loss = self.config["loss_func"](output, data)
try:
loss.backward()
except RuntimeError as e:
raise RuntimeError(
"AutoQuantize: Error while calling `backward()` on the loss returned by `loss_func`. "
"Please fix this!"
f"error: {e}"
) from e
return forward_backward_step
@torch.enable_grad()
def _estimate_auto_quantize_scores(self, is_param_grad_enabled):
score_modules = self._configurable_score_modules()
with _AutoQuantizeGradientScoringSession(
self.model,
score_modules,
is_param_grad_enabled,
verbose=self.config.get("verbose", False),
) as scoring_session:
gc.collect()
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
report_memory("AutoQuantize: starting score estimation, ")
def score_step(model, data):
try:
return self.config["forward_backward_step"](model, data)
finally:
scoring_session._clear_output_grad_hooks()
self._run_func(
score_step,
num_iters=self.config["num_score_steps"],
desc="Estimating auto_quantize scores",
)
if torch.cuda.is_available():
report_memory("AutoQuantize: After score estimation")
gc.collect()
def run_search_with_stats(self, max_weight_size, verbose=False):
"""Linear Programming Solve for gradient based auto_quantize.
AutoQuantize uses Linear Programming Solver to find the optimal quantization configuration which
minimizes the sum of per-layer auto_quantize scores while meeting the specified constraint.
"""
return self._run_linear_program_search(max_weight_size, verbose)
@torch.compile(dynamic=True)
def _get_log_softmax_dist(logits: torch.Tensor, tp_group) -> torch.Tensor:
dtype = logits.dtype
max_logits = torch.amax(logits, dim=-1, keepdim=True)
torch.distributed.all_reduce(max_logits, op=torch.distributed.ReduceOp.MAX, group=tp_group)
logits = (logits - max_logits).float()
sum_exp_logits = torch.exp(torch.logsumexp(logits, dim=-1, keepdim=True))
torch.distributed.all_reduce(sum_exp_logits, op=torch.distributed.ReduceOp.SUM, group=tp_group)
return (logits - torch.log(sum_exp_logits)).to(dtype)
def _get_log_prob(logits: torch.Tensor, lm_head: nn.Module = None) -> torch.Tensor:
parallel_state: ParallelState | None = (
getattr(lm_head, "parallel_state", None) if lm_head is not None else None
)
if parallel_state is not None and parallel_state.tensor_parallel_group.is_initialized():
return _get_log_softmax_dist(logits, parallel_state.tensor_parallel_group.group)
return torch.log_softmax(logits.float(), dim=-1)
def _get_kl_div_loss(
log_prob_unquant: torch.Tensor, logits_quant: torch.Tensor, lm_head: nn.Module = None
) -> torch.Tensor:
log_prob_quant = _get_log_prob(logits_quant, lm_head=lm_head)
return F.kl_div(log_prob_quant, log_prob_unquant, reduction="sum", log_target=True)
def _get_lm_head(model: nn.Module) -> nn.Module:
# HF models do allgather of logits to at lm_head
# Hence lm_head outputs are not TP sharded - so we dont need to return the lm_head for TP KLDiv
# Loss
for name, module in model.named_modules():
if name.endswith("output_layer"): # Megatron models
return module
return None
class AutoQuantizeKLDivSearcher(_AutoQuantizeBaseSearcher):
"""A searcher for AutoQuantize algorithm that uses KL-Divergence loss based score estimation."""
method_name = "kl_div"
@property
def default_search_config(self):
"""Get the default config for the searcher."""
config = super().default_search_config
config.update(
{
"forward_step": None,
}
)
return config
def sanitize_search_config(self, config: SearchConfig | None) -> SearchConfig:
"""Sanitize the search config dict."""
config = config or {}
for ignored_key in ["score_func", "loss_func", "forward_backward_step"]:
if ignored_key in config:
if config[ignored_key] is not None:
warnings.warn(
f"`{ignored_key}` is ignored for KL-Divergence loss based `auto_quantize`."
)
config.pop(ignored_key)
config = super().sanitize_search_config(config)
assert config["forward_step"] is not None, (
"`forward_step` must be provided for KL-Divergence loss based `auto_quantize`. "
"`forward_step(model, data)` should return model logits."
)
return config
@torch.inference_mode()
def estimate_sensitivity_scores(self):
"""Estimate the sensitivity scores for the model.
Higher score means more sensitive to quantization.
"""
def set_to_unquantized():
for name, hparam in named_hparams(self.model, unique=True):
if not isinstance(hparam, QuantRecipeHparam):
continue
if hparam.is_configurable:
hparam.active = QuantRecipe(quant_cfg=None)
self.model.eval()
num_iters = self.config["num_score_steps"]
for _, data in tqdm(
zip(range(num_iters), self.config["data_loader"]),
desc="Estimating KLDivergence loss",
total=num_iters,
):
set_to_unquantized()
logits_unquant = self.config["forward_step"](self.model, data)
log_prob_unquant = _get_log_prob(logits_unquant, lm_head=_get_lm_head(self.model))
for name, hparam in tqdm(
list(named_hparams(self.model, configurable=True)), desc="Evaluating hparams"
):
if not isinstance(hparam, QuantRecipeHparam):
continue
for recipe in hparam.choices:
if recipe == QuantRecipe(quant_cfg=None):
continue
hparam.active = recipe
logits_quant = self.config["forward_step"](self.model, data)
score = _get_kl_div_loss(
log_prob_unquant, logits_quant, _get_lm_head(self.model)
)
if hparam._importance_dict[recipe][hparam.score_modules[0]] is None:
hparam._importance_dict[recipe][hparam.score_modules[0]] = score
else:
hparam._importance_dict[recipe][hparam.score_modules[0]] += score
hparam.active = QuantRecipe(quant_cfg=None)
def run_search_with_stats(self, max_weight_size, verbose=False):
"""Run threshold-based binary search for KLDivergence loss based auto_quantize.
We use binary search to minimize the max(per-layer score) while meeting the constraint.
"""
# Collect all sensitivity scores to determine initial threshold bounds
all_scores = [
score for name in self.candidate_stats for score in self.candidate_stats[name]["scores"]
]
if not all_scores:
warnings.warn("No scores available for threshold-based search!")
is_satisfied = False
return {}, is_satisfied
# Initialize binary search bounds
min_score = min(all_scores)
max_score = max(all_scores)
threshold = (min_score + max_score) / 2.0
lower_bound = min_score
upper_bound = max_score
# Run for fixed number of iterations
max_iterations = 100
if verbose:
print_rank_0("AutoQuantize: Starting threshold-based binary search")
print_rank_0(f" Score range: [{min_score:.6e}, {max_score:.6e}]")
print_rank_0(f" Target weight size: {max_weight_size:.2f}")
for iteration in range(max_iterations):
# Select recipes based on current threshold
best_recipes = {}
total_weight_size = 0.0
for name in self.candidate_stats:
formats = self.candidate_stats[name]["formats"]
scores = self.candidate_stats[name]["scores"]
costs = self.candidate_stats[name]["costs"]
selected_idx = 0
for idx in range(len(formats)):
if scores[idx] <= threshold:
selected_idx = idx
break
best_recipes[name] = {
"format": formats[selected_idx],
"costs": costs[selected_idx],
"scores": scores[selected_idx],
}
total_weight_size += costs[selected_idx]
# Check if we meet the constraint
meets_constraint = total_weight_size <= max_weight_size
if verbose:
print_rank_0(
f" Iteration {iteration + 1}: threshold={threshold:.6e}, "
f"weight_size={total_weight_size:.2f}, "
f"meets_constraint={meets_constraint}"
)
# Update binary search bounds
if meets_constraint:
upper_bound = threshold # Threshold was too aggressive, relax it
else:
lower_bound = threshold # Threshold was too lax, tighten it
# Update threshold for next iteration
threshold = (lower_bound + upper_bound) / 2.0
# Final check if constraint is satisfied
is_satisfied = total_weight_size <= max_weight_size
if verbose:
print_rank_0(
f"AutoQuantize: Search complete. "
f"Final weight size: {total_weight_size:.2f} "
f"(target: {max_weight_size:.2f}), "
f"constraint satisfied: {is_satisfied}"
)
return best_recipes, is_satisfied
# Backward compatibility alias (defaults to gradient-based searcher)
AutoQuantizeSearcher = AutoQuantizeGradientSearcher
# Registry of AutoQuantize scoring methods. Optional methods register on import.
AUTO_QUANTIZE_SEARCHERS: dict[str, type[_AutoQuantizeBaseSearcher]] = {
AutoQuantizeGradientSearcher.method_name: AutoQuantizeGradientSearcher,
AutoQuantizeKLDivSearcher.method_name: AutoQuantizeKLDivSearcher,
}
def _as_list(value) -> list:
if value is None:
return []
if isinstance(value, list):
return value
if isinstance(value, tuple):
return list(value)
return [value]
def _get_replay_quantizer_attrs(candidate_stat: dict, module_name: str) -> tuple[str, ...]:
"""Return quantizer attrs that a generated config should target for a searched module."""
quantizer_attrs = candidate_stat.get("quantizer_attrs")
if isinstance(quantizer_attrs, dict):
attrs = quantizer_attrs.get(module_name)
if attrs:
return tuple(attrs)
# Backward-compatible fallback for search checkpoints saved before
# ``quantizer_attrs`` was persisted. Structural HF fused experts are searched
# as modules named ``...mlp.experts`` and expose gate/up + down quantizers.
if module_name.endswith((".mlp.experts", ".mixer.experts")):
return _FUSED_EXPERTS_REPLAY_QUANTIZER_ATTRS
return _STD_QUANTIZER_ATTRS
def get_auto_quantize_config(search_state, constraints=None, verbose=False):
"""Build a flat quant config dict from auto_quantize search_state.
Re-solves for ``constraints`` if provided, otherwise uses the best recipe from the search.
Args:
search_state: The state dict returned by :func:`auto_quantize`.
constraints: Optional dict with ``effective_bits`` key to re-solve for a new target.
verbose: If True, prints the per-layer recipe assignments.
Returns:
A config dict suitable for :func:`quantize`.
"""
if constraints is not None:
best_recipe = _resolve_best_recipe(search_state, constraints, verbose=verbose)
else:
best_recipe = search_state["best"]["recipe"]
def _cfg_to_dict(v):
if isinstance(v, mtq_config.QuantizerAttributeConfig):
return {
"num_bits": v.num_bits,
**v.model_dump(exclude_defaults=True),
}
if isinstance(v, list):
return [_cfg_to_dict(c) for c in v]
return v
fixed_quantization_config = search_state.get("fixed_quantization_config")
quant_cfg: list[dict] = (
copy.deepcopy(fixed_quantization_config["quant_cfg"])
if fixed_quantization_config is not None
else [{"quantizer_name": "*", "enable": False}]
)
quant_cfg.extend(
{"quantizer_name": pattern, "enable": False}
for pattern in _as_list(search_state.get("disabled_layers"))
)
per_module_entries: list[dict] = []
_per_module_attrs = (
*_STD_QUANTIZER_ATTRS,
*_FUSED_EXPERTS_REPLAY_QUANTIZER_ATTRS,
*_NON_GATED_FUSED_EXPERTS_REPLAY_QUANTIZER_ATTRS,
)
# Track global (non per-module) recipe entries. Last recipe wins for each pattern.
global_entries: dict[str, dict] = {}
for hparam_name, recipe in best_recipe.items():
candidate_stat = search_state["candidate_stats"][hparam_name]
if candidate_stat.get("is_fixed", False):
continue
if recipe == QuantRecipe(quant_cfg=None):
continue
module_names = candidate_stat["module_names"]
for module_name in module_names:
for quantizer_attr in _get_replay_quantizer_attrs(candidate_stat, module_name):
matched_cfg, matched_enable = _match_quantizer_cfg(
recipe.config.quant_cfg, quantizer_attr
)
if matched_enable is not None:
entry: dict[str, Any] = {
"quantizer_name": f"{module_name}.{quantizer_attr}",
"enable": matched_enable,
}
if matched_cfg is not None:
entry["cfg"] = _cfg_to_dict(matched_cfg)
per_module_entries.append(entry)
# Collect non-per-module entries (e.g. *[kv]_bmm_quantizer) from winning recipes.
for recipe_entry in recipe.config.quant_cfg:
pattern = recipe_entry["quantizer_name"]
if pattern == "*" or any(
fnmatch.fnmatch(attr, pattern) or pattern.endswith(attr)
for attr in _per_module_attrs
):
continue
cfg = recipe_entry.get("cfg")
enable = recipe_entry.get("enable", True)
ge: dict[str, Any] = {"quantizer_name": pattern, "enable": enable}
if cfg is not None:
ge["cfg"] = _cfg_to_dict(cfg)
global_entries[pattern] = ge
# Keep path-scoped recipe entries before explicit module entries so selected
# modules override default disables such as ``*lm_head*``.
quant_cfg.extend(global_entries.values())
quant_cfg.extend(per_module_entries)
warnings.warn(
"get_auto_quantize_config: returned config uses algorithm='max'. "
"Per-recipe calibration algorithms (e.g. smoothquant, awq) are not preserved. "
"Update config['algorithm'] if a different calibration algorithm is needed (e.g. 'gptq')."
)
return {"quant_cfg": quant_cfg, "algorithm": "max"}
def _resolve_best_recipe(search_state, constraints, verbose=False):
effective_bits = constraints["effective_bits"]
compression = effective_bits / 16.0
candidate_stats = search_state["candidate_stats"]
total_weight_size = search_state.get("cost_denominator") or sum(
s.get("uncompressed_cost", max(s["costs"])) for s in candidate_stats.values()
)
max_weight_size = total_weight_size * compression
method = search_state["method"]
if method not in AUTO_QUANTIZE_SEARCHERS:
raise ValueError(
f"Unknown autoquant search method: {method!r}. "
f"Expected one of {sorted(AUTO_QUANTIZE_SEARCHERS)}."
)
searcher = AUTO_QUANTIZE_SEARCHERS[method]()
searcher.candidate_stats = candidate_stats
searcher.cost_model = search_state.get("cost_model", COST_MODEL_WEIGHT)
searcher.cost = search_state.get("cost", {})
searcher.active_moe_expert_ratio = search_state.get("active_moe_expert_ratio")
if (
searcher.cost_model == COST_MODEL_ACTIVE_MOE
and not searcher.cost
and searcher.active_moe_expert_ratio is not None
):
searcher.cost = {ACTIVE_MOE_EXPERT_RATIO_KEY: searcher.active_moe_expert_ratio}
searcher.config = {
**searcher.default_search_config,
"cost_model": searcher.cost_model,
"cost": searcher.cost,
"active_moe_expert_ratio": searcher.active_moe_expert_ratio,
}
for key in searcher.default_state_dict:
if key in search_state and not hasattr(searcher, key):
setattr(searcher, key, search_state[key])
best_recipe_info, _ = searcher.run_search_with_stats(max_weight_size, verbose=verbose)
best_recipe = {name: info["format"] for name, info in best_recipe_info.items()}
if verbose:
total_cost = sum(info["costs"] for info in best_recipe_info.values())
_AutoQuantizeBaseSearcher._print_recipe_summary(
best_recipe, total_cost, total_weight_size, prefix="get_auto_quantize_config"
)
return best_recipe
def _match_quantizer_cfg(quant_cfg, quantizer_attr):
# Last-match-wins to mirror set_quantizer_by_cfg behavior.
# Patterns may be path-scoped (e.g. "*mlp*weight_quantizer") while quantizer_attr
# is a bare name like "weight_quantizer". We match if the bare name matches directly
# OR if the pattern ends with the bare quantizer_attr (path-scoped match).
matched = None
matched_enable = None
for entry in quant_cfg:
parent_class = entry.get("parent_class") if hasattr(entry, "get") else entry.parent_class
if parent_class is not None:
continue
pattern = entry["quantizer_name"]
cfg = entry.get("cfg")
enable = entry.get("enable", True)
# Direct match: the bare quantizer_attr matches the whole pattern (e.g. "*weight_quantizer")
if fnmatch.fnmatch(quantizer_attr, pattern) or pattern.endswith(quantizer_attr):
matched = cfg
matched_enable = enable
return matched, matched_enable
# Late import avoids a circular dependency: the Aumann-Shapley searcher builds on the
# AutoQuantize classes above and registers itself in AUTO_QUANTIZE_SEARCHERS on import.
from . import _auto_quantize_shapley as _auto_quantize_shapley # noqa: E402