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:
Joshua
2026-09-23 20:51:55 -07:00
committed by GitHub
co-authored by Cursor
parent 400498d82d
commit 63c4b660bd
6 changed files with 2669 additions and 124 deletions
File diff suppressed because it is too large Load Diff
+155 -99
View File
@@ -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
+71 -23
View File
@@ -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()