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>