Repository navigation
[OMNIML-5858] Add a 4.9-bit fixed-MTP AutoQuant recipe #2582
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
91e1a31
7d1a095
1930f16
db8765b
696117b
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,144 @@ | ||
| # 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. | ||
|
|
||
| # modelopt-schema: modelopt.recipe.config.ModelOptAutoQuantizeRecipe | ||
| imports: | ||
| base_cost_excluded_layers: configs/auto_quantize/units/base_cost_excluded_layers | ||
| nvfp4: configs/numerics/nvfp4 | ||
| nvfp4_static: configs/numerics/nvfp4_static | ||
| fp8: configs/numerics/fp8 | ||
| kv_fp8_cast: configs/ptq/units/kv_fp8_cast | ||
|
|
||
| metadata: | ||
| recipe_type: auto_quantize | ||
| description: >- | ||
| Nemotron-H AutoQuantize at 4.9 effective bits over weight-MSE NVFP4, FP8, and BF16, | ||
| with MSE-calibrated MTP routed/shared experts fixed at NVFP4 W4A4, calibrated MTP FP8 KV, | ||
| and base-model FP8 KV casting. Other MTP modules and vision/audio remain in BF16. | ||
|
|
||
| quantize: | ||
| algorithm: &mse | ||
| method: mse | ||
| fp8_scale_sweep: true | ||
| quant_cfg: | ||
| - quantizer_name: "*" | ||
| enable: false | ||
| - quantizer_name: "*mtp.layers.1.mixer.experts.*weight_quantizer" | ||
| enable: true | ||
| cfg: &nvfp4_weight | ||
| $import: nvfp4_static | ||
| # Four value bits plus one FP8 scale per 16-element block. | ||
| effective_bits: 4.5 | ||
| - quantizer_name: "*mtp.layers.1.mixer.experts.*input_quantizer" | ||
| enable: true | ||
| cfg: {$import: nvfp4} | ||
| - quantizer_name: "*mtp.layers.1.mixer.shared_experts.*weight_quantizer" | ||
| enable: true | ||
| cfg: *nvfp4_weight | ||
| - quantizer_name: "*mtp.layers.1.mixer.shared_experts.*input_quantizer" | ||
| enable: true | ||
| cfg: {$import: nvfp4} | ||
| # Collect MTP KV scales in the fixed calibration pass, before base-model KV casting. | ||
| - quantizer_name: "*mtp.*[kv]_bmm_quantizer" | ||
| enable: true | ||
| cfg: &mtp_kv_fp8 | ||
| $import: fp8 | ||
| use_constant_amax: false | ||
|
|
||
| auto_quantize: | ||
| constraints: | ||
| effective_bits: 4.9 | ||
| module_search_spaces: | ||
| - module_name_patterns: | ||
| - "model.layers.*" | ||
| - "backbone.layers.*" | ||
| - "*language_model.layers*" | ||
| - "*language_model.model.layers*" | ||
| - "*language_model.backbone.layers*" | ||
| candidate_formats: &search_candidates | ||
| - algorithm: *mse | ||
| quant_cfg: | ||
| - quantizer_name: "*" | ||
| enable: false | ||
| - quantizer_name: "*weight_quantizer" | ||
| enable: true | ||
| cfg: *nvfp4_weight | ||
| - quantizer_name: "*input_quantizer" | ||
| enable: true | ||
| cfg: {$import: nvfp4} | ||
| - algorithm: max | ||
| quant_cfg: | ||
| - quantizer_name: "*" | ||
| enable: false | ||
| - quantizer_name: "*weight_quantizer" | ||
| enable: true | ||
| cfg: {$import: fp8} | ||
| - quantizer_name: "*input_quantizer" | ||
| enable: true | ||
| cfg: {$import: fp8} | ||
| allow_no_quant: true | ||
| - module_name_patterns: | ||
| - "*lm_head*" | ||
| - "*output_layer*" | ||
| candidate_formats: *search_candidates | ||
| allow_no_quant: true | ||
| auto_quantize_method: gradient | ||
| score_size: 128 | ||
| kv_cache: | ||
| quant_cfg: | ||
| - $import: kv_fp8_cast | ||
| - quantizer_name: "*mtp.*[kv]_bmm_quantizer" | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. hi @meenchen - the shipped recipe has kv scales as fp8 calibrated scales for MTP as we only dropped the scales in backbone- https://huggingface.co/nvidia/Nemotron-Super-3.5-GA-FINAL-row105-QAD-PreStitched-BoostedMTP/blob/main/model.safetensors.index.json
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Fixed in |
||
| enable: true | ||
| cfg: *mtp_kv_fp8 | ||
| disabled_layers: | ||
| - "*block_sparse_moe.gate*" | ||
| - "*linear_attn.conv1d*" | ||
| - "*linear_attn.in_proj_a*" | ||
| - "*linear_attn.in_proj_b*" | ||
| - "*mixer.conv1d*" | ||
| - "*mlp.gate.*" | ||
| - "*mlp.shared_expert_gate.*" | ||
| - "*shared_expert_gate*" | ||
| - "*proj_out.*" | ||
| - "*router*" | ||
| - "*vision*" | ||
| - "*visual*" | ||
| - "*image*" | ||
| - "*radio*" | ||
| - "*encoder*" | ||
| - "*model_encoder*" | ||
| - "*video*" | ||
| - "*audio*" | ||
| - "*speech*" | ||
| - "*multimodal_projector*" | ||
| - "*multi_modal_projector*" | ||
| - "*vision_projector*" | ||
| - "*audio_projector*" | ||
| cost_excluded_layers: | ||
| - $import: base_cost_excluded_layers | ||
| - "*vision*" | ||
| - "*visual*" | ||
| - "*image*" | ||
| - "*radio*" | ||
| - "*encoder*" | ||
| - "*model_encoder*" | ||
| - "*video*" | ||
| - "*audio*" | ||
| - "*speech*" | ||
| - "*multimodal_projector*" | ||
| - "*multi_modal_projector*" | ||
| - "*vision_projector*" | ||
| - "*audio_projector*" | ||
| - "*mtp*" | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,221 @@ | ||
| # 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. | ||
|
|
||
| """Resolve the fixed-MTP recipe on native modules without running GPU calibration.""" | ||
|
|
||
| import pytest | ||
| import torch | ||
|
|
||
| native = pytest.importorskip("transformers.models.nemotron_h.modeling_nemotron_h") | ||
|
|
||
| import modelopt.torch.quantization as mtq | ||
| from modelopt.recipe import load_recipe | ||
| from modelopt.torch.export.quant_utils import postprocess_state_dict | ||
| from modelopt.torch.models.nemotron_h.mtp import _NemotronHMTP, prepare_for_calibration | ||
| from modelopt.torch.opt.conversion import apply_mode | ||
| from modelopt.torch.opt.utils import named_hparams | ||
| from modelopt.torch.quantization._auto_quantize_cost import get_auto_quantize_cost_model | ||
| from modelopt.torch.quantization.algorithms import AutoQuantizeGradientSearcher, QuantRecipe | ||
| from modelopt.torch.quantization.mode import QuantizeModeRegistry | ||
| from modelopt.torch.quantization.model_calib import max_calibrate | ||
| from modelopt.torch.quantization.nn import TensorQuantizer | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("wrapped", [False, True], ids=["standalone", "wrapped"]) | ||
| @pytest.mark.parametrize("decoder_name", ["model", "backbone"]) | ||
| @pytest.mark.parametrize("with_mtp", [False, True], ids=["no-mtp", "mtp"]) | ||
| def test_4p9_mse_recipe_resolves_search_and_fixed_mtp(wrapped, decoder_name, with_mtp): | ||
| """Search decoder/head groups, fix MTP experts, and preserve calibrated MTP KV scales.""" | ||
| config = native.NemotronHConfig( | ||
| vocab_size=32, | ||
| hidden_size=32, | ||
| intermediate_size=32, | ||
| layers_block_type=["attention", "moe"], | ||
| num_attention_heads=2, | ||
| num_key_value_heads=1, | ||
| head_dim=16, | ||
| n_routed_experts=2, | ||
| num_experts_per_tok=1, | ||
| moe_intermediate_size=32, | ||
| moe_shared_expert_intermediate_size=32, | ||
| use_mamba_kernels=False, | ||
| attn_implementation="eager", | ||
| mtp_layers_block_type=["attention", "moe"], | ||
| ) | ||
| language_model = native.NemotronHForCausalLM(config).eval() | ||
| # Keep real decoder modules while exercising both native and remote-code namespaces. | ||
| current_name = "model" if "model" in language_model._modules else "backbone" | ||
| if current_name != decoder_name: | ||
| language_model.add_module(decoder_name, language_model._modules.pop(current_name)) | ||
| if with_mtp: | ||
| language_model.mtp = _NemotronHMTP(config) | ||
| model = language_model | ||
| prefix = "" | ||
| if wrapped: | ||
| model = torch.nn.Module() | ||
| model.language_model = language_model | ||
| model.vision_model = torch.nn.Linear(32, 32) | ||
| prefix = "language_model." | ||
|
|
||
| recipe = load_recipe("model_type/nemotron_h/auto_quantize/nvfp4_mse_fp8_mtp_fixed_at_4p9bits") | ||
| aq = recipe.auto_quantize | ||
| fixed = QuantRecipe(recipe.quantize.model_dump(), name="fixed_mtp") | ||
| assert fixed.config.algorithm == {"method": "mse", "fp8_scale_sweep": True} | ||
| assert aq.auto_quantize_method == "gradient" | ||
| assert aq.constraints.effective_bits == 4.9 | ||
|
|
||
| model = apply_mode(model, mode="auto_quantize", registry=QuantizeModeRegistry) | ||
| mtq.set_quantizer_by_cfg(model, fixed.config.quant_cfg) | ||
| searcher = AutoQuantizeGradientSearcher() | ||
| searcher.model = model | ||
| searcher.constraints = aq.constraints.model_dump(exclude_none=True) | ||
| searcher.config = {"cost": {"excluded_module_name_patterns": aq.cost_excluded_layers}} | ||
| searcher._cost_model = get_auto_quantize_cost_model("weight") | ||
| spaces = searcher._normalize_module_search_spaces( | ||
| [ | ||
| { | ||
| "module_name_patterns": space.module_name_patterns, | ||
| "quantization_formats": [ | ||
| (fmt.model_dump(), f"candidate_{index}") | ||
| for index, fmt in enumerate(space.candidate_formats) | ||
| ], | ||
| "allow_no_quant": space.allow_no_quant, | ||
| } | ||
| for space in aq.module_search_spaces | ||
| ] | ||
| ) | ||
| searcher.insert_hparams_after_merge_rules(model, [], aq.disabled_layers, spaces, fixed) | ||
| groups = [hparam for _, hparam in named_hparams(model, unique=True)] | ||
| searcher._verify_resolved_constraint(groups) | ||
|
|
||
| searched_names, fixed_expert_names, bf16_mtp_names = set(), set(), set() | ||
| for group in groups: | ||
| names = group.quant_module_names | ||
| bits = [choice.compression * 16 for choice in group.solver_choices] | ||
| if any(".mixer.router" in name or name.startswith("vision_model") for name in names): | ||
| assert bits == [16] | ||
| elif all( | ||
| name.startswith(prefix + decoder_name + ".layers.") for name in names | ||
| ) or names == [prefix + "lm_head"]: | ||
| assert bits == [4.5, 8, 16] | ||
| assert group.allow_no_quant and not group.is_fixed | ||
| assert group.cost_weight == 1 | ||
| assert group.solver_choices[0].config.algorithm == fixed.config.algorithm | ||
| searched_names.update(names) | ||
| elif all( | ||
| "mtp.layers.1.mixer.experts" in name or "mtp.layers.1.mixer.shared_experts." in name | ||
| for name in names | ||
| ): | ||
| assert bits == [4.5] | ||
| assert group.is_fixed and not group.allow_no_quant | ||
| assert group.cost_weight == 0 | ||
| fixed_expert_names.update(names) | ||
| for module in group.quant_modules: | ||
| quantizers = [ | ||
| (name, q) | ||
| for name, q in module.named_modules() | ||
| if isinstance(q, TensorQuantizer) and ("weight" in name or "input" in name) | ||
| ] | ||
| assert quantizers | ||
| for name, quantizer in quantizers: | ||
| assert quantizer.is_enabled and quantizer.num_bits == (2, 1) | ||
| assert quantizer.block_sizes[-1] == 16 | ||
| assert quantizer.block_sizes["type"] == ( | ||
| "static" if "weight" in name else "dynamic" | ||
| ) | ||
| else: | ||
| assert all(name.startswith(prefix + "mtp.") for name in names) | ||
| assert bits == [16] and group.is_fixed and group.cost_weight == 0 | ||
| bf16_mtp_names.update(names) | ||
|
|
||
| assert prefix + "lm_head" in searched_names | ||
| assert any(".layers.0.mixer.q_proj" in name for name in searched_names) | ||
| assert any(".layers.1.mixer.experts" in name for name in searched_names) | ||
| assert bool(fixed_expert_names) == with_mtp | ||
| if with_mtp: | ||
| assert any(".mixer.experts" in name for name in fixed_expert_names) | ||
| for projection in ("up_proj", "down_proj"): | ||
| assert prefix + f"mtp.layers.1.mixer.shared_experts.{projection}" in fixed_expert_names | ||
| assert bf16_mtp_names | ||
|
|
||
| calibrated_amax = {} | ||
| # Exercise actual KV activation collection without CPU NVFP4 weight calibration. | ||
| if with_mtp and decoder_name == "model": | ||
| prepare_for_calibration(model) | ||
| with mtq.set_quantizer_by_cfg_context( | ||
| model, | ||
| [ | ||
| {"quantizer_name": "*weight_quantizer*", "enable": False}, | ||
| {"quantizer_name": "*input_quantizer", "enable": False}, | ||
| ], | ||
| ): | ||
| max_calibrate( | ||
| model, | ||
| lambda _: language_model(torch.tensor([[1, 2, 3, 4]]), use_cache=False), | ||
| distributed_sync=False, | ||
| ) | ||
| calibrated_amax = { | ||
| name: quantizer.amax.clone() | ||
| for name, quantizer in model.named_modules() | ||
| if "mtp." in name and name.endswith(("k_bmm_quantizer", "v_bmm_quantizer")) | ||
| } | ||
| assert len(calibrated_amax) == 2 | ||
| for amax in calibrated_amax.values(): | ||
| assert torch.isfinite(amax).all() and amax.max() > 0 | ||
|
|
||
| # Post-search casting must retain the fixed MTP quantizers and their collected scales. | ||
| for stage in ("baseline", "post-search"): | ||
| if stage == "post-search": | ||
| mtq.set_quantizer_by_cfg(model, aq.kv_cache.quant_cfg) | ||
| for name, quantizer in model.named_modules(): | ||
| if isinstance(quantizer, TensorQuantizer) and name.startswith(prefix + "mtp."): | ||
| expert_io = name.startswith( | ||
| ( | ||
| prefix + "mtp.layers.1.mixer.experts.", | ||
| prefix + "mtp.layers.1.mixer.shared_experts.", | ||
| ) | ||
| ) and ("weight_quantizer" in name or "input_quantizer" in name) | ||
| mtp_kv = name.endswith(("k_bmm_quantizer", "v_bmm_quantizer")) | ||
| assert quantizer.is_enabled == (expert_io or mtp_kv) | ||
| kv = { | ||
| name: q | ||
| for name, q in model.named_modules() | ||
| if name.endswith(("k_bmm_quantizer", "v_bmm_quantizer")) | ||
| } | ||
| assert len(kv) == 2 + 2 * with_mtp | ||
| for name, quantizer in kv.items(): | ||
| is_mtp = "mtp." in name | ||
| enabled = stage == "post-search" or is_mtp | ||
| assert quantizer.is_enabled == enabled | ||
| if enabled: | ||
| assert quantizer.num_bits == (4, 3) | ||
| assert quantizer._use_constant_amax == (not is_mtp) | ||
| if not is_mtp: | ||
| assert quantizer._get_amax(torch.tensor([0.01, 1000.0])).item() == 448 | ||
| assert not hasattr(quantizer, "_amax") | ||
| elif name in calibrated_amax: | ||
| torch.testing.assert_close(quantizer.amax, calibrated_amax[name]) | ||
|
|
||
| if calibrated_amax: | ||
| exported = postprocess_state_dict(model.state_dict(), maxbound=448, quantization="FP8") | ||
| for name, amax in calibrated_amax.items(): | ||
| projection = "k_proj.k_scale" if name.endswith("k_bmm_quantizer") else "v_proj.v_scale" | ||
| torch.testing.assert_close( | ||
| exported[name.rsplit(".", 1)[0] + "." + projection], amax / 448 | ||
| ) | ||
| assert not any( | ||
| name.startswith(prefix + decoder_name + ".") and name.endswith(("k_scale", "v_scale")) | ||
| for name in exported | ||
| ) |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
should we have the shared experts in nvfp4 too?
@yeyu-nvidia please help confirm
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Updated in
696117b8f: MTP shared-expertup_proj/down_projnow joins routed experts in the fixed NVFP4 W4A4 MSE baseline. Runtime regressions verify both shared projections use static block-16 NVFP4 weights and NVFP4 activations, remain fixed, and are excluded from the search cost. Routers/gates and other MTP modules remain BF16. This implements the requested candidate policy; speculative acceptance quality has not yet been measured.