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