mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Add Aumann-Shapley sensitivity scoring method to auto_quantize (#2183)
[Paper](https://arxiv.org/abs/2607.12266) · [Overview](https://x.com/waterloo_intern/status/2076460984475263401) · [Implementation thread](https://x.com/the_joshua_hill/status/2076427869388255635) Depends on #2231, which provides the shared AutoQuantize backward-scoring infrastructure. Until that PR merges, the focused PR B diff is available [here](https://github.com/joshua-hill/Model-Optimizer/compare/fix/autoquant-scoring-infrastructure...feat/aumann-shapley-autoquant). ### What does this PR do? Type of change: new feature This PR adds `method="aumann_shapley"` to `mtq.auto_quantize`. It is a label-free scoring method that measures how each candidate quantization format affects the model across the path from full precision to quantized. The new method is opt-in; existing `gradient` and `kl_div` behavior is unchanged. For each calibration batch, the method: 1. Runs the baseline model and saves its next-token distribution. 2. Measures each candidate format at a configurable number of points along the quantization path. 3. Uses the KL-divergence gradients at those points to assign a damage contribution to every runtime group and candidate format. 4. Measures the most aggressive candidate configuration once and uses that value to calibrate the per-group damage model. The resulting scores use the existing AutoQuantize linear-program solver. The search can either: - choose the least damaging configuration that meets an `effective_bits` target; or - choose the smallest configuration whose predicted damage stays below `max_predicted_damage`. The selected recipe records `predicted_damage` in mean per-token KL units together with its validity and fit diagnostics. ### Public API `auto_quantize` gains an optional `method_options` dictionary. For `method="aumann_shapley"`, it accepts: | Option | Default | Meaning | |---|---:|---| | `num_path_nodes` | `2` | Number of points used to average gradients along the quantization path. | | `damage_link` | `"coverage"` | How per-group scores combine. `"coverage"` uses `damage = c * (1 - exp(-sum(b)))`; `"additive"` sums the path contributions. | | `max_predicted_damage` | `None` | Replaces the bit target with a maximum predicted mean per-token KL. | Method options are validated before the model is modified. Unknown options and incompatible targets fail early. ### Usage Select a configuration for a target effective bit width: ```python import modelopt.torch.quantization as mtq model, search_state = mtq.auto_quantize( model, constraints={"effective_bits": 4.8}, quantization_formats=["NVFP4_DEFAULT_CFG", "FP8_DEFAULT_CFG"], data_loader=calib_loader, forward_step=lambda model, batch: model(**batch), method="aumann_shapley", ) print(search_state["best"]["predicted_damage"]) print(search_state["best"]["predicted_damage_valid"]) ``` Or let the search choose the bit width for a predicted-damage target: ```python model, search_state = mtq.auto_quantize( model, constraints={}, quantization_formats=["NVFP4_DEFAULT_CFG", "FP8_DEFAULT_CFG"], data_loader=calib_loader, forward_step=forward_step, method="aumann_shapley", method_options={"max_predicted_damage": 0.05}, ) ``` ### Implementation - Reuses the candidate-replay and backward-scoring lifecycle introduced in #2231. - Reuses the existing AutoQuantize linear-program solver for both search directions. - Numerically integrates the coverage path when converting measured contributions into per-group damage costs. - Preserves deterministic runtime-group and candidate ordering. - Keeps raw measurements, solver scores, and damage-model diagnostics distinct in the search state. - Rejects incompatible checkpoint resumes while allowing the same scores to be re-solved for a new bit budget. - Retains the shared MoE score-module rules so routed experts are scored at their enclosing block. Recipe integration will follow separately. ### Testing Focused tests: ```text pytest -q \ tests/unit/torch/quantization/test_autoquant.py::test_backward_scoring_session_restores_partial_setup \ tests/unit/torch/quantization/test_autoquant_shapley.py ``` Result: **51 passed**. The tests cover: - end-to-end scoring and configuration generation; - agreement between path contributions and measured quantization damage; - exact allocation checks against exhaustive search; - effective-bits and predicted-damage search modes; - checkpoint resume and offline re-solving; - custom formats and heterogeneous candidate ladders; - distributed reductions and nested MoE score modules; - reused score modules and model-specific backward support; - non-finite measurements and invalid-fit reporting; and - input validation before model conversion. End-to-end checks with NVFP4 and FP8 candidates at a 6.0-bit target: - `Qwen/Qwen2.5-0.5B-Instruct` reaches 5.998 effective bits, with the summed path contributions reproducing 98% of the directly measured lowest-precision KL. - `Qwen/Qwen3-30B-A3B` reaches 6.000 effective bits and reproduces 99%, with all 48 MoE layers scored once at the sparse-MoE block rather than per expert. ### Production use We use this method in production for NVFP4 checkpoints of Kimi-K3, MiniMax-M3, and GLM-5.2. ### Before your PR is "Ready for review" - Is this change backward compatible?: ✅ - If you copied code from any other sources or added a new PIP dependency, did you follow the contributing guidance?: ✅ No copied code and no new dependencies. - Did you write the necessary tests?: ✅ - Did you update `CHANGELOG.rst`?: ✅ - Are the commits signed and signed off?: ✅ <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **New Features** - Added label-free Aumann–Shapley scoring for automatic quantization. - Added configurable path sampling, damage modeling, effective-bit targets, and predicted-damage bounds. - Added temporary weight-folding support with automatic state restoration. - Added method-specific search options, checkpoint resumption, and distributed scoring. - **Bug Fixes** - Improved cleanup and restoration of quantizer state, gradients, hooks, and forward behavior after scoring or failures. - Added validation and clearer handling for unsupported configurations and invalid measurements. - **Tests** - Expanded coverage for scoring, solver behavior, distributed execution, checkpointing, and custom quantization formats. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Joshua Hill <joshua.hill@baseten.co> Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -17,6 +17,7 @@
|
||||
|
||||
import copy
|
||||
import fnmatch
|
||||
import functools
|
||||
import gc
|
||||
import types
|
||||
import warnings
|
||||
@@ -268,6 +269,12 @@ def estimate_quant_compression(quant_cfg: QuantizeConfig) -> float:
|
||||
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.
|
||||
|
||||
@@ -308,6 +315,11 @@ class QuantRecipe(CustomHPType):
|
||||
"""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."""
|
||||
@@ -332,8 +344,10 @@ class QuantRecipe(CustomHPType):
|
||||
return self._str_repr
|
||||
|
||||
def __lt__(self, other: "QuantRecipe"):
|
||||
return (self.compression, self.checkpoint_signature) < (
|
||||
# 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,
|
||||
)
|
||||
|
||||
@@ -659,10 +673,11 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
# 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, list[float]]]
|
||||
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
|
||||
@@ -736,6 +751,9 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
)
|
||||
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)
|
||||
|
||||
@@ -1083,7 +1101,7 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
if not isinstance(hparam, QuantRecipeHparam):
|
||||
continue
|
||||
|
||||
formats, scores, costs = [], [], []
|
||||
formats, raw_scores, scores, costs = [], [], [], []
|
||||
prev_score = float("inf")
|
||||
for recipe in hparam.solver_choices:
|
||||
formats.append(recipe)
|
||||
@@ -1091,6 +1109,7 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
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)
|
||||
@@ -1098,6 +1117,7 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
|
||||
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
|
||||
@@ -1385,6 +1405,46 @@ class _AutoQuantizeBaseSearcher(BaseSearcher, ABC):
|
||||
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."""
|
||||
@@ -1461,10 +1521,6 @@ def _get_auto_quantize_score(grad_output, output_diff):
|
||||
return x.clamp(-1e10, 1e10).square().sum()
|
||||
|
||||
|
||||
def _add_auto_quantize_score(grad_output, output_diff, score_tensor):
|
||||
score_tensor += _get_auto_quantize_score(grad_output, output_diff)
|
||||
|
||||
|
||||
class _AutoQuantizeBackwardScoringSession(ABC):
|
||||
"""Manage temporary model state used by activation-backward scoring."""
|
||||
|
||||
@@ -1563,47 +1619,48 @@ class _AutoQuantizeBackwardScoringSession(ABC):
|
||||
"""Run a score module forward pass and collect method-specific state."""
|
||||
|
||||
|
||||
class _AutoQuantizeGradientScoringSession(_AutoQuantizeBackwardScoringSession):
|
||||
"""Collect gradient-based scores while candidate recipes are replayed."""
|
||||
class _AutoQuantizeCandidateReplayScoringSession(_AutoQuantizeBackwardScoringSession):
|
||||
"""Share baseline execution, candidate replay, and score accumulation."""
|
||||
|
||||
def forward(self, module: nn.Module, *args, **kwargs):
|
||||
"""Run the reference forward and cache each recipe's output perturbation."""
|
||||
no_quant_recipe = QuantRecipe(quant_cfg=None)
|
||||
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 = no_quant_recipe
|
||||
|
||||
hparam.active = self.no_quant
|
||||
output = self.original_forward(module)(*args, **kwargs)
|
||||
base = output[0] if isinstance(output, tuple) else output
|
||||
return output, base
|
||||
|
||||
# Checkpointed modules recompute with gradients enabled during backward.
|
||||
base_output = output[0] if isinstance(output, tuple) else output
|
||||
if not torch.is_grad_enabled() or not base_output.requires_grad:
|
||||
return output
|
||||
|
||||
output_diffs = {hparam: {} for hparam in module._hparams_for_scoring}
|
||||
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
|
||||
for recipe in hparam.choices:
|
||||
if recipe == no_quant_recipe:
|
||||
recipe_diffs = {}
|
||||
for recipe in candidate_recipes(hparam):
|
||||
if recipe == self.no_quant:
|
||||
continue
|
||||
hparam.active = recipe
|
||||
replay = self.original_forward(module)(*args, **kwargs)
|
||||
output_diff = (
|
||||
replay[0] - output[0] if isinstance(replay, tuple) else replay - output
|
||||
)
|
||||
output_diffs[hparam][recipe] = output_diff.detach()
|
||||
hparam.active = no_quant_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
|
||||
|
||||
self._register_output_grad_hook(
|
||||
base_output,
|
||||
lambda grad_output: self._accumulate_scores(module, output_diffs, grad_output),
|
||||
)
|
||||
return output
|
||||
|
||||
def _accumulate_scores(self, module, invocation_diffs, grad_output) -> None:
|
||||
"""Accumulate scores for the invocation that produced ``grad_output``."""
|
||||
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
|
||||
@@ -1614,15 +1671,48 @@ class _AutoQuantizeGradientScoringSession(_AutoQuantizeBackwardScoringSession):
|
||||
"cuDNN SDPA backward on fully masked attention rows is one possible cause; "
|
||||
"try torch.backends.cuda.enable_cudnn_sdp(False) before rerunning auto_quantize."
|
||||
)
|
||||
for hparam, output_diffs in invocation_diffs.items():
|
||||
for recipe, output_diff in output_diffs.items():
|
||||
importance = hparam._importance_dict[recipe][module]
|
||||
if importance is None:
|
||||
hparam._importance_dict[recipe][module] = _get_auto_quantize_score(
|
||||
grad_output, output_diff
|
||||
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
|
||||
)
|
||||
else:
|
||||
_add_auto_quantize_score(grad_output, output_diff, importance)
|
||||
|
||||
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):
|
||||
@@ -1741,7 +1831,8 @@ class AutoQuantizeGradientSearcher(_AutoQuantizeBackwardScoringSearcher):
|
||||
"""Sanitize the search config dict."""
|
||||
config = config or {}
|
||||
if "score_func" in config:
|
||||
warnings.warn("`score_func` is ignored for gradient based `auto_quantize`.")
|
||||
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:
|
||||
@@ -1804,52 +1895,7 @@ class AutoQuantizeGradientSearcher(_AutoQuantizeBackwardScoringSearcher):
|
||||
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.
|
||||
"""
|
||||
# TODO: Do this only for rank 0 in the respective pipeline group
|
||||
|
||||
for lower_bound in self._get_search_lower_bounds():
|
||||
# The LP solver for auto_quantize sometimes fails to find a solution if a lower bound is not
|
||||
# specified. I dont know why this happens.
|
||||
# As a workaround, lets specify a lower bound for the weight compression if previous
|
||||
# search without lower bound fails.
|
||||
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
|
||||
|
||||
if self.status != "Optimal":
|
||||
warnings.warn(
|
||||
"AutoQuantize FAILED to find a solution! The searched model might not meet all constraints. "
|
||||
)
|
||||
is_satisfied = False
|
||||
else:
|
||||
is_satisfied = True
|
||||
|
||||
best_recipes = {}
|
||||
for name, selected_idx in zip(self.candidate_stats.keys(), selections):
|
||||
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
|
||||
return self._run_linear_program_search(max_weight_size, verbose)
|
||||
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
@@ -2055,6 +2101,12 @@ class AutoQuantizeKLDivSearcher(_AutoQuantizeBaseSearcher):
|
||||
# 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:
|
||||
@@ -2187,15 +2239,12 @@ def _resolve_best_recipe(search_state, constraints, verbose=False):
|
||||
max_weight_size = total_weight_size * compression
|
||||
method = search_state["method"]
|
||||
|
||||
if method == "gradient":
|
||||
searcher = AutoQuantizeGradientSearcher()
|
||||
elif method == "kl_div":
|
||||
searcher = AutoQuantizeKLDivSearcher()
|
||||
else:
|
||||
if method not in AUTO_QUANTIZE_SEARCHERS:
|
||||
raise ValueError(
|
||||
f"Unknown autoquant search method: {method!r}. Expected 'gradient' or 'kl_div'."
|
||||
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", {})
|
||||
@@ -2212,8 +2261,10 @@ def _resolve_best_recipe(search_state, constraints, verbose=False):
|
||||
"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())
|
||||
@@ -2244,3 +2295,8 @@ def _match_quantizer_cfg(quant_cfg, quantizer_attr):
|
||||
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
|
||||
|
||||
@@ -39,7 +39,7 @@ from modelopt.torch.quantization.conversion import (
|
||||
from modelopt.torch.utils import atomic_print
|
||||
|
||||
from ._auto_quantize_cost import COST_MODEL_KV_CACHE
|
||||
from .algorithms import AutoQuantizeGradientSearcher, AutoQuantizeKLDivSearcher, QuantRecipe
|
||||
from .algorithms import AUTO_QUANTIZE_SEARCHERS, QuantRecipe
|
||||
from .algorithms import get_auto_quantize_config as _get_auto_quantize_config
|
||||
from .config import QuantizeAlgoCfgType
|
||||
from .kv_cache_auto_quant import AutoQuantizeKVSearcher, get_kv_cache_auto_quantize_config
|
||||
@@ -378,6 +378,24 @@ def _process_quantization_formats(formats, custom_name_prefix):
|
||||
return processed
|
||||
|
||||
|
||||
def _parse_auto_quantize_method(
|
||||
method: str | dict[str, Any] | None,
|
||||
) -> tuple[str | None, dict[str, Any]]:
|
||||
"""Split an AutoQuantize method config into its name and method-specific options."""
|
||||
if method is None or isinstance(method, str):
|
||||
return method, {}
|
||||
if not isinstance(method, dict):
|
||||
raise TypeError(f"`method` must be a string or dict, got {type(method).__name__}.")
|
||||
if "method" not in method:
|
||||
raise ValueError("An AutoQuantize method dictionary must contain a 'method' key.")
|
||||
|
||||
method_config = dict(method)
|
||||
method_name = method_config.pop("method")
|
||||
if not isinstance(method_name, str):
|
||||
raise TypeError("The 'method' value in an AutoQuantize method dictionary must be a string.")
|
||||
return method_name, method_config
|
||||
|
||||
|
||||
def _auto_quantize_kv_cache(
|
||||
model: nn.Module,
|
||||
constraints: dict[str, Any],
|
||||
@@ -488,7 +506,7 @@ def auto_quantize(
|
||||
num_calib_steps: int = 512,
|
||||
num_score_steps: int = 128,
|
||||
verbose: bool = False,
|
||||
method: str | None = None,
|
||||
method: str | dict[str, Any] | None = None,
|
||||
checkpoint: str | None = None,
|
||||
module_search_spaces: list[dict[str, Any]] | None = None,
|
||||
fixed_quantization_config: dict[str, Any] | str | None = None,
|
||||
@@ -496,8 +514,9 @@ def auto_quantize(
|
||||
r"""Perform optimal per-layer quantization by searching for the best quantization formats per-layer.
|
||||
|
||||
``auto_quantize`` uses sensitivity scores to rank the per-layer quantization formats and search
|
||||
for the best quantization formats per-layer. The sensitivity score can be computed using gradient-based
|
||||
methods (default) or KL divergence loss, controlled by the ``method`` parameter.
|
||||
for the best quantization formats per-layer. The sensitivity score can be computed with
|
||||
gradient-based methods (default), KL divergence loss, or Aumann-Shapley path-integral
|
||||
attributions, controlled by the ``method`` parameter.
|
||||
|
||||
Internally this API runs two main phases:
|
||||
|
||||
@@ -652,11 +671,21 @@ def auto_quantize(
|
||||
A higher value could increase the time taken for performing ``auto_quantize``; reducing it speeds up the
|
||||
sensitivity score estimation phase and typically affects accuracy less than lowering ``num_calib_steps``.
|
||||
verbose: If True, prints the search progress/intermediate results.
|
||||
method: Method to use for estimating sensitivity loss. Higher loss indicates greater sensitivity
|
||||
to quantization. Options are ``"gradient"`` (default; uses gradient-based loss estimation,
|
||||
linear programming search, and requires ``loss_func`` or ``forward_backward_step``) and
|
||||
``"kl_div"`` (uses KL divergence between unquantized and quantized outputs, relies on
|
||||
threshold-based binary search, and only requires ``forward_step`` returning logits).
|
||||
method: Method to use for estimating sensitivity loss, either as a string or a dictionary
|
||||
whose ``"method"`` entry selects the method and whose remaining entries configure it.
|
||||
Higher loss indicates greater sensitivity to quantization. Options are ``"gradient"``
|
||||
(default; uses gradient-based loss estimation, linear programming search, and requires
|
||||
``loss_func`` or ``forward_backward_step``), ``"kl_div"`` (uses KL divergence between
|
||||
unquantized and quantized outputs, relies on threshold-based binary search, and only
|
||||
requires ``forward_step`` returning logits), and ``"aumann_shapley"`` (path-integral
|
||||
damage attributions, calibrated against a directly measured reference point; label-free
|
||||
like ``"kl_div"``, and additionally reports a ``predicted_damage`` estimate for the
|
||||
selected recipe). For example, use
|
||||
``{"method": "aumann_shapley", "num_path_nodes": 2}`` or
|
||||
``{"method": "aumann_shapley", "max_predicted_damage": 1e-3}``. Scoring passes grow
|
||||
with the number of candidate formats and path nodes, not with the number of whole-model
|
||||
configurations the search considers -- see
|
||||
:mod:`modelopt.torch.quantization._auto_quantize_shapley`.
|
||||
checkpoint: (Optional) Path to checkpoint file for saving/restoring auto_quantize search state.
|
||||
If the checkpoint file exists, the search state will be restored from it, skipping the
|
||||
expensive score estimation step.
|
||||
@@ -726,10 +755,13 @@ def auto_quantize(
|
||||
raise TypeError("`quantization_formats` must be a sequence of formats.")
|
||||
quantization_formats = list(quantization_formats)
|
||||
|
||||
method_name, method_config = _parse_auto_quantize_method(method)
|
||||
is_kv_search = constraints is not None and constraints.get("cost_model") == COST_MODEL_KV_CACHE
|
||||
if is_kv_search:
|
||||
assert constraints is not None
|
||||
assert quantization_formats is not None
|
||||
if method_config:
|
||||
raise ValueError("cost_model='kv_cache' does not accept method-specific options.")
|
||||
return _auto_quantize_kv_cache(
|
||||
model,
|
||||
constraints,
|
||||
@@ -742,13 +774,13 @@ def auto_quantize(
|
||||
num_calib_steps=num_calib_steps,
|
||||
num_score_steps=num_score_steps,
|
||||
verbose=verbose,
|
||||
method=method,
|
||||
method=method_name,
|
||||
checkpoint=checkpoint,
|
||||
module_search_spaces=module_search_spaces,
|
||||
fixed_quantization_config=fixed_quantization_config,
|
||||
)
|
||||
|
||||
method = method or "gradient"
|
||||
method_name = method_name or "gradient"
|
||||
|
||||
if fixed_quantization_config is None and quantization_formats is None:
|
||||
quantization_formats = [mtq.NVFP4_AWQ_LITE_CFG, mtq.FP8_DEFAULT_CFG]
|
||||
@@ -843,18 +875,12 @@ def auto_quantize(
|
||||
)
|
||||
|
||||
# Select the appropriate searcher based on method
|
||||
if method == "gradient":
|
||||
searcher = AutoQuantizeGradientSearcher()
|
||||
elif method == "kl_div":
|
||||
searcher = AutoQuantizeKLDivSearcher()
|
||||
else:
|
||||
raise ValueError(f"Invalid method: {method}. Valid options are 'gradient' or 'kl_div'.")
|
||||
if method_name not in AUTO_QUANTIZE_SEARCHERS:
|
||||
raise ValueError(
|
||||
f"Invalid method: {method_name}. Valid options are {sorted(AUTO_QUANTIZE_SEARCHERS)}."
|
||||
)
|
||||
searcher = AUTO_QUANTIZE_SEARCHERS[method_name]()
|
||||
|
||||
model = apply_mode(
|
||||
model,
|
||||
mode="auto_quantize",
|
||||
registry=QuantizeModeRegistry,
|
||||
)
|
||||
search_config = {
|
||||
"quantization_formats": processed_quantization_formats,
|
||||
"fixed_quantization_config": processed_fixed_quantization_config,
|
||||
@@ -869,13 +895,35 @@ def auto_quantize(
|
||||
"verbose": verbose,
|
||||
"checkpoint": checkpoint,
|
||||
}
|
||||
if method_config:
|
||||
# Only the selected method's declared options are accepted; core inputs (loaders,
|
||||
# steps, checkpoint, ...) cannot be overridden here.
|
||||
invalid = set(method_config) - searcher.method_config_keys
|
||||
if invalid:
|
||||
raise ValueError(
|
||||
f"Invalid options {sorted(invalid)} for method={method_name!r}. "
|
||||
f"Supported options: {sorted(searcher.method_config_keys)}."
|
||||
)
|
||||
search_config.update(method_config)
|
||||
# Validate the full search config (including method-option values and cross-field
|
||||
# consistency with the constraints) before the model is converted, so a rejected
|
||||
# configuration leaves the model untouched. The searcher re-sanitizes the
|
||||
# already-sanitized config inside search(), which is a no-op.
|
||||
search_config = searcher.sanitize_search_config(search_config)
|
||||
search_constraints = cast("ConstraintsDict", constraints or {})
|
||||
searcher.validate_search_input(search_constraints, search_config)
|
||||
|
||||
model = apply_mode(
|
||||
model,
|
||||
mode="auto_quantize",
|
||||
registry=QuantizeModeRegistry,
|
||||
)
|
||||
# Disable all quantizers; AutoQuantize will enable the needed ones
|
||||
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
|
||||
if processed_fixed_quantization_config is not None:
|
||||
fixed_cfg, fixed_name = processed_fixed_quantization_config
|
||||
fixed_recipe = QuantRecipe(fixed_cfg, name=fixed_name)
|
||||
set_quantizer_by_cfg(model, fixed_recipe.config.quant_cfg)
|
||||
search_constraints = cast("ConstraintsDict", constraints or {})
|
||||
searcher.search(model, search_constraints, config=search_config)
|
||||
|
||||
return model, searcher.state_dict()
|
||||
|
||||
Reference in New Issue
Block a user