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

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

### What does this PR do?

Type of change: new feature

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

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

For each calibration batch, the method:

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

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

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

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

### Public API

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

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

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

### Usage

Select a configuration for a target effective bit width:

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

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

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

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

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

### Implementation

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

Recipe integration will follow separately.

### Testing

Focused tests:

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

Result: **51 passed**.

The tests cover:

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

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

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

### Production use

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

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

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


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

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

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

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

---------

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