[OMNIML-5570] 2/2 Compose GEMM and KV-cache AutoQuant workflows (#2273)

### What does this PR do?

Type of change: new feature.

Follow-up to merged #2272. Adds composition of existing GEMM
quantization with KV-cache AutoQuantize:

- fixed FP8 GEMM PTQ followed by mixed-KV AutoQuantize;
- gradient-based NVFP4/FP8 GEMM AutoQuantize followed by independent
mixed-KV AutoQuantize;
- an optional `kv_auto_quantize` recipe stage with independent method,
constraints, candidates, and checkpoint path;
- ordered `hf_ptq.py` orchestration that keeps selected
weight/activation QDQ active while its calibration state remains frozen
during KV candidate calibration;
- fail-closed validation when a preceding stage leaves actual K/V
quantizers enabled; and
- unified export of a uniform-weight or mixed-weight checkpoint together
with the selected per-layer KV map.

The KV search still uses the public `mtq.auto_quantize(...,
constraints={"cost_model": "kv_cache", ...})` API from #2272. On a
converted model, the API preserves existing non-KV quantizers and
requires K/V to be disabled before search. Fresh-model behavior is
unchanged and starts from a deny-all quantizer baseline.

#### Why a follow-up field instead of a generic stage list?

This PR deliberately supports the two composition forms required by
`hf_ptq.py` without replacing the stable recipe schema. Existing recipes
already express a fixed `quantize` baseline plus one primary
`auto_quantize` search. A generic ordered `stages` list would require a
broader recipe/API migration, indexed checkpoint semantics, and
compatibility rules for arbitrary stage sequences. There is not yet a
demonstrated third search stage that justifies that surface-area change.

The two searches are not combined inside `mtq.auto_quantize`: each
invocation owns one search domain, constraint model, scoring method, and
resumable checkpoint. Their ordering and independent checkpoint paths
are orchestration concerns, while candidate calibration, scoring,
selection, and state application remain in the shared public API. A
general stage pipeline can be considered separately if more than this
one optional KV follow-up is needed.

Both solvers and scoring protocols are unchanged. The KV checkpoint
compatibility signature additionally fingerprints the preceding
quantizer configuration and calibrated state. Unsupported uniform-weight
plus mixed-KV exports record `kv_cache_deployment_supported: false` in
both ModelOpt and converted HF metadata.

### Usage

Fixed FP8 GEMM PTQ followed by KV AutoQuantize:

```bash
python examples/hf_ptq/hf_ptq.py \
  --pyt_ckpt_path Qwen/Qwen3-8B \
  --recipe general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \
  --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \
  --export_path /path/to/qwen3-8b-fp8-and-mixed-kv
```

Weight AutoQuantize followed by KV AutoQuantize:

```bash
python examples/hf_ptq/hf_ptq.py \
  --pyt_ckpt_path Qwen/Qwen3-8B \
  --recipe general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \
  --auto_quantize_checkpoint /path/to/weight_autoquant.pth \
  --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \
  --export_path /path/to/qwen3-8b-autoquant-and-mixed-kv
```

KV checkpoint resume requires identical preceding non-K/V quantizer
configuration and calibrated state. If rerunning the preceding stage
changes that state, use a new KV checkpoint path to recompute
sensitivities; configuration identity alone is insufficient to reuse the
scores safely.

### Testing

- Latest changed-area validation: 126 tests passed across `hf_ptq.py`
orchestration, KV checkpoint compatibility, export metadata, and HF
configuration conversion.
- A broader local run had 604 passes, one skip, and six failures: two
socket-binding failures under the sandbox and four local Transformers
API incompatibilities. This is not a full-suite pass.
- The fixed-PTQ→KV recipe executes end to end on a tiny offline Qwen
fixture.
- Public API coverage verifies that composed KV search preserves
preceding weight quantization and rejects enabled K/V state.
- Changed-file pre-commit hooks passed; the isolated recipe validator
also passed after dependency bootstrap.

### 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 guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update Changelog?: ✅ (0.48.0 composition feature and KV
checkpoint flag deprecation)
- Did you get Claude approval on this PR?: ❌

### Additional Information

- This follow-up targets `main`, which contains merged #2272.
- `--auto_quantize_checkpoint` and `--kv_auto_quantize_checkpoint` are
intentionally separate because KV sensitivities depend on the preceding
GEMM state.
- Uniform-weight plus mixed-KV exports are for artifact inspection until
the runtime's uniform-weight ModelOpt configuration consumes
`kv_cache_quantized_layers`. Export emits an actionable warning and
records `kv_cache_deployment_supported: false`; this marker does not
itself add runtime support.


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

- **New Features**
- Added staged post-training quantization workflows for weights and KV
caches, including dedicated KV-cache checkpoints.
- Added FP8/NVFP4 recipes with configurable bit constraints and scoring.
  - KV-cache quantization now supports pre-quantized models.

- **Bug Fixes**
- Mixed weight and KV-cache quantization now exports with a warning
instead of failing.
- Improved validation and checkpoint compatibility for staged
configurations.
- Added safeguards for configurations without enabled weight quantizers.

- **Documentation**
- Clarified staged KV-cache workflows, checkpoint options, configuration
behavior, and unsupported deployment combinations.
- Documented deprecated legacy quantization options and their
replacement behavior.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: weimingc <17592131+meenchen@users.noreply.github.com>
This commit is contained in:
Wei-Ming Chen
2026-09-29 00:15:40 +00:00
committed by GitHub
parent de2e810219
commit 1d392999b4
18 changed files with 1711 additions and 664 deletions
+7
View File
@@ -13,6 +13,10 @@ Changelog
*Quantization*
- Add composed Hugging Face AutoQuantize recipes that run fixed PTQ or weight AutoQuantize before
a separate KV-cache AutoQuantize stage, with independent resumable checkpoints for the weight and
KV searches. Uniform-weight plus layer-wise mixed-KV exports are marked unsupported for deployment
until a runtime consumes that metadata combination.
- Add IQ1_S and IQ2_XS weight-only quantization with GGML-compatible 256-value block encoders, built-in ``iq1_s`` / ``iq2_xs`` PTQ recipes, and unified HF and Megatron export of the packed blocks. Quantized weights must have a final dimension divisible by 256, and Megatron export requires tensor and pipeline parallel sizes of 1.
- Add ``iq2_xxs`` weight-only quantization with a CUDA encoder and a ``general/ptq`` recipe, at 2.0625 bits per weight between ``iq1_s`` and ``iq2_xs``. The same 256-value block constraint applies.
- A recipe can now **delegate its whole body to another recipe** with a top-level ``$import``; any top-level key given alongside it overrides the imported one. ``metadata.recipe_type`` became optional along with it: a recipe states its kind with a ``# modelopt-schema:`` comment, with ``metadata.recipe_type``, or by delegating to a recipe that does, and only a recipe that another file imports has to carry the schema comment. Whatever a recipe does state must be true: a schema comment and a ``recipe_type`` must agree, and so must a recipe and the recipe it delegates to. ``modelopt_recipes/models/`` uses this for checkpoint entries that a portable recipe already reproduces: the entry aliases that recipe instead of copying it.
@@ -69,6 +73,9 @@ Changelog
**Deprecations**
- KV-domain searches in ``examples/hf_ptq`` now use ``--kv_auto_quantize_checkpoint``;
``--auto_quantize_checkpoint`` remains a deprecated fallback for a KV-primary recipe for one
release.
- Rename the architecture-specific recipe tier from ``modelopt_recipes/huggingface/`` to ``modelopt_recipes/model_type/`` to clarify that it holds recipes shared across every checkpoint of a Hugging Face ``model_type``. Saved ``--recipe huggingface/<model_type>/...`` paths still resolve via a backward-compatibility alias but now emit a ``FutureWarning``, so update them to ``model_type/<model_type>/...`` as the ``huggingface/`` prefix is deprecated.
- The single-format quantization CLI flags are deprecated in favour of ``--recipe`` and will be removed in a future release; passing one now emits a ``FutureWarning``. ``examples/hf_ptq``: ``--qformat`` and ``--kv_cache_qformat``. ``examples/megatron_bridge/quantize.py``: ``--quant_cfg``, ``--kv_cache_quant`` and ``--weight_only``. ``examples/torch_onnx/torch_quant_to_onnx.py``: ``--qformat``. A recipe carries the quantization config, the calibration algorithm and the KV-cache setting in one file, so they cannot drift apart the way separate flags can -- and ``--recipe`` already took precedence over all six, silently on ``hf_ptq`` and with a warning on ``megatron_bridge`` -- with one gap the recipe closes rather than inherits: a weight AutoQuantize recipe that omits ``kv_cache`` still falls back to ``--kv_cache_qformat``, so set ``kv_cache`` in the recipe when migrating. Use a recipe from ``modelopt_recipes/general/ptq/``, an architecture-specific one under ``modelopt_recipes/model_type/<model_type>/``, or a checkpoint-specific one under ``modelopt_recipes/models/``. The warning fires only when a flag is passed explicitly: ``--qformat`` defaults to ``fp8`` and ``--kv_cache_qformat`` to ``fp8_cast``, so warning on the defaults would fire on every run, including runs that correctly use ``--recipe``. ``examples/speculative_decoding/scripts/quantize_drafter.py`` keeps ``--qformat`` undeprecated: it has no ``--recipe`` alternative yet.
- The TensorRT-LLM checkpoint export format is deprecated and will be removed in 0.49.0: ``export_tensorrt_llm_checkpoint`` and ``torch_to_tensorrt_llm_checkpoint`` now emit a ``DeprecationWarning`` on use. Use ``export_hf_checkpoint``, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang. Its implementation moved to ``modelopt.torch.export.trtllm``, so import those two functions from there and the ``ModelConfig`` dataclasses from ``modelopt.torch.export.trtllm.model_config``; both functions remain importable from ``modelopt.torch.export`` for this release only.
+26 -7
View File
@@ -203,7 +203,7 @@ python hf_ptq.py \
Built-in recipes are located in `modelopt_recipes/general/ptq/` for model-agnostic recipes and in `modelopt_recipes/model_type/<model_type>/ptq/` for recipes tuned to a specific Hugging Face `model_type` (see [`modelopt_recipes/model_type/README.md`](../../modelopt_recipes/model_type/README.md)). You can also provide a path to your own custom YAML recipe file or directory. See the [recipe documentation](https://nvidia.github.io/Model-Optimizer) for details on the YAML schema and available recipes.
> *When `--recipe` is specified, `--qformat` is ignored. KV cache handling depends on the recipe type: a **PTQ** recipe bakes KV cache into its config and ignores `--kv_cache_qformat`; an **AutoQuantize** recipe falls back to `--kv_cache_qformat` unless it sets an explicit `kv_cache` field.*
> *When `--recipe` is specified, `--qformat` is ignored. KV cache handling depends on the recipe type: a **PTQ** recipe bakes KV cache into its config and ignores `--kv_cache_qformat`; an **AutoQuantize** recipe falls back to `--kv_cache_qformat` unless it sets an explicit `kv_cache` field or `kv_auto_quantize` follow-up.*
#### KV Cache Quantization
@@ -480,6 +480,19 @@ For models without backprop support (e.g. Llama-4), use the `kl_div` scoring met
Weight AutoQuantize recipes still apply KV cache as a uniform post-step and fall back to
`--kv_cache_qformat` (default `fp8_cast`) unless they set an explicit `kv_cache` field.
To optimize GEMM and KV cache in one invocation, compose ordered stages in the same recipe. A fixed
`quantize` block followed by a KV-domain `auto_quantize` first calibrates the GEMM weight/activation
configuration, then searches K/V while the existing GEMM QDQ remains enabled with calibration
frozen. See `general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits`.
A weight-domain `auto_quantize` can instead add a `kv_auto_quantize` follow-up with its own method,
constraints, candidates, score size, and disabled layers. This supports, for example, a
gradient-based GEMM search followed by a KL-divergence KV search; see
`general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits`. When the
follow-up is present, the recipe owns KV configuration and suppresses the CLI's uniform
`--kv_cache_qformat` fallback. Use `--auto_quantize_checkpoint` for the weight search and
`--kv_auto_quantize_checkpoint` for the KV search.
KV-cache AutoQuantize recipes use the same `mtq.auto_quantize` API and set
`constraints.cost_model: kv_cache` with an `effective_bits` target. Their
`candidate_formats` are complete K/V cache configs whose config-level `effective_bits` includes
@@ -493,14 +506,14 @@ companion vLLM implementation does not support that asymmetric per-layer format:
python hf_ptq.py \
--pyt_ckpt_path Qwen/Qwen3.8-27B \
--recipe general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits \
--auto_quantize_checkpoint /path/to/kv_autoquant.pth \
--kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \
--export_path /path/to/qwen3.8-27b-mixed-kv
```
Each candidate uses an explicit constant scale, avoiding an additional calibration pass while
keeping persistent K/V scales in the unified HF checkpoint. Unified export records the selected
formats in `kv_cache_quantized_layers`. `mtq.auto_quantize` returns the sensitivity scores and
selected recipe in its search state; `--auto_quantize_checkpoint` stores that resumable state,
selected recipe in its search state; `--kv_auto_quantize_checkpoint` stores that resumable state,
including the candidate quantizer tensors needed for replay.
KV sensitivity scoring runs one reference forward plus one forward per eligible-layer candidate
@@ -513,11 +526,17 @@ times the model vocabulary because reference and candidate log probabilities are
> [vLLM mixed-KV metadata consumer](https://github.com/vllm-project/vllm/pull/52813) or a later
> vLLM release containing it. The repository's currently pinned vLLM 0.26.0 does not consume
> `kv_cache_quantized_layers`, so these checkpoints are export-only in that stock environment.
> Do not deploy them with the pinned runtime. Full FP8 K/V and full NVFP4 K/V use existing vLLM
> kernels once the layer-wise metadata consumer is available.
> The companion consumer also does not yet apply the layer map when uniform FP8/NVFP4 weights are
> present; export warns for that composed combination and records
> `kv_cache_deployment_supported: false`. Do not deploy either unsupported case. Full FP8 K/V and
> full NVFP4 K/V use existing vLLM kernels once the relevant metadata path is available.
The one runtime flag is `--auto_quantize_checkpoint` — save/restore the search state to resume an
interrupted search (skips re-scoring):
`--auto_quantize_checkpoint` saves/restores weight-search state. Every KV-domain search uses
`--kv_auto_quantize_checkpoint`; a KV-primary recipe temporarily accepts the former flag as a
deprecated fallback. Composed recipes therefore keep the weight and KV search states separate. A
KV checkpoint is compatible only with the exact preceding non-K/V quantizer configuration and
calibrated state; if that state cannot be reproduced exactly, use a new checkpoint path to recompute
the KV sensitivities:
```bash
scripts/huggingface_example.sh --model $HF_PATH --recipe general/auto_quantize/nvfp4_fp8_at_5p4bits \
+384
View File
@@ -0,0 +1,384 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""Recipe translation and search helpers for the Hugging Face AutoQuantize example."""
import argparse
import copy
import warnings
from fnmatch import fnmatch
from typing import Any
import torch
from torch.utils.data import DataLoader
import modelopt.torch.quantization as mtq
from modelopt.recipe import ModelOptAutoQuantizeRecipe, load_recipe
from modelopt.recipe.presets import KV_CACHE_NONE, KV_QUANT_CFG_CHOICES, QUANT_CFG_CHOICES
from modelopt.torch.utils.dataset_utils import create_forward_loop
__all__ = ["auto_quantize"]
_FSDP2_KV_AUTOQUANT_ERROR = (
"KV-cache AutoQuantize does not support --use_fsdp2 until distributed sensitivity scoring, "
"selection, and checkpoint writes are synchronized across ranks."
)
_FSDP2_AUTOQUANT_WARNING = (
"AutoQuantize with --use_fsdp2 has not been validated end-to-end yet "
"(distributed calibration, sensitivity scoring, and recipe/checkpoint "
"synchronization across ranks); use at your own risk."
)
# Presets safe to mix into an AutoQuantize search *and* write via the unified HF checkpoint
# exporter. Export-compatibility is a property of the export path, not of a preset's validity for
# plain PTQ, so this is a curated set rather than something derived from QUANT_CFG_CHOICES.
# TODO: drop the partial-model presets (e.g. nvfp4_mlp_only, nvfp4_experts_only) from this set as future work.
_AUTO_QUANTIZE_QFORMATS: frozenset[str] = frozenset(
{
"fp8",
"int8_smoothquant",
"int8_weight_only",
"int4_awq",
"nvfp4",
"nvfp4_awq_lite",
"nvfp4_w4a4_weight_mse_fp8_sweep",
"w4a8_awq_beta",
"w4a16_nvfp4",
"fp8_2d_blockwise_weight_only",
"w4a8_mxfp4_fp8",
"nvfp4_mlp_only",
"nvfp4_experts_only",
"nvfp4_omlp_only",
"nvfp4_w4a4_weight_local_hessian",
"mxfp8",
}
)
def auto_quantize(
args: argparse.Namespace,
language_model: torch.nn.Module,
calib_dataloader: DataLoader,
aq_config,
full_model: torch.nn.Module | None = None,
fixed_quantize_config=None,
allow_uniform_kv: bool = True,
checkpoint: str | None = None,
):
"""Recipe-driven auto_quantize, organized around an AutoQuantizeConfig.
The sole AutoQuantize entry point: it is driven by the recipe's AutoQuantizeConfig and optional
fixed PTQ config, then wraps ``mtq.auto_quantize``.
"""
if args.calib_with_images:
raise NotImplementedError(
"AutoQuantize with image-text calibration is not supported yet. "
"Please run plain PTQ (e.g., --qformat nvfp4) with --calib_with_images."
)
assert args.inference_pipeline_parallel <= 1, (
"Auto Quantization is not supported for pipeline parallel size > 1"
)
inputs = _mtq_inputs_from_auto_quantize_config(
aq_config,
args,
fixed_quantize_config=fixed_quantize_config,
allow_uniform_kv=allow_uniform_kv,
)
if args.use_fsdp2:
if inputs["search_domain"] == "kv_cache":
raise NotImplementedError(_FSDP2_KV_AUTOQUANT_ERROR)
warnings.warn(_FSDP2_AUTOQUANT_WARNING)
# base-model lm_head handling (mirrors the CLI helper)
is_base_model = (
full_model is not None
and language_model is not full_model
and not hasattr(language_model, "lm_head")
and hasattr(full_model, "lm_head")
)
if is_base_model:
assert full_model is not None
lm_head = full_model.lm_head
def loss_func(output, data):
logits = lm_head(output.last_hidden_state)
labels = data["labels"]
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
return torch.nn.functional.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)
)
else:
def loss_func(output, data):
return output.loss
if inputs["method"] == "gradient":
def forward_step(model, batch):
inputs_ = {k: v for k, v in batch.items() if k != "labels"} if is_base_model else batch
return model(**inputs_)
elif inputs["method"] == "kl_div":
def forward_step(model, batch):
inputs_ = {k: v for k, v in batch.items() if k != "labels"} if is_base_model else batch
output = model(**inputs_)
if is_base_model:
assert full_model is not None
logits = full_model.lm_head(output.last_hidden_state)
else:
logits = output.logits
if inputs["search_domain"] == "kv_cache":
return _select_unpadded_logits(logits, batch)
return logits
else:
raise ValueError(
f"Invalid auto_quantize method: {inputs['method']}. Must be 'gradient' or 'kl_div'"
)
auto_quantize_kwargs: dict[str, Any] = {
"constraints": inputs["constraints"],
"data_loader": calib_dataloader,
"forward_step": forward_step,
"quantization_formats": inputs["quantization_formats"],
"num_calib_steps": len(calib_dataloader),
"num_score_steps": min(
len(calib_dataloader), max(inputs["score_size"] // args.batch_size, 1)
),
"verbose": True,
"disabled_layers": inputs["disabled_layers"],
"method": inputs["method"],
"checkpoint": checkpoint,
}
if inputs["search_domain"] == "weight":
auto_quantize_kwargs.update(
{
"loss_func": loss_func,
"fixed_quantization_config": inputs["fixed_quantization_config"],
"module_search_spaces": inputs["module_search_spaces"],
}
)
language_model, _ = mtq.auto_quantize(
language_model,
**auto_quantize_kwargs,
)
if inputs["search_domain"] == "kv_cache":
return language_model
# KV cache quantization is uniform; applied after the LP search.
kv_cache_quant_cfg = inputs["kv_cache_quant_cfg"]
calibrate_loop = create_forward_loop(dataloader=calib_dataloader)
print(f"{'Enable' if kv_cache_quant_cfg is not None else 'Disable'} KV cache quantization")
if kv_cache_quant_cfg is not None:
kv_entries = [
e for e in copy.deepcopy(kv_cache_quant_cfg["quant_cfg"]) if e["quantizer_name"] != "*"
]
mtq.set_quantizer_by_cfg(language_model, quant_cfg=kv_entries)
if not _kv_cfg_uses_constant_amax(kv_entries):
with mtq.set_quantizer_by_cfg_context(
language_model,
[{"quantizer_name": "*", "enable": False}, *kv_entries],
):
mtq.calibrate(language_model, algorithm="max", forward_loop=calibrate_loop)
return language_model
def _mtq_inputs_from_auto_quantize_config(
aq_config,
args: argparse.Namespace,
fixed_quantize_config=None,
allow_uniform_kv: bool = True,
) -> dict:
"""Map a resolved AutoQuantizeConfig to mtq.auto_quantize inputs.
Single, testable place where a recipe maps to mtq inputs. ``fixed_quantize_config`` is the
optional normal PTQ baseline for modules outside explicit search spaces. ``disabled_layers``
and candidate cost come entirely from the recipe (no model introspection). KV cache falls back
to ``--kv_cache_qformat`` when the recipe omits it.
"""
constraints = aq_config.constraints.model_dump(exclude_none=True)
is_kv_search = aq_config.constraints.cost_model == "kv_cache"
if is_kv_search:
return {
"search_domain": "kv_cache",
"constraints": constraints,
"quantization_formats": [
fmt.model_dump(exclude_none=True) for fmt in aq_config.candidate_formats
],
"disabled_layers": aq_config.disabled_layers,
"method": aq_config.auto_quantize_method,
"score_size": aq_config.score_size,
}
# cost_excluded_layers (sibling of disabled_layers) maps to the mtq cost key: these layers are
# kept out of the bit-budget denominator (cost_weight 0) — e.g. VL vision towers — distinct from
# disabled_layers, which removes them from the search.
if aq_config.cost_excluded_layers:
constraints.setdefault("cost", {})["excluded_module_name_patterns"] = (
aq_config.cost_excluded_layers
)
if not allow_uniform_kv:
kv_cache_quant_cfg = None
elif aq_config.kv_cache is not None:
kv_cache_quant_cfg = aq_config.kv_cache.model_dump()
elif args.kv_cache_qformat == KV_CACHE_NONE:
kv_cache_quant_cfg = None
else:
kv_cache_quant_cfg = copy.deepcopy(KV_QUANT_CFG_CHOICES[args.kv_cache_qformat])
# Translate each candidate to its mtq preset dict and, in the same pass, guard export
# compatibility (fails fast, before the expensive search). Custom configs matching no shipped
# preset can't be verified, so warn rather than block.
quantization_formats = _mtq_candidate_formats(aq_config.candidate_formats)
fixed_quantization_config = (
_mtq_candidate_formats([fixed_quantize_config])[0]
if fixed_quantize_config is not None
else None
)
module_search_spaces = [
{
"module_name_patterns": search_space.module_name_patterns,
"quantization_formats": _mtq_candidate_formats(search_space.candidate_formats),
"allow_no_quant": search_space.allow_no_quant,
}
for search_space in aq_config.module_search_spaces
]
return {
"search_domain": "weight",
"constraints": constraints,
"quantization_formats": quantization_formats,
"fixed_quantization_config": fixed_quantization_config,
"module_search_spaces": module_search_spaces,
"disabled_layers": aq_config.disabled_layers,
"kv_cache_quant_cfg": kv_cache_quant_cfg,
"method": aq_config.auto_quantize_method,
"score_size": aq_config.score_size,
}
def _mtq_candidate_formats(formats) -> list[dict]:
"""Translate recipe candidate formats to export-compatible mtq configs."""
quantization_formats = []
for fmt in formats:
preset_name, quant_cfg = _match_candidate_to_preset(fmt)
if preset_name is not None and preset_name not in _AUTO_QUANTIZE_QFORMATS:
raise ValueError(
f"AutoQuantize candidate_formats entry '{preset_name}' is not supported for "
"unified checkpoint export. Use an export-compatible format."
)
if preset_name is None:
warnings.warn(
"An AutoQuantize candidate_formats entry matches no shipped preset; its export "
"compatibility cannot be verified. Ensure it is safe for HF checkpoint export."
)
quantization_formats.append(quant_cfg)
return quantization_formats
def _match_candidate_to_preset(fmt) -> tuple[str | None, dict]:
"""Match a recipe candidate against the shipped QUANT_CFG_CHOICES presets by value.
Returns ``(preset_name, quant_cfg)``: ``preset_name`` is the matched preset (or None for a
custom config matching none), and ``quant_cfg`` is the dict passed to mtq.auto_quantize.
Passing the matched preset dict (rather than the candidate's own dump) keeps the search naming
the candidate after the preset (e.g. FP8_DEFAULT_CFG), consistent with CLI-produced checkpoints.
``effective_bits`` is cost-only metadata (it does not affect export), so it is excluded when
identifying the preset — otherwise a per-candidate override would make a shipped preset look
"custom" and slip past the export-compat whitelist. Any override is preserved in the return.
"""
stripped = fmt.model_dump(exclude_unset=True)
match_key = {k: v for k, v in stripped.items() if k != "effective_bits"}
for name, preset in QUANT_CFG_CHOICES.items():
if preset == match_key:
if "effective_bits" in stripped:
return name, {**preset, "effective_bits": stripped["effective_bits"]}
return name, preset
return None, fmt.model_dump()
def _quantize_config_explicitly_enables_kv(quant_cfg: dict[str, Any]) -> bool:
"""Detect explicit K/V rules while preserving their ordered override semantics."""
names = ("k_bmm_quantizer", "v_bmm_quantizer")
enabled_by_parent = {None: dict.fromkeys(names, False)}
for entry in quant_cfg["quant_cfg"]:
pattern = entry["quantizer_name"]
if pattern != "*" and not any(marker in pattern for marker in ("bmm", "attn", "attention")):
continue
basename_pattern = pattern.rsplit(".", 1)[-1]
matched_names = [
name for name in names if fnmatch(name, basename_pattern) or pattern.endswith(name)
]
if not matched_names:
continue
parent_class = entry.get("parent_class")
if parent_class is None:
scopes = enabled_by_parent.values()
else:
scopes = [enabled_by_parent.setdefault(parent_class, enabled_by_parent[None].copy())]
for enabled in scopes:
for name in matched_names:
enabled[name] = entry["enable"]
return any(any(enabled.values()) for enabled in enabled_by_parent.values())
def _recipe_is_auto_quantize(recipe: str | None) -> bool:
"""True if ``recipe`` resolves to an AutoQuantize recipe (peeked before model load)."""
return recipe is not None and isinstance(load_recipe(recipe), ModelOptAutoQuantizeRecipe)
def _recipe_is_kv_auto_quantize(recipe: str | None) -> bool:
"""True if ``recipe`` resolves to a KV AutoQuantize recipe (peeked before model load)."""
if recipe is None:
return False
loaded_recipe = load_recipe(recipe)
return isinstance(loaded_recipe, ModelOptAutoQuantizeRecipe) and any(
stage is not None and stage.constraints.cost_model == "kv_cache"
for stage in (loaded_recipe.auto_quantize, loaded_recipe.kv_auto_quantize)
)
def _select_unpadded_logits(logits: torch.Tensor, batch: dict[str, Any]) -> torch.Tensor:
"""Return logits only for token positions selected by ``attention_mask``."""
attention_mask = batch.get("attention_mask")
if attention_mask is None:
return logits
if logits.shape[:-1] != attention_mask.shape:
raise ValueError(
"AutoQuantize KL logits and attention_mask must have matching token dimensions; "
f"got {tuple(logits.shape[:-1])} and {tuple(attention_mask.shape)}."
)
return logits[attention_mask.bool()]
def _kv_cfg_uses_constant_amax(kv_quant_cfg: list[dict[str, Any]]) -> bool:
"""Return True if this KV cfg pins ``use_constant_amax`` on the bmm quantizer.
Cast-style KV presets (e.g. ``fp8_cast`` / ``nvfp4_cast``) set
``use_constant_amax: true`` on the ``*[kv]_bmm_quantizer`` entry; that flag
means there is no data-driven calibration to run, so callers should skip
the KV-only calibration pass. Detect the property from the YAML contents
rather than from the preset name so new cast-style presets work
automatically.
"""
for entry in kv_quant_cfg:
if entry.get("quantizer_name") != "*[kv]_bmm_quantizer":
continue
cfg = entry.get("cfg") or {}
return bool(cfg.get("use_constant_amax"))
return False
+86
View File
@@ -33,6 +33,7 @@ import torch
import transformers
from accelerate import infer_auto_device_map, init_empty_weights
from accelerate.utils import get_max_memory
from cast_mxfp4_to_nvfp4 import force_weight_quantizers_static
from safetensors import safe_open
from transformers import (
AutoConfig,
@@ -44,6 +45,7 @@ from transformers import (
ProcessorMixin,
)
from modelopt.torch.export import has_spec_opt
from modelopt.torch.export.model_utils import is_multimodal_model
from modelopt.torch.utils.plugins.hf_checkpoint_utils import (
copy_non_safetensor_files_from_ckpt,
@@ -1074,6 +1076,90 @@ def save_processor_config(args, export_path) -> None:
print("This is normal for some VLM architectures that don't use AutoProcessor")
def _prepare_quant_cfg(
args: argparse.Namespace, quant_cfg: dict[str, Any], full_model: torch.nn.Module
) -> dict[str, Any]:
"""Apply shared checkpoint-local adjustments to a PTQ configuration."""
# Resolve the real export directory before resolve_checkpoint_dir hashes the config; otherwise
# distinct --export_path values containing the placeholder would share one checkpoint path.
if args.layerwise_export:
assert_layerwise_export_compatible(args, full_model, quant_cfg.get("algorithm"))
quant_cfg = set_layerwise_export_dir(quant_cfg, args.export_path)
print(f"Layerwise export enabled: writing quantized shards to {args.export_path}")
# Shards are resumable only while the manifest naming their resume point remains beside
# them; default the calibration checkpoint directory accordingly.
quant_cfg, moved = default_layerwise_resume_dir(quant_cfg, args.export_path)
if moved:
print(
"Layerwise checkpoint_dir co-located with the export path so a resumed run "
"finds its manifest next to the shards it must not overwrite."
)
if needs_checkpoint_path_update(quant_cfg):
quant_cfg, resolved_dir = resolve_checkpoint_dir(quant_cfg, args.pyt_ckpt_path)
print(f"Auto-resolved layerwise checkpoint_dir: {resolved_dir}")
if args.cast_mxfp4_to_nvfp4:
quant_cfg = copy.deepcopy(quant_cfg)
force_weight_quantizers_static(quant_cfg["quant_cfg"])
return quant_cfg
def assert_layerwise_export_compatible(args, full_model, algorithm) -> None:
"""Refuse layerwise export before calibration starts, not after the run is paid for.
Layerwise export writes each layer's shard during calibration and finishes the checkpoint
in finalize() afterwards, so anything that would rewrite or contradict that checkpoint has
to be caught here -- once calibration begins, the user has already paid for the whole run.
"""
block = layerwise_export_block(algorithm)
if block is not None:
entries = algorithm if isinstance(algorithm, list) else [algorithm]
owner = next(e for e in entries if isinstance(e, dict) and e.get("layerwise") is block)
if not owner.get("method"):
raise NotImplementedError(
"layerwise.export_dir needs a calibration method: without one there is no "
"per-layer pass to write the shards, so the export would find nothing. Set "
"algorithm.method, or export without layerwise.export_dir."
)
if has_spec_opt(full_model):
raise NotImplementedError(
"layerwise.export_dir does not support speculative-decoding models: "
"export_speculative_decoding() would write a second checkpoint over the same "
"--export_path."
)
if args.cast_mxfp4_to_nvfp4:
raise NotImplementedError(
"layerwise.export_dir is not compatible with --cast_mxfp4_to_nvfp4: the cast "
"rewrites weights after calibration, by which point every shard is written."
)
# Mirrors export_quantized's branches: a second exporter would overwrite --export_path.
for flag, value, exporter in (
("--vllm_fakequant_export", args.vllm_fakequant_export, "export_hf_vllm_fq_checkpoint()"),
("--sparsity_fmt", args.sparsity_fmt != "dense", "export_tensorrt_llm_checkpoint()"),
(
# int8_sq is the export-format constant, int8_smoothquant the qformat preset.
"--qformat int8_smoothquant",
any(t in args.qformat for t in ("int8_sq", "int8_smoothquant")),
"export_tensorrt_llm_checkpoint()",
),
(
"an encoder-decoder model_type (t5/bart/whisper)",
getattr(full_model.config, "model_type", None) in ("t5", "bart", "whisper"),
"export_tensorrt_llm_checkpoint()",
),
):
if value:
raise NotImplementedError(
f"layerwise.export_dir is not compatible with {flag}: {exporter} would write a "
"second checkpoint over the same --export_path that layerwise calibration "
"already populated."
)
def _layerwise_blocks(algorithm) -> list[dict]:
"""Every ``layerwise`` block in the algorithm, which may be one entry or a list."""
entries = algorithm if isinstance(algorithm, list) else [algorithm]
+97 -400
View File
@@ -14,7 +14,6 @@
# limitations under the License.
import argparse
import copy
import os
import random
import time
@@ -25,30 +24,32 @@ from typing import Any
import numpy as np
import torch
from accelerate.hooks import remove_hook_from_module
from autoquant_utils import (
_FSDP2_KV_AUTOQUANT_ERROR,
_quantize_config_explicitly_enables_kv,
_recipe_is_auto_quantize,
_recipe_is_kv_auto_quantize,
auto_quantize,
)
from cast_mxfp4_to_nvfp4 import apply_to_model as apply_cast_mxfp4_to_nvfp4
from cast_mxfp4_to_nvfp4 import force_weight_quantizers_static
from example_utils import (
HF_PTQ,
_prepare_quant_cfg,
_resolve_model_path,
build_quant_cfg,
cleanup_distributed,
copy_custom_model_files,
create_vlm_calibration_loop,
default_layerwise_resume_dir,
get_model,
get_processor,
get_tokenizer,
is_enc_dec,
is_nemotron_vl,
layerwise_export_block,
mlflow_run,
needs_checkpoint_path_update,
recipe_layerwise_blocks,
resolve_checkpoint_dir,
run_nemotron_vl_preview,
save_processor_config,
save_source_config,
set_layerwise_export_dir,
setup_distributed_args,
validate_fsdp2_supported,
)
@@ -106,47 +107,6 @@ from modelopt.torch.utils.speech_dataset_utils import get_speech_dataset_dataloa
from modelopt.torch.utils.vlm_dataset_utils import get_vlm_dataset_dataloader
RAND_SEED = 1234
_FSDP2_KV_AUTOQUANT_ERROR = (
"KV-cache AutoQuantize does not support --use_fsdp2 until distributed sensitivity scoring, "
"selection, and checkpoint writes are synchronized across ranks."
)
_FSDP2_AUTOQUANT_WARNING = (
"AutoQuantize with --use_fsdp2 has not been validated end-to-end yet "
"(distributed calibration, sensitivity scoring, and recipe/checkpoint "
"synchronization across ranks); use at your own risk."
)
def _select_unpadded_logits(logits: torch.Tensor, batch: dict[str, Any]) -> torch.Tensor:
"""Return logits only for token positions selected by ``attention_mask``."""
attention_mask = batch.get("attention_mask")
if attention_mask is None:
return logits
if logits.shape[:-1] != attention_mask.shape:
raise ValueError(
"AutoQuantize KL logits and attention_mask must have matching token dimensions; "
f"got {tuple(logits.shape[:-1])} and {tuple(attention_mask.shape)}."
)
return logits[attention_mask.bool()]
def _kv_cfg_uses_constant_amax(kv_quant_cfg: list[dict[str, Any]]) -> bool:
"""Return True if this KV cfg pins ``use_constant_amax`` on the bmm quantizer.
Cast-style KV presets (e.g. ``fp8_cast`` / ``nvfp4_cast``) set
``use_constant_amax: true`` on the ``*[kv]_bmm_quantizer`` entry; that flag
means there is no data-driven calibration to run, so callers should skip
the KV-only calibration pass. Detect the property from the YAML contents
rather than from the preset name so new cast-style presets work
automatically.
"""
for entry in kv_quant_cfg:
if entry.get("quantizer_name") != "*[kv]_bmm_quantizer":
continue
cfg = entry.get("cfg") or {}
return bool(cfg.get("use_constant_amax"))
return False
mto.enable_huggingface_checkpointing()
@@ -303,280 +263,18 @@ def make_calib_dataloader(
return calib_dataloader, first_text_speech_dataset
# Presets safe to mix into an AutoQuantize search *and* write via the unified HF checkpoint
# exporter. Export-compatibility is a property of the export path, not of a preset's validity for
# plain PTQ, so this is a curated set rather than something derived from QUANT_CFG_CHOICES.
# TODO: drop the partial-model presets (e.g. nvfp4_mlp_only, nvfp4_experts_only) from this set as future work.
_AUTO_QUANTIZE_QFORMATS: frozenset[str] = frozenset(
{
"fp8",
"int8_smoothquant",
"int8_weight_only",
"int4_awq",
"nvfp4",
"nvfp4_awq_lite",
"nvfp4_w4a4_weight_mse_fp8_sweep",
"w4a8_awq_beta",
"w4a16_nvfp4",
"fp8_2d_blockwise_weight_only",
"w4a8_mxfp4_fp8",
"nvfp4_mlp_only",
"nvfp4_experts_only",
"nvfp4_omlp_only",
"nvfp4_w4a4_weight_local_hessian",
"mxfp8",
}
)
def _match_candidate_to_preset(fmt) -> tuple[str | None, dict]:
"""Match a recipe candidate against the shipped QUANT_CFG_CHOICES presets by value.
Returns ``(preset_name, quant_cfg)``: ``preset_name`` is the matched preset (or None for a
custom config matching none), and ``quant_cfg`` is the dict passed to mtq.auto_quantize.
Passing the matched preset dict (rather than the candidate's own dump) keeps the search naming
the candidate after the preset (e.g. FP8_DEFAULT_CFG), consistent with CLI-produced checkpoints.
``effective_bits`` is cost-only metadata (it does not affect export), so it is excluded when
identifying the preset — otherwise a per-candidate override would make a shipped preset look
"custom" and slip past the export-compat whitelist. Any override is preserved in the return.
"""
stripped = fmt.model_dump(exclude_unset=True)
match_key = {k: v for k, v in stripped.items() if k != "effective_bits"}
for name, preset in QUANT_CFG_CHOICES.items():
if preset == match_key:
if "effective_bits" in stripped:
return name, {**preset, "effective_bits": stripped["effective_bits"]}
return name, preset
return None, fmt.model_dump()
def _mtq_candidate_formats(formats) -> list[dict]:
"""Translate recipe candidate formats to export-compatible mtq configs."""
quantization_formats = []
for fmt in formats:
preset_name, quant_cfg = _match_candidate_to_preset(fmt)
if preset_name is not None and preset_name not in _AUTO_QUANTIZE_QFORMATS:
raise ValueError(
f"AutoQuantize candidate_formats entry '{preset_name}' is not supported for "
"unified checkpoint export. Use an export-compatible format."
)
if preset_name is None:
warnings.warn(
"An AutoQuantize candidate_formats entry matches no shipped preset; its export "
"compatibility cannot be verified. Ensure it is safe for HF checkpoint export."
)
quantization_formats.append(quant_cfg)
return quantization_formats
def _mtq_inputs_from_auto_quantize_config(
aq_config, args: argparse.Namespace, fixed_quantize_config=None
) -> dict:
"""Map a resolved AutoQuantizeConfig to mtq.auto_quantize inputs.
Single, testable place where a recipe maps to mtq inputs. ``fixed_quantize_config`` is the
optional normal PTQ baseline for modules outside explicit search spaces. ``disabled_layers``
and candidate cost come entirely from the recipe (no model introspection). KV cache falls back
to ``--kv_cache_qformat`` when the recipe omits it.
"""
constraints = aq_config.constraints.model_dump(exclude_none=True)
is_kv_search = aq_config.constraints.cost_model == "kv_cache"
if is_kv_search:
return {
"search_domain": "kv_cache",
"constraints": constraints,
"quantization_formats": [
fmt.model_dump(exclude_none=True) for fmt in aq_config.candidate_formats
],
"disabled_layers": aq_config.disabled_layers,
"method": aq_config.auto_quantize_method,
"score_size": aq_config.score_size,
}
# cost_excluded_layers (sibling of disabled_layers) maps to the mtq cost key: these layers are
# kept out of the bit-budget denominator (cost_weight 0) — e.g. VL vision towers — distinct from
# disabled_layers, which removes them from the search.
if aq_config.cost_excluded_layers:
constraints.setdefault("cost", {})["excluded_module_name_patterns"] = (
aq_config.cost_excluded_layers
def _resolve_kv_auto_quantize_checkpoint(args: argparse.Namespace) -> str | None:
"""Resolve a KV-primary checkpoint with a one-release legacy fallback."""
if args.kv_auto_quantize_checkpoint is not None:
return args.kv_auto_quantize_checkpoint
if args.auto_quantize_checkpoint is not None:
warnings.warn(
"Using --auto_quantize_checkpoint for a KV-cache search is deprecated; use "
"--kv_auto_quantize_checkpoint instead.",
FutureWarning,
)
if aq_config.kv_cache is not None:
kv_cache_quant_cfg = aq_config.kv_cache.model_dump()
elif args.kv_cache_qformat == KV_CACHE_NONE:
kv_cache_quant_cfg = None
else:
kv_cache_quant_cfg = copy.deepcopy(KV_QUANT_CFG_CHOICES[args.kv_cache_qformat])
# Translate each candidate to its mtq preset dict and, in the same pass, guard export
# compatibility (fails fast, before the expensive search). Custom configs matching no shipped
# preset can't be verified, so warn rather than block.
quantization_formats = _mtq_candidate_formats(aq_config.candidate_formats)
fixed_quantization_config = (
_mtq_candidate_formats([fixed_quantize_config])[0]
if fixed_quantize_config is not None
else None
)
module_search_spaces = [
{
"module_name_patterns": search_space.module_name_patterns,
"quantization_formats": _mtq_candidate_formats(search_space.candidate_formats),
"allow_no_quant": search_space.allow_no_quant,
}
for search_space in aq_config.module_search_spaces
]
return {
"search_domain": "weight",
"constraints": constraints,
"quantization_formats": quantization_formats,
"fixed_quantization_config": fixed_quantization_config,
"module_search_spaces": module_search_spaces,
"disabled_layers": aq_config.disabled_layers,
"kv_cache_quant_cfg": kv_cache_quant_cfg,
"method": aq_config.auto_quantize_method,
"score_size": aq_config.score_size,
}
def auto_quantize(
args: argparse.Namespace,
language_model: torch.nn.Module,
calib_dataloader: DataLoader,
aq_config,
full_model: torch.nn.Module | None = None,
fixed_quantize_config=None,
):
"""Recipe-driven auto_quantize, organized around an AutoQuantizeConfig.
The sole AutoQuantize entry point: it is driven by the recipe's AutoQuantizeConfig and optional
fixed PTQ config, then wraps ``mtq.auto_quantize``.
"""
if args.calib_with_images:
raise NotImplementedError(
"AutoQuantize with image-text calibration is not supported yet. "
"Please run plain PTQ (e.g., --qformat nvfp4) with --calib_with_images."
)
assert args.inference_pipeline_parallel <= 1, (
"Auto Quantization is not supported for pipeline parallel size > 1"
)
inputs = _mtq_inputs_from_auto_quantize_config(
aq_config, args, fixed_quantize_config=fixed_quantize_config
)
if args.use_fsdp2:
if inputs["search_domain"] == "kv_cache":
raise NotImplementedError(_FSDP2_KV_AUTOQUANT_ERROR)
warnings.warn(_FSDP2_AUTOQUANT_WARNING)
# base-model lm_head handling (mirrors the CLI helper)
is_base_model = (
full_model is not None
and language_model is not full_model
and not hasattr(language_model, "lm_head")
and hasattr(full_model, "lm_head")
)
if is_base_model:
assert full_model is not None
lm_head = full_model.lm_head
def loss_func(output, data):
logits = lm_head(output.last_hidden_state)
labels = data["labels"]
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
return torch.nn.functional.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)
)
else:
def loss_func(output, data):
return output.loss
if inputs["method"] == "gradient":
def forward_step(model, batch):
inputs_ = {k: v for k, v in batch.items() if k != "labels"} if is_base_model else batch
return model(**inputs_)
elif inputs["method"] == "kl_div":
def forward_step(model, batch):
inputs_ = {k: v for k, v in batch.items() if k != "labels"} if is_base_model else batch
output = model(**inputs_)
if is_base_model:
assert full_model is not None
logits = full_model.lm_head(output.last_hidden_state)
else:
logits = output.logits
if inputs["search_domain"] == "kv_cache":
return _select_unpadded_logits(logits, batch)
return logits
else:
raise ValueError(
f"Invalid auto_quantize method: {inputs['method']}. Must be 'gradient' or 'kl_div'"
)
auto_quantize_kwargs: dict[str, Any] = {
"constraints": inputs["constraints"],
"data_loader": calib_dataloader,
"forward_step": forward_step,
"quantization_formats": inputs["quantization_formats"],
"num_calib_steps": len(calib_dataloader),
"num_score_steps": min(
len(calib_dataloader), max(inputs["score_size"] // args.batch_size, 1)
),
"verbose": True,
"disabled_layers": inputs["disabled_layers"],
"method": inputs["method"],
"checkpoint": args.auto_quantize_checkpoint,
}
if inputs["search_domain"] == "weight":
auto_quantize_kwargs.update(
{
"loss_func": loss_func,
"fixed_quantization_config": inputs["fixed_quantization_config"],
"module_search_spaces": inputs["module_search_spaces"],
}
)
language_model, _ = mtq.auto_quantize(
language_model,
**auto_quantize_kwargs,
)
if inputs["search_domain"] == "kv_cache":
return language_model
# KV cache quantization is uniform; applied after the LP search.
kv_cache_quant_cfg = inputs["kv_cache_quant_cfg"]
calibrate_loop = create_forward_loop(dataloader=calib_dataloader)
print(f"{'Enable' if kv_cache_quant_cfg is not None else 'Disable'} KV cache quantization")
if kv_cache_quant_cfg is not None:
kv_entries = [
e for e in copy.deepcopy(kv_cache_quant_cfg["quant_cfg"]) if e["quantizer_name"] != "*"
]
mtq.set_quantizer_by_cfg(language_model, quant_cfg=kv_entries)
if not _kv_cfg_uses_constant_amax(kv_entries):
with mtq.set_quantizer_by_cfg_context(
language_model,
[{"quantizer_name": "*", "enable": False}, *kv_entries],
):
mtq.calibrate(language_model, algorithm="max", forward_loop=calibrate_loop)
return language_model
def _recipe_is_auto_quantize(recipe: str | None) -> bool:
"""True if ``recipe`` resolves to an AutoQuantize recipe (peeked before model load)."""
return recipe is not None and isinstance(load_recipe(recipe), ModelOptAutoQuantizeRecipe)
def _recipe_is_kv_auto_quantize(recipe: str | None) -> bool:
"""True if ``recipe`` resolves to a KV AutoQuantize recipe (peeked before model load)."""
if recipe is None:
return False
loaded_recipe = load_recipe(recipe)
return (
isinstance(loaded_recipe, ModelOptAutoQuantizeRecipe)
and loaded_recipe.auto_quantize.constraints.cost_model == "kv_cache"
)
return args.auto_quantize_checkpoint
return None
def _validate_recipe_calibration(args: argparse.Namespace, recipe) -> None:
@@ -863,60 +561,71 @@ def mono_quantize(
warnings.warn("Skipping quantization: model is already quantized.")
def assert_layerwise_export_compatible(args, full_model, algorithm) -> None:
"""Refuse layerwise export before calibration starts, not after the run is paid for.
def _run_auto_quantize_recipe(
args: argparse.Namespace,
recipe: ModelOptAutoQuantizeRecipe,
full_model: torch.nn.Module,
language_model: torch.nn.Module,
model_type: str | None,
calibration_only: bool,
calib_dataloader: DataLoader,
is_nemotron_vl_model: bool,
) -> None:
"""Run the recipe's fixed PTQ, weight search, and KV search in order."""
primary = recipe.auto_quantize
followup_kv = recipe.kv_auto_quantize
primary_is_kv = primary.constraints.cost_model == "kv_cache"
fixed_quantize_config = recipe.quantize
Layerwise export writes each layer's shard during calibration and finishes the checkpoint
in finalize() afterwards, so anything that would rewrite or contradict that checkpoint has
to be caught here -- once calibration begins, the user has already paid for the whole run.
"""
block = layerwise_export_block(algorithm)
if block is not None:
entries = algorithm if isinstance(algorithm, list) else [algorithm]
owner = next(e for e in entries if isinstance(e, dict) and e.get("layerwise") is block)
if not owner.get("method"):
raise NotImplementedError(
"layerwise.export_dir needs a calibration method: without one there is no "
"per-layer pass to write the shards, so the export would find nothing. Set "
"algorithm.method, or export without layerwise.export_dir."
if fixed_quantize_config is not None and (primary_is_kv or followup_kv is not None):
if _quantize_config_explicitly_enables_kv(fixed_quantize_config.model_dump()):
raise ValueError(
"The fixed quantize stage explicitly enables K/V quantizers before KV-cache "
"AutoQuantize. Disable them in the fixed stage."
)
if has_spec_opt(full_model):
raise NotImplementedError(
"layerwise.export_dir does not support speculative-decoding models: "
"export_speculative_decoding() would write a second checkpoint over the same "
"--export_path."
if primary_is_kv and fixed_quantize_config is not None:
quant_cfg = _prepare_quant_cfg(args, fixed_quantize_config.model_dump(), full_model)
mono_quantize(
args,
quant_cfg,
full_model,
language_model,
model_type,
calibration_only,
calib_dataloader,
is_nemotron_vl_model,
)
fixed_quantize_config = None
if args.cast_mxfp4_to_nvfp4:
raise NotImplementedError(
"layerwise.export_dir is not compatible with --cast_mxfp4_to_nvfp4: the cast "
"rewrites weights after calibration, by which point every shard is written."
auto_quantize(
args,
full_model,
calib_dataloader,
aq_config=primary,
full_model=full_model,
fixed_quantize_config=fixed_quantize_config,
allow_uniform_kv=followup_kv is None,
checkpoint=(
_resolve_kv_auto_quantize_checkpoint(args)
if primary_is_kv
else args.auto_quantize_checkpoint
),
)
if followup_kv is not None:
auto_quantize(
args,
full_model,
calib_dataloader,
aq_config=followup_kv,
full_model=full_model,
allow_uniform_kv=False,
# The weight search owns --auto_quantize_checkpoint, so a follow-up KV search must
# never use the KV-primary legacy fallback and collide with the weight state.
checkpoint=args.kv_auto_quantize_checkpoint,
)
# Mirrors export_quantized's branches: a second exporter would overwrite --export_path.
for flag, value, exporter in (
("--vllm_fakequant_export", args.vllm_fakequant_export, "export_hf_vllm_fq_checkpoint()"),
("--sparsity_fmt", args.sparsity_fmt != "dense", "export_tensorrt_llm_checkpoint()"),
(
# int8_sq is the export-format constant, int8_smoothquant the qformat preset.
"--qformat int8_smoothquant",
any(t in args.qformat for t in ("int8_sq", "int8_smoothquant")),
"export_tensorrt_llm_checkpoint()",
),
(
"an encoder-decoder model_type (t5/bart/whisper)",
getattr(full_model.config, "model_type", None) in ("t5", "bart", "whisper"),
"export_tensorrt_llm_checkpoint()",
),
):
if value:
raise NotImplementedError(
f"layerwise.export_dir is not compatible with {flag}: {exporter} would write a "
"second checkpoint over the same --export_path that layerwise calibration "
"already populated."
)
def export_quantized(
args: argparse.Namespace,
@@ -1274,10 +983,8 @@ def quantize_main(
# AutoQuantize is recipe-driven: everything downstream reads the resolved AutoQuantizeConfig.
if isinstance(recipe, ModelOptAutoQuantizeRecipe):
aq_config = recipe.auto_quantize
fixed_quantize_config = recipe.quantize
else:
aq_config = None
fixed_quantize_config = None
layerwise_cfgs = recipe_layerwise_blocks(recipe)
is_layerwise = any(cfg.get("enable", False) for cfg in layerwise_cfgs)
@@ -1380,16 +1087,16 @@ def quantize_main(
)
if aq_config is not None:
# AutoQuantize (recipe-driven). For VL models the search walks the OUTER CausalLM (which
# carries lm_head and the LM-head forward path); architecture-specific exclusions come
# from aq_config.disabled_layers.
auto_quantize(
assert isinstance(recipe, ModelOptAutoQuantizeRecipe)
_run_auto_quantize_recipe(
args,
recipe,
full_model,
language_model,
model_type,
calibration_only,
calib_dataloader,
aq_config,
full_model=full_model,
fixed_quantize_config=fixed_quantize_config,
is_nemotron_vl_model,
)
else:
@@ -1424,28 +1131,7 @@ def quantize_main(
KV_QUANT_CFG_CHOICES[args.kv_cache_qformat]["quant_cfg"],
)
# Before resolve_checkpoint_dir, which hashes the config: with the placeholder
# still in it, two --export_path values would share one checkpoint dir.
if args.layerwise_export:
assert_layerwise_export_compatible(args, full_model, quant_cfg.get("algorithm"))
quant_cfg = set_layerwise_export_dir(quant_cfg, args.export_path)
print(f"Layerwise export enabled: writing quantized shards to {args.export_path}")
# The shards are only a resume artifact if the manifest that names the resume
# point survives alongside them; see default_layerwise_resume_dir.
quant_cfg, moved = default_layerwise_resume_dir(quant_cfg, args.export_path)
if moved:
print(
"Layerwise checkpoint_dir co-located with the export path so a resumed "
"run finds its manifest next to the shards it must not overwrite."
)
if needs_checkpoint_path_update(quant_cfg):
quant_cfg, resolved_dir = resolve_checkpoint_dir(quant_cfg, args.pyt_ckpt_path)
print(f"Auto-resolved layerwise checkpoint_dir: {resolved_dir}")
if args.cast_mxfp4_to_nvfp4:
quant_cfg = copy.deepcopy(quant_cfg)
force_weight_quantizers_static(quant_cfg["quant_cfg"])
quant_cfg = _prepare_quant_cfg(args, quant_cfg, full_model)
if quant_cfg:
mono_quantize(
@@ -1684,10 +1370,21 @@ def parse_args() -> argparse.Namespace:
type=str,
default=None,
help=(
"Path to checkpoint file for saving/restoring auto_quantize search state "
"Path to checkpoint file for saving/restoring weight AutoQuantize search state "
"(sensitivity scores, costs, etc.). Used with an AutoQuantize --recipe."
),
)
parser.add_argument(
"--kv_auto_quantize_checkpoint",
type=str,
default=None,
help=(
"Path for saving/restoring any KV-cache AutoQuantize search checkpoint. Use a new "
"path whenever the preceding weight/activation quantization stage changes. "
"KV-primary recipes temporarily accept --auto_quantize_checkpoint as a deprecated "
"fallback."
),
)
parser.add_argument(
"--moe_calib_experts_ratio",
type=float,
+26 -6
View File
@@ -351,9 +351,9 @@ class ModelOptAutoQuantizeRecipe(ModelOptRecipeBase):
quantize: QuantizeConfig | None = ModeloptField(
default=None,
title="Fixed PTQ baseline",
description="Optional normal PTQ QuantizeConfig for modules outside the explicit "
"AutoQuantize module_search_spaces. Fixed and searched modules are calibrated, scored, "
"costed, and exported in one integrated AutoQuantize operation.",
description="Optional normal PTQ QuantizeConfig. A weight AutoQuantize stage uses it for "
"modules outside explicit module_search_spaces; a KV AutoQuantize stage applies it first "
"as the fixed GEMM weight/activation configuration.",
)
auto_quantize: AutoQuantizeConfig = Field(
@@ -361,22 +361,42 @@ class ModelOptAutoQuantizeRecipe(ModelOptRecipeBase):
description="AutoQuantize search configuration. Required.",
)
kv_auto_quantize: AutoQuantizeConfig | None = ModeloptField(
default=None,
title="Follow-up KV-cache AutoQuantize config",
description="Optional KV-cache search run after the primary weight AutoQuantize search.",
)
@model_validator(mode="after")
def _validate_fixed_and_searched_spaces(self):
primary_is_kv = self.auto_quantize.constraints.cost_model == "kv_cache"
if self.kv_auto_quantize is not None:
if primary_is_kv:
raise ValueError(
"kv_auto_quantize cannot follow an auto_quantize stage that already searches "
"the KV cache."
)
if self.kv_auto_quantize.constraints.cost_model != "kv_cache":
raise ValueError("kv_auto_quantize must use cost_model=kv_cache.")
if self.auto_quantize.kv_cache is not None:
raise ValueError(
"A weight AutoQuantize stage followed by kv_auto_quantize must omit the "
"uniform auto_quantize.kv_cache post-step."
)
has_fixed_baseline = self.quantize is not None
has_global_search = bool(self.auto_quantize.candidate_formats)
if has_fixed_baseline and has_global_search:
if not primary_is_kv and has_fixed_baseline and has_global_search:
raise ValueError(
"An AutoQuantize recipe with a fixed quantize baseline must omit top-level "
"auto_quantize.candidate_formats and explicitly list searched modules under "
"auto_quantize.module_search_spaces."
)
if has_fixed_baseline and not self.auto_quantize.module_search_spaces:
if not primary_is_kv and has_fixed_baseline and not self.auto_quantize.module_search_spaces:
raise ValueError(
"An AutoQuantize recipe with a fixed quantize baseline requires at least one "
"auto_quantize.module_search_spaces entry."
)
if not has_fixed_baseline and not has_global_search:
if not primary_is_kv and not has_fixed_baseline and not has_global_search:
raise ValueError(
"An AutoQuantize recipe without a fixed quantize baseline requires top-level "
"auto_quantize.candidate_formats for unmatched modules."
@@ -308,6 +308,10 @@ def convert_hf_quant_config_format(input_config: dict[str, Any]) -> dict[str, An
new_config["kv_cache_schema_version"] = original_quantization_details.get(
"kv_cache_schema_version", 1
)
if "kv_cache_deployment_supported" in original_quantization_details:
new_config["kv_cache_deployment_supported"] = original_quantization_details[
"kv_cache_deployment_supported"
]
producer_info = input_config.get("producer")
if producer_info:
+9 -3
View File
@@ -1906,10 +1906,16 @@ def get_quant_config(
)
if needs_layerwise_kv_metadata:
if weight_quant_algo not in (None, "MIXED_PRECISION"):
raise NotImplementedError(
"Mixed-precision KV-cache export with a uniform quantized-weight format is "
"not supported yet. Use BF16 weights or a mixed-weight AutoQuantize recipe."
warn(
"The exported checkpoint combines uniform quantized weights with a mixed-precision "
"KV-cache layer map. Released runtimes do not yet consume "
"kv_cache_quantized_layers for uniform-weight ModelOpt checkpoints. Export succeeds "
"for artifact inspection only; do not deploy this checkpoint until the runtime "
"adds that metadata path. The exported metadata records "
"kv_cache_deployment_supported=false.",
stacklevel=2,
)
quant_config["quantization"]["kv_cache_deployment_supported"] = False
# KV metadata is orthogonal to weight metadata. In particular, a KV-only search
# must preserve BF16 weights instead of synthesizing a weight quantization algorithm.
quant_config["quantization"]["kv_cache_quant_algo"] = (
@@ -18,6 +18,8 @@
from __future__ import annotations
import fnmatch
import hashlib
import json
import math
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, cast
@@ -435,6 +437,7 @@ def _search_signature(
layers: list[tuple[str, nn.Module, int]],
num_calib_steps: int,
num_score_steps: int,
preceding_quantizers: list[dict[str, Any]],
) -> dict[str, Any]:
return {
"schema_version": _KV_AUTOQUANT_SCHEMA_VERSION,
@@ -455,11 +458,82 @@ def _search_signature(
}
for name, module, _ in layers
],
"preceding_quantizers": preceding_quantizers,
}
def _checkpoint_state_is_compatible(state: dict[str, Any], signature: dict[str, Any]) -> bool:
return state.get("search_signature") == signature
checkpoint_signature = state.get("search_signature")
if checkpoint_signature == signature:
return True
if not isinstance(checkpoint_signature, dict) or signature["preceding_quantizers"]:
return False
# Checkpoints written before composed GEMM -> KV searches had no preceding quantizers.
# Preserve their compatibility with an unquantized model while rejecting them for a
# quantized baseline, whose sensitivity scores depend on that baseline.
legacy_signature = signature.copy()
legacy_signature.pop("preceding_quantizers")
return checkpoint_signature == legacy_signature
def _fingerprint_value(value: Any) -> Any:
"""Convert quantizer configuration and tensor state into a stable JSON value."""
if isinstance(value, torch.Tensor):
if value.device.type == "meta":
raise ValueError("Cannot fingerprint a meta-device preceding quantizer state.")
tensor = value.detach().contiguous().cpu()
raw = tensor.reshape(-1).view(torch.uint8).numpy().tobytes()
return {
"dtype": str(tensor.dtype),
"shape": list(tensor.shape),
"sha256": hashlib.sha256(raw).hexdigest(),
}
if hasattr(value, "model_dump"):
return _fingerprint_value(value.model_dump(mode="json"))
if isinstance(value, dict):
return [
[_fingerprint_value(key), _fingerprint_value(item)]
for key, item in sorted(
value.items(), key=lambda entry: (type(entry[0]).__qualname__, repr(entry[0]))
)
]
if isinstance(value, (list, tuple)):
return [_fingerprint_value(item) for item in value]
if isinstance(value, (torch.dtype, torch.device)):
return str(value)
if value is None or isinstance(value, (bool, int, float, str)):
return value
raise TypeError(f"Unsupported preceding quantizer state value: {type(value).__qualname__}.")
def _quantizer_fingerprint(module: TensorQuantizer) -> str:
payload = {
"type": f"{type(module).__module__}.{type(module).__qualname__}",
"properties": module.get_modelopt_state(properties_only=True),
"state_dict": module.state_dict(),
}
serialized = json.dumps(
_fingerprint_value(payload), sort_keys=True, separators=(",", ":"), allow_nan=False
)
return hashlib.sha256(serialized.encode()).hexdigest()
def _preceding_quantizer_signature(model: nn.Module) -> list[dict[str, Any]]:
"""Fingerprint enabled non-K/V configuration and state that affect KV scores."""
return sorted(
(
{
"name": name,
"fingerprint": _quantizer_fingerprint(module),
}
for name, module in model.named_modules(remove_duplicate=False)
if isinstance(module, TensorQuantizer)
and module.is_enabled
and not name.endswith(_KV_QUANTIZER_ATTRS)
),
key=lambda entry: entry["name"],
)
def _quantizer_state_dict(
@@ -691,13 +765,16 @@ class AutoQuantizeKVSearcher(BaseSearcher):
layers,
self.config["num_calib_steps"],
self.config["num_score_steps"],
_preceding_quantizer_signature(self.model),
)
if self.search_signature is not None and not _checkpoint_state_is_compatible(
self.state_dict(), signature
):
raise ValueError(
"KV-cache AutoQuantize checkpoint does not match the current candidates, scoring "
"setup, or eligible layers. Use a different checkpoint path."
"setup, eligible layers, or preceding non-K/V quantizer configuration or calibrated "
"state. Recreate the exact preceding quantized state to resume, or use a different "
"checkpoint path to recompute KV sensitivities."
)
self.search_signature = signature
self._hparams = [
+22 -8
View File
@@ -42,7 +42,11 @@ from ._auto_quantize_cost import COST_MODEL_KV_CACHE
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
from .kv_cache_auto_quant import (
_KV_QUANTIZER_ATTRS,
AutoQuantizeKVSearcher,
get_kv_cache_auto_quantize_config,
)
from .kv_cache_auto_quant import _validate_search_inputs as _validate_kv_cache_search_inputs
from .mode import QuantizeModeRegistry, get_modelike_from_algo_cfg
from .nn import QuantModule, SequentialQuantizer, TensorQuantizer
@@ -424,11 +428,6 @@ def _auto_quantize_kv_cache(
"KV-cache AutoQuantize is single-process only; distributed scoring, selection, "
"and checkpoint writes are not synchronized."
)
if is_quantized(model):
raise NotImplementedError(
"KV-cache AutoQuantize requires an unquantized model; composing it after GEMM "
"PTQ or AutoQuantize is not supported yet."
)
if method not in (None, "kl_div"):
raise ValueError("cost_model='kv_cache' requires method='kl_div'.")
if fixed_quantization_config is not None or module_search_spaces:
@@ -446,6 +445,20 @@ def _auto_quantize_kv_cache(
if data_loader is None or forward_step is None:
raise ValueError("data_loader and forward_step must be provided for KV-cache AutoQuantize.")
converted_for_search = not is_quantized(model)
if not converted_for_search:
enabled_kv_quantizers = [
name
for name, module in model.named_modules(remove_duplicate=False)
if name.endswith(_KV_QUANTIZER_ATTRS) and getattr(module, "is_enabled", False)
]
if enabled_kv_quantizers:
raise ValueError(
"The preceding quantization stage left K/V quantizers enabled: "
f"{enabled_kv_quantizers}. Disable them before running KV-cache AutoQuantize; "
"clearing them now would not undo prior calibration or sensitivity measurements."
)
processed_kv_formats: list[tuple[dict[str, Any], str | None]] = []
for candidate in quantization_formats:
if isinstance(candidate, tuple):
@@ -474,8 +487,9 @@ def _auto_quantize_kv_cache(
num_calib_steps,
num_score_steps,
)
model = apply_mode(model, mode="auto_quantize", registry=QuantizeModeRegistry)
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
if converted_for_search:
model = apply_mode(model, mode="auto_quantize", registry=QuantizeModeRegistry)
set_quantizer_by_cfg(model, [{"quantizer_name": "*", "enable": False}])
searcher = AutoQuantizeKVSearcher()
searcher.search(
model,
@@ -0,0 +1,64 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
# Fixed FP8 GEMM PTQ followed by layer-wise KV-cache AutoQuantize.
# modelopt-schema: modelopt.recipe.config.ModelOptAutoQuantizeRecipe
imports:
base_disable_all: configs/ptq/units/base_disable_all
base_disabled_layers: configs/auto_quantize/units/base_disabled_layers
default_disabled_quantizers: configs/ptq/units/default_disabled_quantizers
fp8: configs/numerics/fp8
nvfp4: configs/numerics/nvfp4
w8a8_fp8_fp8: configs/ptq/units/w8a8_fp8_fp8
metadata:
recipe_type: auto_quantize
description: Fixed FP8 GEMM PTQ followed by mixed FP8/NVFP4 KV-cache search.
quantize:
algorithm: max
quant_cfg:
- $import: base_disable_all
- $import: w8a8_fp8_fp8
- $import: default_disabled_quantizers
auto_quantize:
constraints:
effective_bits: 5.4
cost_model: kv_cache
candidate_formats:
- quant_cfg:
- quantizer_name: "*[kv]_bmm_quantizer"
cfg:
$import: fp8
constant_amax: 448.0
algorithm:
effective_bits: 8.0
- quant_cfg:
- quantizer_name: "*[kv]_bmm_quantizer"
cfg:
$import: nvfp4
constant_amax: 448.0
algorithm:
effective_bits: 4.5
auto_quantize_method: kl_div
score_size: 128
disabled_layers:
- $import: base_disabled_layers
- "*mtp*"
@@ -0,0 +1,74 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
# Gradient-based GEMM AutoQuantize followed by layer-wise KV-cache AutoQuantize.
# modelopt-schema: modelopt.recipe.config.ModelOptAutoQuantizeRecipe
imports:
base_cost_excluded_layers: configs/auto_quantize/units/base_cost_excluded_layers
base_disabled_layers: configs/auto_quantize/units/base_disabled_layers
fp8: configs/ptq/presets/model/fp8
kv_fp8: configs/numerics/fp8
kv_nvfp4: configs/numerics/nvfp4
nvfp4: configs/ptq/presets/model/nvfp4
metadata:
recipe_type: auto_quantize
description: Gradient GEMM search followed by KL-divergence mixed-KV search at 5.4 bits.
auto_quantize:
constraints:
effective_bits: 5.4
candidate_formats:
- $import: nvfp4
- $import: fp8
auto_quantize_method: gradient
score_size: 128
disabled_layers:
- $import: base_disabled_layers
cost_excluded_layers:
- $import: base_cost_excluded_layers
kv_auto_quantize:
constraints:
effective_bits: 5.4
cost_model: kv_cache
candidate_formats:
- quant_cfg:
- quantizer_name: "*[kv]_bmm_quantizer"
cfg:
$import: kv_fp8
constant_amax: 448.0
algorithm:
effective_bits: 8.0
- quant_cfg:
- quantizer_name: "*[kv]_bmm_quantizer"
cfg:
$import: kv_nvfp4
constant_amax: 448.0
algorithm:
effective_bits: 4.5
auto_quantize_method: kl_div
score_size: 128
disabled_layers:
- $import: base_disabled_layers
- "*mtp*"
@@ -0,0 +1,365 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""Offline coverage for the Hugging Face AutoQuantize example helpers."""
import importlib
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from modelopt.recipe import load_recipe
from modelopt.recipe.config import AutoQuantizeConfig, AutoQuantizeConstraints
from modelopt.recipe.presets import QUANT_CFG_CHOICES
from modelopt.torch.quantization.config import QuantizeConfig
@pytest.fixture
def autoquant_utils(monkeypatch):
examples_dir = Path(__file__).resolve().parents[3] / "examples" / "hf_ptq"
monkeypatch.syspath_prepend(str(examples_dir))
return importlib.import_module("autoquant_utils")
def test_autoquant_recipe_builds_mtq_inputs(autoquant_utils):
"""The recipe path maps an AutoQuantizeConfig to the expected mtq.auto_quantize inputs."""
args = SimpleNamespace(kv_cache_qformat="none")
aq = load_recipe("general/auto_quantize/nvfp4_fp8_at_5p4bits").auto_quantize
inputs = autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
# The shared base cost-excluded unit is spliced into every general AutoQuantize recipe, so it
# reaches mtq under constraints.cost (VL vision tower / MTP out of the bit-budget denominator).
assert inputs["constraints"] == {
"effective_bits": 5.4,
"cost_model": "weight",
"cost": {"excluded_module_name_patterns": ["*visual*", "*mtp*", "*vision_tower*"]},
}
assert inputs["kv_cache_quant_cfg"] is None
assert inputs["method"] == "gradient"
assert inputs["score_size"] == 128
assert inputs["fixed_quantization_config"] is None
assert inputs["module_search_spaces"] == []
# disabled_layers come straight from the recipe (no model introspection).
assert inputs["disabled_layers"] == aq.disabled_layers
assert "*output_layer*" in inputs["disabled_layers"]
# Candidates resolve to the exact preset dicts mtq expects (preset identity preserved).
assert inputs["quantization_formats"][0] == QUANT_CFG_CHOICES["nvfp4"]
assert inputs["quantization_formats"][1] == QUANT_CFG_CHOICES["fp8"]
def test_kv_autoquant_recipe_builds_kv_search_inputs(autoquant_utils):
args = SimpleNamespace(kv_cache_qformat="fp8_cast")
aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
inputs = autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
assert inputs["search_domain"] == "kv_cache"
assert inputs["constraints"] == {"effective_bits": 5.4, "cost_model": "kv_cache"}
assert inputs["method"] == "kl_div"
assert [config["effective_bits"] for config in inputs["quantization_formats"]] == [8.0, 4.5]
assert aq.cost_excluded_layers == []
assert "*mtp*" in inputs["disabled_layers"]
assert "kv_cache_quant_cfg" not in inputs
def test_followup_kv_autoquant_suppresses_uniform_kv_fallback(autoquant_utils):
args = SimpleNamespace(kv_cache_qformat="fp8_cast")
aq = load_recipe("general/auto_quantize/nvfp4_fp8_at_5p4bits").auto_quantize
inputs = autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args, allow_uniform_kv=False)
assert inputs["kv_cache_quant_cfg"] is None
def test_fixed_ptq_kv_precheck_does_not_widen_scoped_gemm_rule(autoquant_utils):
fixed = QuantizeConfig(
quant_cfg=[
{
"quantizer_name": "model.layers.*.mlp.*",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
algorithm="max",
)
assert not autoquant_utils._quantize_config_explicitly_enables_kv(fixed.model_dump())
@pytest.mark.parametrize(
"kv_pattern",
[
"model.layers.*.self_attn.*[kv]_bmm_quantizer",
"*self_attn*k_bmm_quantizer",
"*.language_model.*.attention.*_bmm_quantizer",
"*k_bmm*",
"*self_attn.*",
"*[kv]_bmm*",
],
)
def test_fixed_ptq_kv_precheck_detects_scoped_kv_rules(autoquant_utils, kv_pattern):
fixed = QuantizeConfig(
quant_cfg=[
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": kv_pattern,
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
},
{
"parent_class": "nn.Embedding",
"quantizer_name": "*",
"enable": False,
},
],
algorithm="max",
)
assert autoquant_utils._quantize_config_explicitly_enables_kv(fixed.model_dump())
def test_fixed_ptq_kv_precheck_detects_parent_scoped_kv_rule(autoquant_utils):
fixed = QuantizeConfig(
quant_cfg=[
{"quantizer_name": "*", "enable": False},
{
"parent_class": "LlamaAttention",
"quantizer_name": "*_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
},
],
algorithm="max",
)
assert autoquant_utils._quantize_config_explicitly_enables_kv(fixed.model_dump())
def test_kv_autoquant_kl_excludes_padding_positions(autoquant_utils):
logits = torch.arange(2 * 4 * 3).reshape(2, 4, 3)
attention_mask = torch.tensor([[1, 1, 0, 0], [0, 1, 1, 0]])
selected = autoquant_utils._select_unpadded_logits(logits, {"attention_mask": attention_mask})
assert torch.equal(selected, logits[attention_mask.bool()])
def test_kv_autoquant_kl_rejects_misaligned_attention_mask(autoquant_utils):
with pytest.raises(ValueError, match="matching token dimensions"):
autoquant_utils._select_unpadded_logits(
torch.zeros(2, 4, 3), {"attention_mask": torch.ones(2, 3)}
)
@pytest.mark.parametrize(
("search_domain", "expected_shape"),
[("weight", (2, 4, 3)), ("kv_cache", (4, 3))],
)
def test_kl_padding_exclusion_is_scoped_to_kv_autoquant(
autoquant_utils, monkeypatch, search_domain, expected_shape
):
inputs = {
"search_domain": search_domain,
"constraints": {"effective_bits": 8.0},
"quantization_formats": [],
"fixed_quantization_config": None,
"module_search_spaces": [],
"disabled_layers": [],
"kv_cache_quant_cfg": None,
"method": "kl_div",
"score_size": 1,
}
monkeypatch.setattr(
autoquant_utils, "_mtq_inputs_from_auto_quantize_config", lambda *_args, **_kwargs: inputs
)
logits = torch.arange(2 * 4 * 3).reshape(2, 4, 3).float()
batch = {
"input_ids": torch.ones(2, 4, dtype=torch.long),
"attention_mask": torch.tensor([[1, 1, 0, 0], [0, 1, 1, 0]]),
}
class Model(torch.nn.Module):
def forward(self, **_kwargs):
return SimpleNamespace(logits=logits, loss=torch.tensor(0.0))
observed = {}
def fake_auto_quantize(search_model, **kwargs):
observed["shape"] = tuple(kwargs["forward_step"](search_model, batch).shape)
return search_model, {}
monkeypatch.setattr(autoquant_utils.mtq, "auto_quantize", fake_auto_quantize)
args = SimpleNamespace(
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=False,
batch_size=1,
auto_quantize_checkpoint=None,
)
model = Model()
autoquant_utils.auto_quantize(args, model, [batch], SimpleNamespace(), full_model=model)
assert observed["shape"] == expected_shape
def test_kv_autoquant_rejects_fsdp2(autoquant_utils, monkeypatch):
monkeypatch.setattr(
autoquant_utils,
"_mtq_inputs_from_auto_quantize_config",
lambda *_args, **_kwargs: {"search_domain": "kv_cache"},
)
args = SimpleNamespace(
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=True,
)
with pytest.raises(NotImplementedError, match="KV-cache AutoQuantize does not support"):
autoquant_utils.auto_quantize(args, torch.nn.Module(), [], SimpleNamespace())
def test_weight_autoquant_retains_fsdp2_warning(autoquant_utils, monkeypatch):
model = torch.nn.Module()
inputs = {
"search_domain": "weight",
"constraints": {"effective_bits": 8.0},
"quantization_formats": [],
"fixed_quantization_config": None,
"module_search_spaces": [],
"disabled_layers": [],
"kv_cache_quant_cfg": None,
"method": "gradient",
"score_size": 1,
}
monkeypatch.setattr(
autoquant_utils, "_mtq_inputs_from_auto_quantize_config", lambda *_args, **_kwargs: inputs
)
monkeypatch.setattr(
autoquant_utils.mtq, "auto_quantize", lambda search_model, **_kwargs: (search_model, {})
)
args = SimpleNamespace(
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=True,
batch_size=1,
auto_quantize_checkpoint=None,
)
with pytest.warns(UserWarning, match="use at your own risk"):
assert autoquant_utils.auto_quantize(args, model, [], SimpleNamespace()) is model
def test_fsdp2_preload_guard_distinguishes_weight_and_kv_autoquant(autoquant_utils):
assert autoquant_utils._recipe_is_kv_auto_quantize(
"general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits"
)
assert autoquant_utils._recipe_is_kv_auto_quantize(
"general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits"
)
assert not autoquant_utils._recipe_is_kv_auto_quantize(
"general/auto_quantize/nvfp4_fp8_at_5p4bits"
)
def test_autoquant_recipe_cost_excluded_layers_map_into_cost(autoquant_utils):
"""Top-level cost_excluded_layers maps to the mtq constraints.cost.excluded_module_name_patterns
key (distinct from disabled_layers), so a cost-exclusion recipe matches the nested mtq dict."""
args = SimpleNamespace(kv_cache_qformat="none")
aq = load_recipe(
"model_type/qwen3_6_moe/auto_quantize/w4a16_nvfp4_fp8_at_6p0bits-active_moe"
).auto_quantize
inputs = autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
# cost-exclusion is hoisted to a sibling of disabled_layers but still reaches the mtq cost dict.
assert aq.cost_excluded_layers == ["*visual*", "*mtp*", "*vision_tower*"]
assert inputs["constraints"]["cost"] == {
"active_moe_expert_ratio": 0.03125,
"excluded_module_name_patterns": ["*visual*", "*mtp*", "*vision_tower*"],
}
# The two exclusions are independent: cost-excluded patterns are also disabled here, but the
# roles (cost-accounting vs search) are tracked separately.
assert "*visual*" in inputs["disabled_layers"]
def test_autoquant_recipe_maps_module_search_spaces(autoquant_utils):
"""Fixed PTQ baseline and explicit recipe candidates map to mtq inputs."""
args = SimpleNamespace(kv_cache_qformat="none")
recipe = load_recipe(
"model_type/qwen3_6_moe/auto_quantize/w4a16_nvfp4_fp8_module_spaces_at_6p0bits-active_moe"
)
inputs = autoquant_utils._mtq_inputs_from_auto_quantize_config(
recipe.auto_quantize, args, fixed_quantize_config=recipe.quantize
)
model_ptq = load_recipe("model_type/qwen3_5_moe/ptq/w4a16_nvfp4-fp8_attn-kv_fp8_cast")
assert inputs["quantization_formats"] == []
assert inputs["fixed_quantization_config"] == model_ptq.quantize.model_dump()
(searched,) = inputs["module_search_spaces"]
assert searched["module_name_patterns"] == [
"*mlp.shared_expert*",
"*linear_attn*",
"*self_attn*",
"*lm_head*",
]
assert searched["quantization_formats"] == [
QUANT_CFG_CHOICES["w4a16_nvfp4"],
QUANT_CFG_CHOICES["fp8"],
]
assert searched["allow_no_quant"] is False
def test_autoquant_rejects_non_export_safe_candidate(autoquant_utils):
"""A candidate that resolves to a preset outside the export-safe set is rejected before search."""
args = SimpleNamespace(kv_cache_qformat="none")
non_safe = next(
k for k in QUANT_CFG_CHOICES if k not in autoquant_utils._AUTO_QUANTIZE_QFORMATS
)
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=4.8),
candidate_formats=[
QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]),
QuantizeConfig(**QUANT_CFG_CHOICES[non_safe]),
],
)
with pytest.raises(ValueError, match="not supported for unified checkpoint export"):
autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
def test_autoquant_warns_on_custom_candidate(autoquant_utils):
"""A candidate matching no shipped preset can't be export-verified, so it warns (not blocks)."""
args = SimpleNamespace(kv_cache_qformat="none")
custom = QuantizeConfig(quant_cfg=[{"quantizer_name": "*", "enable": False}])
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=4.8),
candidate_formats=[QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]), custom],
)
with pytest.warns(UserWarning, match="export compatibility cannot be verified"):
autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
def test_autoquant_export_guard_not_bypassed_by_effective_bits(autoquant_utils):
"""A non-export-safe preset can't dodge the guard by adding a cost-only effective_bits override."""
args = SimpleNamespace(kv_cache_qformat="none")
non_safe = next(
k for k in QUANT_CFG_CHOICES if k not in autoquant_utils._AUTO_QUANTIZE_QFORMATS
)
tampered = QuantizeConfig(**{**QUANT_CFG_CHOICES[non_safe], "effective_bits": 4.5})
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=5.4),
candidate_formats=[QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]), tampered],
)
with pytest.raises(ValueError, match="not supported for unified checkpoint export"):
autoquant_utils._mtq_inputs_from_auto_quantize_config(aq, args)
+210 -232
View File
@@ -28,7 +28,11 @@ from _test_utils.mlflow import clean_env # noqa: F401
from _test_utils.torch.transformers_models import get_tiny_qwen3
from modelopt.recipe import load_recipe
from modelopt.recipe.config import AutoQuantizeConfig, AutoQuantizeConstraints
from modelopt.recipe.config import (
AutoQuantizeConfig,
AutoQuantizeConstraints,
ModelOptAutoQuantizeRecipe,
)
from modelopt.recipe.presets import QUANT_CFG_CHOICES, RecipeSupersededAction
from modelopt.torch.quantization import tensor_quant
from modelopt.torch.quantization.config import QuantizeConfig
@@ -76,50 +80,6 @@ def test_recipe_help_distinguishes_weight_and_kv_autoquant(monkeypatch, capsys):
assert "KV-cache AutoQuantize recipes select per-layer K/V formats" in help_text
def test_autoquant_recipe_builds_mtq_inputs(monkeypatch):
"""The recipe path maps an AutoQuantizeConfig to the expected mtq.auto_quantize inputs."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
aq = load_recipe("general/auto_quantize/nvfp4_fp8_at_5p4bits").auto_quantize
inputs = hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
# The shared base cost-excluded unit is spliced into every general AutoQuantize recipe, so it
# reaches mtq under constraints.cost (VL vision tower / MTP out of the bit-budget denominator).
assert inputs["constraints"] == {
"effective_bits": 5.4,
"cost_model": "weight",
"cost": {"excluded_module_name_patterns": ["*visual*", "*mtp*", "*vision_tower*"]},
}
assert inputs["kv_cache_quant_cfg"] is None
assert inputs["method"] == "gradient"
assert inputs["score_size"] == 128
assert inputs["fixed_quantization_config"] is None
assert inputs["module_search_spaces"] == []
# disabled_layers come straight from the recipe (no model introspection).
assert inputs["disabled_layers"] == aq.disabled_layers
assert "*output_layer*" in inputs["disabled_layers"]
# Candidates resolve to the exact preset dicts mtq expects (preset identity preserved).
assert inputs["quantization_formats"][0] == QUANT_CFG_CHOICES["nvfp4"]
assert inputs["quantization_formats"][1] == QUANT_CFG_CHOICES["fp8"]
def test_kv_autoquant_recipe_builds_kv_search_inputs(monkeypatch):
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "fp8_cast"
)
aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
inputs = hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
assert inputs["search_domain"] == "kv_cache"
assert inputs["constraints"] == {"effective_bits": 5.4, "cost_model": "kv_cache"}
assert inputs["method"] == "kl_div"
assert [config["effective_bits"] for config in inputs["quantization_formats"]] == [8.0, 4.5]
assert aq.cost_excluded_layers == []
assert "*mtp*" in inputs["disabled_layers"]
assert "kv_cache_quant_cfg" not in inputs
def test_hf_ptq_kv_autoquant_invokes_public_api(monkeypatch):
"""The HF entry point runs the real public KV AutoQuant path on an offline Qwen fixture."""
hf_ptq = _import_hf_ptq(monkeypatch)
@@ -131,12 +91,14 @@ def test_hf_ptq_kv_autoquant_invokes_public_api(monkeypatch):
model = get_tiny_qwen3(num_hidden_layers=1)
aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
args = SimpleNamespace(
qformat="fp8",
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=False,
kv_cache_qformat="none",
batch_size=1,
auto_quantize_checkpoint=None,
kv_auto_quantize_checkpoint=None,
)
data = [{"input_ids": torch.randint(0, model.config.vocab_size, (1, 8))}]
@@ -149,121 +111,243 @@ def test_hf_ptq_kv_autoquant_invokes_public_api(monkeypatch):
assert attention.v_bmm_quantizer.amax == 448.0
def test_kv_autoquant_kl_excludes_padding_positions(monkeypatch):
def test_hf_ptq_runs_weight_then_kv_autoquantize_stages(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
logits = torch.arange(2 * 4 * 3).reshape(2, 4, 3)
attention_mask = torch.tensor([[1, 1, 0, 0], [0, 1, 1, 0]])
selected = hf_ptq._select_unpadded_logits(logits, {"attention_mask": attention_mask})
assert torch.equal(selected, logits[attention_mask.bool()])
def test_kv_autoquant_kl_rejects_misaligned_attention_mask(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
with pytest.raises(ValueError, match="matching token dimensions"):
hf_ptq._select_unpadded_logits(torch.zeros(2, 4, 3), {"attention_mask": torch.ones(2, 3)})
@pytest.mark.parametrize(
("search_domain", "expected_shape"),
[("weight", (2, 4, 3)), ("kv_cache", (4, 3))],
)
def test_kl_padding_exclusion_is_scoped_to_kv_autoquant(monkeypatch, search_domain, expected_shape):
hf_ptq = _import_hf_ptq(monkeypatch)
inputs = {
"search_domain": search_domain,
"constraints": {"effective_bits": 8.0},
"quantization_formats": [],
"fixed_quantization_config": None,
"module_search_spaces": [],
"disabled_layers": [],
"kv_cache_quant_cfg": None,
"method": "kl_div",
"score_size": 1,
}
monkeypatch.setattr(
hf_ptq, "_mtq_inputs_from_auto_quantize_config", lambda *_args, **_kwargs: inputs
weight_aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=8.0),
candidate_formats=[QuantizeConfig(**QUANT_CFG_CHOICES["fp8"])],
)
logits = torch.arange(2 * 4 * 3).reshape(2, 4, 3).float()
batch = {
"input_ids": torch.ones(2, 4, dtype=torch.long),
"attention_mask": torch.tensor([[1, 1, 0, 0], [0, 1, 1, 0]]),
}
kv_aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=8.0, cost_model="kv_cache"),
candidate_formats=[
QuantizeConfig(
quant_cfg=[
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
algorithm=None,
effective_bits=8.0,
)
],
auto_quantize_method="kl_div",
)
recipe = ModelOptAutoQuantizeRecipe(auto_quantize=weight_aq, kv_auto_quantize=kv_aq)
calls = []
monkeypatch.setattr(hf_ptq, "auto_quantize", lambda *_args, **kwargs: calls.append(kwargs))
class Model(torch.nn.Module):
def forward(self, **_kwargs):
return SimpleNamespace(logits=logits, loss=torch.tensor(0.0))
hf_ptq._run_auto_quantize_recipe(
SimpleNamespace(
auto_quantize_checkpoint="weight-search.pth",
kv_auto_quantize_checkpoint="kv-search.pth",
),
recipe,
torch.nn.Module(),
torch.nn.Module(),
None,
False,
[],
False,
)
observed = {}
assert [call["aq_config"] for call in calls] == [weight_aq, kv_aq]
assert calls[0]["allow_uniform_kv"] is False
assert calls[0]["checkpoint"] == "weight-search.pth"
assert calls[1]["checkpoint"] == "kv-search.pth"
def fake_auto_quantize(search_model, **kwargs):
observed["shape"] = tuple(kwargs["forward_step"](search_model, batch).shape)
return search_model, {}
monkeypatch.setattr(hf_ptq.mtq, "auto_quantize", fake_auto_quantize)
def test_hf_ptq_runs_real_weight_then_kv_autoquantize_stages(monkeypatch):
"""Exercise the shipped gradient-weight -> KL-div KV composition without mocked stages."""
hf_ptq = _import_hf_ptq(monkeypatch)
monkeypatch.setattr(
tensor_quant,
"dynamic_block_quantize_op",
lambda inputs, *_args, **_kwargs: torch.zeros_like(inputs),
)
recipe = load_recipe(
"general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits"
)
model = get_tiny_qwen3(num_hidden_layers=1)
input_ids = torch.arange(8).unsqueeze(0) % model.config.vocab_size
data = [{"input_ids": input_ids, "labels": input_ids.clone()}]
args = SimpleNamespace(
qformat="fp8",
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=False,
kv_cache_qformat="none",
batch_size=1,
auto_quantize_checkpoint=None,
kv_auto_quantize_checkpoint=None,
)
hf_ptq._run_auto_quantize_recipe(args, recipe, model, model, None, False, data, False)
enabled_weight_quantizers = [
module
for name, module in model.named_modules()
if name.endswith("weight_quantizer") and getattr(module, "is_enabled", False)
]
assert enabled_weight_quantizers
assert all(module.num_bits in ((2, 1), (4, 3)) for module in enabled_weight_quantizers)
attention = model.model.layers[0].self_attn
assert attention.k_bmm_quantizer.is_enabled
assert attention.v_bmm_quantizer.is_enabled
assert attention.k_bmm_quantizer.num_bits in ((2, 1), (4, 3))
assert attention.v_bmm_quantizer.num_bits in ((2, 1), (4, 3))
def test_hf_ptq_runs_fixed_ptq_before_kv_autoquantize(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
monkeypatch.setattr(
tensor_quant,
"dynamic_block_quantize_op",
lambda inputs, *_args, **_kwargs: torch.zeros_like(inputs),
)
recipe = load_recipe("general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits")
model = get_tiny_qwen3(num_hidden_layers=1)
data = [{"input_ids": torch.randint(0, model.config.vocab_size, (1, 8))}]
args = SimpleNamespace(
qformat="fp8",
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=False,
batch_size=1,
auto_quantize_checkpoint=None,
kv_auto_quantize_checkpoint=None,
pyt_ckpt_path="dummy",
cast_mxfp4_to_nvfp4=False,
layerwise_export=False,
specdec_offline_dataset=None,
)
model = Model()
hf_ptq.auto_quantize(args, model, [batch], SimpleNamespace(), full_model=model)
hf_ptq._run_auto_quantize_recipe(args, recipe, model, model, None, False, data, False)
assert observed["shape"] == expected_shape
attention = model.model.layers[0].self_attn
assert attention.q_proj.weight_quantizer.is_enabled
assert attention.q_proj.weight_quantizer.num_bits == (4, 3)
assert attention.k_bmm_quantizer.is_enabled
assert attention.v_bmm_quantizer.is_enabled
def test_kv_autoquant_rejects_fsdp2(monkeypatch):
def test_kv_autoquantize_checkpoint_uses_dedicated_flag_with_legacy_fallback(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
args = SimpleNamespace(
auto_quantize_checkpoint="legacy.pth",
kv_auto_quantize_checkpoint="kv.pth",
)
assert hf_ptq._resolve_kv_auto_quantize_checkpoint(args) == "kv.pth"
args.kv_auto_quantize_checkpoint = None
with pytest.warns(FutureWarning, match="deprecated"):
assert hf_ptq._resolve_kv_auto_quantize_checkpoint(args) == "legacy.pth"
def test_fixed_ptq_then_kv_rejects_explicit_kv_before_calibration(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
fixed = QuantizeConfig(
quant_cfg=[
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "model.layers.*.self_attn.*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
},
],
algorithm="max",
)
kv_aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
recipe = ModelOptAutoQuantizeRecipe(quantize=fixed, auto_quantize=kv_aq)
args = SimpleNamespace(
auto_quantize_checkpoint=None,
kv_auto_quantize_checkpoint=None,
pyt_ckpt_path="dummy",
cast_mxfp4_to_nvfp4=False,
layerwise_export=False,
)
monkeypatch.setattr(
hf_ptq,
"_mtq_inputs_from_auto_quantize_config",
lambda *_args, **_kwargs: {"search_domain": "kv_cache"},
)
args = SimpleNamespace(
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=True,
"mono_quantize",
lambda *_args, **_kwargs: pytest.fail("fixed PTQ must not start"),
)
with pytest.raises(NotImplementedError, match="KV-cache AutoQuantize does not support"):
hf_ptq.auto_quantize(args, torch.nn.Module(), [], SimpleNamespace())
with pytest.raises(ValueError, match="fixed quantize stage explicitly enables K/V"):
hf_ptq._run_auto_quantize_recipe(
args, recipe, torch.nn.Module(), torch.nn.Module(), None, False, [], False
)
def test_weight_autoquant_retains_fsdp2_warning(monkeypatch):
def test_weight_autoquant_then_kv_rejects_fixed_kv_before_weight_search(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
model = torch.nn.Module()
inputs = {
"search_domain": "weight",
"constraints": {"effective_bits": 8.0},
"quantization_formats": [],
"fixed_quantization_config": None,
"module_search_spaces": [],
"disabled_layers": [],
"kv_cache_quant_cfg": None,
"method": "gradient",
"score_size": 1,
}
monkeypatch.setattr(
hf_ptq, "_mtq_inputs_from_auto_quantize_config", lambda *_args, **_kwargs: inputs
fixed = QuantizeConfig(
quant_cfg=[
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*self_attn.*",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
},
],
algorithm="max",
)
weight_aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=8.0),
module_search_spaces=[
{
"module_name_patterns": ["*mlp*"],
"candidate_formats": [QuantizeConfig(**QUANT_CFG_CHOICES["fp8"])],
}
],
)
kv_aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
recipe = ModelOptAutoQuantizeRecipe(
quantize=fixed, auto_quantize=weight_aq, kv_auto_quantize=kv_aq
)
monkeypatch.setattr(
hf_ptq.mtq, "auto_quantize", lambda search_model, **_kwargs: (search_model, {})
hf_ptq,
"auto_quantize",
lambda *_args, **_kwargs: pytest.fail("weight AutoQuantize must not start"),
)
with pytest.raises(ValueError, match="fixed quantize stage explicitly enables K/V"):
hf_ptq._run_auto_quantize_recipe(
SimpleNamespace(),
recipe,
torch.nn.Module(),
torch.nn.Module(),
None,
False,
[],
False,
)
def test_composed_kv_autoquant_rejects_enabled_actual_kv_quantizers(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
model = get_tiny_qwen3(num_hidden_layers=1)
hf_ptq.mtq.quantize(
model,
{
"quant_cfg": [
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
"algorithm": None,
},
)
args = SimpleNamespace(
calib_with_images=False,
inference_pipeline_parallel=1,
use_fsdp2=True,
use_fsdp2=False,
kv_cache_qformat="none",
batch_size=1,
auto_quantize_checkpoint=None,
)
aq = load_recipe("general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits").auto_quantize
with pytest.warns(UserWarning, match="use at your own risk"):
assert hf_ptq.auto_quantize(args, model, [], SimpleNamespace()) is model
with pytest.raises(ValueError, match="preceding quantization stage left K/V"):
hf_ptq.auto_quantize(args, model, [], aq, full_model=model)
def test_fsdp2_kv_autoquant_rejected_before_model_load(monkeypatch):
@@ -279,112 +363,6 @@ def test_fsdp2_kv_autoquant_rejected_before_model_load(monkeypatch):
hf_ptq.load_model(SimpleNamespace(use_fsdp2=True, recipe="autoquant"))
def test_fsdp2_preload_guard_distinguishes_weight_and_kv_autoquant(monkeypatch):
hf_ptq = _import_hf_ptq(monkeypatch)
assert hf_ptq._recipe_is_kv_auto_quantize(
"general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits"
)
assert not hf_ptq._recipe_is_kv_auto_quantize("general/auto_quantize/nvfp4_fp8_at_5p4bits")
def test_autoquant_recipe_cost_excluded_layers_map_into_cost(monkeypatch):
"""Top-level cost_excluded_layers maps to the mtq constraints.cost.excluded_module_name_patterns
key (distinct from disabled_layers), so a cost-exclusion recipe matches the nested mtq dict."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
aq = load_recipe(
"model_type/qwen3_6_moe/auto_quantize/w4a16_nvfp4_fp8_at_6p0bits-active_moe"
).auto_quantize
inputs = hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
# cost-exclusion is hoisted to a sibling of disabled_layers but still reaches the mtq cost dict.
assert aq.cost_excluded_layers == ["*visual*", "*mtp*", "*vision_tower*"]
assert inputs["constraints"]["cost"] == {
"active_moe_expert_ratio": 0.03125,
"excluded_module_name_patterns": ["*visual*", "*mtp*", "*vision_tower*"],
}
# The two exclusions are independent: cost-excluded patterns are also disabled here, but the
# roles (cost-accounting vs search) are tracked separately.
assert "*visual*" in inputs["disabled_layers"]
def test_autoquant_recipe_maps_module_search_spaces(monkeypatch):
"""Fixed PTQ baseline and explicit recipe candidates map to mtq inputs."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
recipe = load_recipe(
"model_type/qwen3_6_moe/auto_quantize/w4a16_nvfp4_fp8_module_spaces_at_6p0bits-active_moe"
)
inputs = hf_ptq._mtq_inputs_from_auto_quantize_config(
recipe.auto_quantize, args, fixed_quantize_config=recipe.quantize
)
model_ptq = load_recipe("model_type/qwen3_5_moe/ptq/w4a16_nvfp4-fp8_attn-kv_fp8_cast")
assert inputs["quantization_formats"] == []
assert inputs["fixed_quantization_config"] == model_ptq.quantize.model_dump()
(searched,) = inputs["module_search_spaces"]
assert searched["module_name_patterns"] == [
"*mlp.shared_expert*",
"*linear_attn*",
"*self_attn*",
"*lm_head*",
]
assert searched["quantization_formats"] == [
QUANT_CFG_CHOICES["w4a16_nvfp4"],
QUANT_CFG_CHOICES["fp8"],
]
assert searched["allow_no_quant"] is False
def test_autoquant_rejects_non_export_safe_candidate(monkeypatch):
"""A candidate that resolves to a preset outside the export-safe set is rejected before search."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
non_safe = next(k for k in QUANT_CFG_CHOICES if k not in hf_ptq._AUTO_QUANTIZE_QFORMATS)
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=4.8),
candidate_formats=[
QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]),
QuantizeConfig(**QUANT_CFG_CHOICES[non_safe]),
],
)
with pytest.raises(ValueError, match="not supported for unified checkpoint export"):
hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
def test_autoquant_warns_on_custom_candidate(monkeypatch):
"""A candidate matching no shipped preset can't be export-verified, so it warns (not blocks)."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
custom = QuantizeConfig(quant_cfg=[{"quantizer_name": "*", "enable": False}])
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=4.8),
candidate_formats=[QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]), custom],
)
with pytest.warns(UserWarning, match="export compatibility cannot be verified"):
hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
def test_autoquant_export_guard_not_bypassed_by_effective_bits(monkeypatch):
"""A non-export-safe preset can't dodge the guard by adding a cost-only effective_bits override."""
hf_ptq, args = _parse_hf_ptq_args(
monkeypatch, "--pyt_ckpt_path", "dummy", "--kv_cache_qformat", "none"
)
non_safe = next(k for k in QUANT_CFG_CHOICES if k not in hf_ptq._AUTO_QUANTIZE_QFORMATS)
tampered = QuantizeConfig(**{**QUANT_CFG_CHOICES[non_safe], "effective_bits": 4.5})
aq = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=5.4),
candidate_formats=[QuantizeConfig(**QUANT_CFG_CHOICES["fp8"]), tampered],
)
with pytest.raises(ValueError, match="not supported for unified checkpoint export"):
hf_ptq._mtq_inputs_from_auto_quantize_config(aq, args)
def test_mlflow_flag_defaults_the_experiment_name(monkeypatch):
monkeypatch.setattr(getpass, "getuser", lambda: "tester")
hf_ptq, args = _parse_hf_ptq_args(
+88
View File
@@ -25,6 +25,7 @@ from importlib.resources import files
from pathlib import Path
import pytest
from pydantic import ValidationError
import modelopt.recipe.loader
import modelopt.torch.quantization.config as qcfg
@@ -2562,8 +2563,10 @@ def test_load_recipe_autoquantize_fixed_baseline_requires_explicit_search(tmp_pa
@pytest.mark.parametrize(
"recipe_path",
[
"general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits",
"general/auto_quantize/nvfp4_fp8_at_5p4bits",
"general/auto_quantize/nvfp4_fp8_kl_div_at_5p4bits",
"general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits",
"general/auto_quantize/kv_fp8_nvfp4_cast_kl_div_at_5p4bits",
"general/auto_quantize/nvfp4_mse_fp8_at_6p0bits",
"general/auto_quantize/w4a8_awq_beta_fp8_at_6p0bits",
@@ -2607,6 +2610,91 @@ def test_load_recipe_kv_autoquantize_contract():
assert fmt.algorithm is None
@pytest.mark.parametrize(
("recipe_path", "kv_stage"),
[
(
"general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits",
"auto_quantize",
),
(
"general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits",
"kv_auto_quantize",
),
],
)
def test_builtin_composed_kv_recipes_use_calibration_free_cast_candidates(recipe_path, kv_stage):
aq = getattr(load_recipe(recipe_path), kv_stage)
assert aq is not None
assert aq.constraints.cost_model == "kv_cache"
assert all(candidate.algorithm is None for candidate in aq.candidate_formats)
assert all(
candidate.quant_cfg[0].cfg.constant_amax == 448.0 for candidate in aq.candidate_formats
)
def _weight_autoquantize_test_config(**updates):
config = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=8.0),
candidate_formats=[qcfg.QuantizeConfig(quant_cfg=[], algorithm="max")],
)
return config.model_copy(update=updates)
def _kv_autoquantize_test_config(**updates):
config = AutoQuantizeConfig(
constraints=AutoQuantizeConstraints(effective_bits=8.0, cost_model="kv_cache"),
candidate_formats=[
qcfg.QuantizeConfig(
quant_cfg=[
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
algorithm=None,
effective_bits=8.0,
)
],
auto_quantize_method="kl_div",
)
return config.model_copy(update=updates)
def test_autoquantize_recipe_rejects_second_kv_search():
with pytest.raises(ValidationError, match=r"cannot follow.*already searches the KV cache"):
ModelOptAutoQuantizeRecipe(
auto_quantize=_kv_autoquantize_test_config(),
kv_auto_quantize=_kv_autoquantize_test_config(),
)
def test_autoquantize_recipe_rejects_non_kv_followup():
with pytest.raises(ValidationError, match="must use cost_model=kv_cache"):
ModelOptAutoQuantizeRecipe(
auto_quantize=_weight_autoquantize_test_config(),
kv_auto_quantize=_weight_autoquantize_test_config(),
)
def test_autoquantize_recipe_rejects_uniform_and_searched_kv():
uniform_kv = qcfg.QuantizeConfig(
quant_cfg=[
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
algorithm=None,
)
with pytest.raises(ValidationError, match=r"must omit.*uniform auto_quantize.kv_cache"):
ModelOptAutoQuantizeRecipe(
auto_quantize=_weight_autoquantize_test_config(kv_cache=uniform_kv),
kv_auto_quantize=_kv_autoquantize_test_config(),
)
def test_kv_autoquantize_rejects_cost_excluded_layers():
with pytest.raises(ValueError, match=r"cost_excluded_layers.*disabled_layers"):
AutoQuantizeConfig(
@@ -45,6 +45,7 @@ def test_convert_mixed_kv_cache_config_preserves_layer_map():
"kv_cache_quant_algo": "MIXED_PRECISION",
"kv_cache_quantized_layers": layer_map,
"kv_cache_schema_version": 1,
"kv_cache_deployment_supported": False,
},
}
)
@@ -55,6 +56,7 @@ def test_convert_mixed_kv_cache_config_preserves_layer_map():
assert converted["kv_cache_quant_algo"] == "MIXED_PRECISION"
assert converted["kv_cache_quantized_layers"] == layer_map
assert converted["kv_cache_schema_version"] == 1
assert converted["kv_cache_deployment_supported"] is False
def test_convert_uniform_kv_cache_config_preserves_layer_map():
@@ -389,6 +389,46 @@ def test_uniform_vlm_export_ignores_disabled_vision_attention():
assert "kv_cache_quantized_layers" not in quantization
def test_uniform_weight_quantization_exports_mixed_kv_cache_map():
model = ToyModel()
mtq.quantize(model, partial_fp8_config, lambda x: x(torch.randn(1, 4, 10)))
model.attn0 = _FakeAttention()
model.attn1 = _FakeAttention()
mtq.set_quantizer_by_cfg(
model.attn0,
[
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {"num_bits": (4, 3), "constant_amax": 1.0},
}
],
)
mtq.set_quantizer_by_cfg(
model.attn1,
[
{
"quantizer_name": "*[kv]_bmm_quantizer",
"cfg": {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
"constant_amax": 1.0,
},
}
],
)
with pytest.warns(UserWarning, match="uniform quantized weights.*mixed-precision KV-cache"):
quantization = get_quant_config(model)["quantization"]
assert quantization["quant_algo"] == "FP8"
assert quantization["kv_cache_quant_algo"] == "MIXED_PRECISION"
assert quantization["kv_cache_deployment_supported"] is False
assert quantization["kv_cache_quantized_layers"] == {
"attn0": {"quant_algo": "FP8"},
"attn1": {"quant_algo": "NVFP4"},
}
def test_quant_config_tolerates_ambiguous_language_model_roots():
model = torch.nn.Module()
model.model = torch.nn.Module()
@@ -761,7 +761,7 @@ def test_public_kv_autoquant_rejects_distributed_execution_before_mutation(monke
assert {name: type(module) for name, module in model.named_modules()} == original_types
def test_public_kv_autoquant_rejects_preceding_quantization_before_search():
def test_public_kv_autoquant_preserves_preceding_weight_quantization():
model = get_tiny_llama(num_hidden_layers=2)
model = mtq.quantize(
model,
@@ -770,7 +770,7 @@ def test_public_kv_autoquant_rejects_preceding_quantization_before_search():
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.weight_quantizer",
"cfg": {"num_bits": (4, 3), "axis": None},
"cfg": {"num_bits": (4, 3), "axis": None, "constant_amax": 1.0},
"enable": True,
},
],
@@ -788,13 +788,135 @@ def test_public_kv_autoquant_rejects_preceding_quantization_before_search():
"effective_bits": 8.0,
}
with pytest.raises(NotImplementedError, match="requires an unquantized model"):
data = [{"input_ids": torch.randint(0, model.config.vocab_size, (1, 8))}]
weight_quantizer = model.model.layers[0].self_attn.q_proj.weight_quantizer
model, _ = mtq.auto_quantize(
model,
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
data,
lambda search_model, batch: search_model(**batch).logits,
num_calib_steps=1,
num_score_steps=1,
)
assert weight_quantizer.is_enabled
assert weight_quantizer.num_bits == (4, 3)
assert weight_quantizer.amax == 1.0
assert all(
layer.self_attn.k_bmm_quantizer.is_enabled and layer.self_attn.v_bmm_quantizer.is_enabled
for layer in model.model.layers
)
def _quantized_weight_baseline(bits, *, constant_amax=1.0, axis=None):
quantizer_cfg = _quantizer_cfg(bits, constant_amax=constant_amax)
quantizer_cfg["axis"] = axis
return mtq.quantize(
get_tiny_llama(num_hidden_layers=1),
{
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*.weight_quantizer",
"cfg": quantizer_cfg,
},
],
"algorithm": None,
},
)
def test_kv_autoquant_checkpoint_rejects_changed_preceding_quantization(
tmp_path, nvfp4_fake_quant_stub
):
candidate = _kv_config((4, 3), 8.0, algorithm=None, constant_amax=1.0).model_dump()
data = [{"input_ids": torch.randint(0, 16, (1, 8))}]
checkpoint = str(tmp_path / "kv_search.pth")
mtq.auto_quantize(
_quantized_weight_baseline((4, 3)),
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
data,
lambda model, batch: model(**batch).logits,
num_calib_steps=1,
num_score_steps=1,
checkpoint=checkpoint,
)
mtq.auto_quantize(
_quantized_weight_baseline((4, 3)),
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
data,
lambda *_: pytest.fail("An identical preceding state must restore without rescoring."),
num_calib_steps=1,
num_score_steps=1,
checkpoint=checkpoint,
)
with pytest.raises(
ValueError, match=r"preceding non-K/V quantizer.*recompute KV sensitivities"
):
mtq.auto_quantize(
model,
_quantized_weight_baseline((2, 1)),
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
[],
lambda *_: pytest.fail("Validation must fail before search."),
data,
lambda *_: pytest.fail("A stale checkpoint must be rejected before scoring."),
num_calib_steps=1,
num_score_steps=1,
checkpoint=checkpoint,
)
@pytest.mark.parametrize(
("first_kwargs", "second_kwargs", "mutate_second_amax"),
[
({"constant_amax": 1.0}, {"constant_amax": 2.0}, False),
(
{"constant_amax": 1.0, "axis": None},
{"constant_amax": 1.0, "axis": 0},
False,
),
({"constant_amax": 1.0}, {"constant_amax": 1.0}, True),
],
ids=("constant-amax", "axis", "calibrated-amax"),
)
def test_kv_autoquant_checkpoint_rejects_changed_preceding_state(
tmp_path, first_kwargs, second_kwargs, mutate_second_amax
):
candidate = _kv_config((4, 3), 8.0, algorithm=None, constant_amax=1.0).model_dump()
data = [{"input_ids": torch.randint(0, 16, (1, 8))}]
checkpoint = str(tmp_path / "kv_search.pth")
mtq.auto_quantize(
_quantized_weight_baseline((4, 3), **first_kwargs),
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
data,
lambda model, batch: model(**batch).logits,
num_calib_steps=1,
num_score_steps=1,
checkpoint=checkpoint,
)
second_model = _quantized_weight_baseline((4, 3), **second_kwargs)
if mutate_second_amax:
for name, module in second_model.named_modules():
if name.endswith("weight_quantizer") and module.is_enabled:
module.amax = module.amax * 2
with pytest.raises(ValueError, match="preceding non-K/V quantizer"):
mtq.auto_quantize(
second_model,
{"effective_bits": 8.0, "cost_model": "kv_cache"},
[candidate],
data,
lambda *_: pytest.fail("A stale checkpoint must be rejected before scoring."),
num_calib_steps=1,
num_score_steps=1,
checkpoint=checkpoint,
)