Skip to content

Refine static NVFP4 MSE calibration - #1536

Merged
realAsma merged 3 commits into
mainfrom
asma/mse_cleanups
May 29, 2026
Merged

realAsma merged 3 commits into
mainfrom
asma/mse_cleanups

Conversation

@realAsma

@realAsma realAsma commented May 22, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

Type of change: Bug fix

Refines static NVFP4 MSE calibration and forces static NVFP4 amax state to stay FP32 across calibration loading, quantizer promotion, dtype casts, and restore paths.

Main changes:

  • Tighten max/MSE calibration bootstrap and static NVFP4 quantizer promotion.
  • Keep static NVFP4 _amax and _global_amax in FP32.
  • Update focused GPU/unit coverage for FP8 sweep calibration, promotion, restore, and FP32 amax preservation.

Usage

algorithm:
  method: mse
  fp8_scale_sweep: true

Testing

pre-commit run --files modelopt/torch/quantization/nn/modules/tensor_quantizer.py tests/gpu/torch/quantization/test_nvfp4_static_quantizer_cuda.py
pytest_pwd tests/gpu/torch/quantization/test_nvfp4_static_quantizer_cuda.py -q

Before your PR is "Ready for review"

Make sure you read and follow Contributor guidelines and your commits are signed (git commit -s -S).

Make sure you read and follow the Security Best Practices.

  • Is this change backward compatible?: Yes
  • 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?: Yes
  • Did you update Changelog?: N/A
  • Did you get Claude approval on this PR?: Yes

Additional Information

N/A

@realAsma
realAsma requested review from a team as code owners May 22, 2026 17:25
@coderabbitai

coderabbitai Bot commented May 22, 2026 •

Copy link
Copy Markdown
Contributor

Note

Reviews paused

It looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the reviews.auto_review.auto_pause_after_reviewed_commits setting.

Use the following commands to manage reviews:

  • @coderabbitai resume to resume automatic reviews.
  • @coderabbitai review to trigger a single review.

Use the checkboxes below for quick actions:

  • ▶️ Resume reviews
  • 🔍 Trigger review
📝 Walkthrough

Walkthrough

This PR refactors NVFP4 MSE calibration to cache final per-block amax immediately as float32 in a one-shot cycle, centralizes max-stat collection and per-weight MSE calibrator dispatch, tightens NVFP4-static promotion and global-amax sync, updates TensorQuantizer amax buffer handling, and expands tests for dtype and one-shot semantics.

Changes

NVFP4 MSE Calibration and Static Promotion Refactor

Layer / File(s) Summary
NVFP4MSECalibrator one-shot fp32 caching and validation
modelopt/torch/quantization/calib/mse.py
NVFP4MSECalibrator documents one-shot Triton sweep; _compute_candidate_amax normalizes candidates/global_amax to float32; collect() enforces single-cycle semantics via _best_amax_fast and routes to Triton fast-path or reference helper; compute_amax() returns cached fp32 result or None.
MSE config, calibrator dispatch, and per-weight factory
modelopt/torch/quantization/config.py, modelopt/torch/quantization/model_calib.py
MseCalibConfig docs corrected; new gating helper and _make_weight_mse_calibrator() centralize per-weight eligibility and calibrator selection (registered FP8 sweep, NVFP4MSECalibrator, or MseCalibrator); mse_calibrate() refactored to run max_calibrate then per-quantizer MSE refinement.
Max calibration refactor with centralized stats collection
modelopt/torch/quantization/model_calib.py
Adds _run_and_load_max_stats() to unify enable-stats → run → finish-load; tightens _is_calibrated_nvfp4_static; bootstraps eligible uncalibrated static weight quantizers via weight-only stat re-runs; max_calibrate reworked to use centralized lifecycle.
Max calibrate ordering, bootstrap, and gptq removal
modelopt/torch/quantization/model_calib.py
Reorders max_calibrate to run core lifecycle, performs local MoE expert amax sync, replaces unconditional promotion with wrapper that runs promote+group preprocessing, and removes initial promote_nvfp4_static_quantizers call from gptq and post-max hook.
TensorQuantizer load_calib_amax buffer handling
modelopt/torch/quantization/nn/modules/tensor_quantizer.py
load_calib_amax always clones/detaches and replaces internal _amax buffer aligned to target device; validates shape when existing amax exists; removes in-place .data.copy and register_buffer path.
NVFP4 static promotion and core utils update
modelopt/torch/quantization/utils/core_utils.py
promote_nvfp4_static_quantizers now promotes enabled TensorQuantizer instances marked is_nvfp4_static with non-None amax, computes global_amax from module.amax (clone/detach + reduce_amax), and constructs NVFP4StaticQuantizer.from_tensor_quantizer while counting only new conversions.
NVFP4StaticQuantizer and NVFP4MSECalibrator CUDA tests
tests/gpu/torch/quantization/test_nvfp4_static_quantizer_cuda.py
Adds CUDA test verifying load_calib_amax preserves fp32 dtype when existing amax buffer lower-precision; updates NVFP4MSECalibrator tests to assert compute_amax() returns None before collect finalization, expects reference-path to cache _best_amax_fast and leave _losses_sum None, and replaces multi-collection test with reset-based recollect scenario.
NVFP4MSECalibrator fp32 amax dtype and fast-path caching tests
tests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.py
Adds Triton-gated test to verify per-block amax stored as float32 and output dtype preserved; strengthens reset/recollect and dispatch tests to assert _best_amax_fast set and _losses_sum None; end-to-end mse test extended to assert NVFP4 static quantizer amax dtypes are torch.float32 across runs.
Test infra, fused-experts dead-expert validation, and mse dispatch tests
tests/gpu/torch/quantization/test_gptq.py, tests/unit/torch/quantization/plugins/test_fused_experts.py, tests/unit/torch/quantization/test_mse_calibrator.py
test_gptq.py imports updated to use calib_utils and no longer calls promote utility in test setup; fused-experts dead-expert test refactored to validate max_calibrate populates dead static NVFP4 quantizers and weight config updated to static NVFP4 style; test_mse_calibrator imports/dispatch tests extended and new TestStaticNVFP4Promotion class added to validate promotion and grouped global_amax sync.
Example and docstring updates
examples/llm_ptq/cast_mxfp4_to_nvfp4.py
Minor docstring wording change to reference static NVFP4 finalization as the pickup mechanism for forced block_sizes['type']='static' entries.

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~60 minutes

Suggested reviewers

  • meenchen
  • cjluo-nv
  • Edwardf0t1
🚥 Pre-merge checks | ✅ 6
✅ Passed checks (6 passed)
Check name Status Explanation
Title check ✅ Passed The title "Refine static NVFP4 MSE calibration" directly and specifically describes the primary purpose of the pull request: refining the MSE calibration behavior for static NVFP4 quantizers.
Docstring Coverage ✅ Passed Docstring coverage is 84.91% which is sufficient. The required threshold is 80.00%.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
Security Anti-Patterns ✅ Passed All 11 modified files passed security review. No unsafe torch.load, numpy.load, trust_remote_code, eval/exec, nosec comments, or restricted-license dependencies found.
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests
  • Commit unit tests in branch asma/mse_cleanups

Comment @coderabbitai help to get the list of available commands and usage tips.

@github-actions

github-actions Bot commented May 22, 2026 •

Copy link
Copy Markdown
Contributor
PR Preview Action v1.8.1
Preview removed because the pull request was closed.
2026-05-29 21:04 UTC

Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py Outdated
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py Outdated

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Warning

CodeRabbit couldn't request changes on this pull request because it doesn't have sufficient GitHub permissions.

Please grant CodeRabbit Pull requests: Read and write permission and re-run the review.

👉 Steps to fix this

Actionable comments posted: 3

🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Inline comments:
In `@modelopt/torch/quantization/model_calib.py`:
- Around line 214-221: The stats lifecycle in _run_and_load_max_stats is not
guarded: call enable_stats_collection(model) then run the forward path (either
weight_only_quantize(model) or forward_loop(model)) inside a try block and call
finish_stats_collection(model) in a finally block so finish_stats_collection
always executes even if the forward path raises; re-raise any caught exception
after the finally to preserve behavior. Reference functions:
_run_and_load_max_stats, enable_stats_collection, weight_only_quantize,
forward_loop, finish_stats_collection.

In `@tests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.py`:
- Line 295: The local import of TensorQuantizer inside the
test_mse_calibrate_end_to_end function should be moved to module scope: remove
the in-function import and add "from modelopt.torch.quantization.nn import
TensorQuantizer" to the top of the test file with the other imports so import
failures surface at collection time; update any references in
test_mse_calibrate_end_to_end to use the now-module-level TensorQuantizer and
ensure there is no justification comment left for an inside-function import.

In `@tests/unit/torch/quantization/test_mse_calibrator.py`:
- Around line 686-700: Move the in-test imports of
_promote_nvfp4_static_quantizers_with_global_amax_sync out of the individual
test methods and place them in the module-level import block (i.e., import
_promote_nvfp4_static_quantizers_with_global_amax_sync from
modelopt.torch.quantization.model_calib at the top of the test file) so tests
follow the guideline that imports belong at file scope; only keep them inside a
test if there is a documented circular/optional dependency reason.
🪄 Autofix (Beta)

Fix all unresolved CodeRabbit comments on this PR:

  • Push a commit to this branch (recommended)
  • Create a new PR with the fixes

ℹ️ Review info
⚙️ Run configuration

Configuration used: Path: .coderabbit.yaml

Review profile: CHILL

Plan: Enterprise

Run ID: 39296e29-e4e1-4049-9048-513405b3ee9d

📥 Commits

Reviewing files that changed from the base of the PR and between 3ff15cc and a263e63.

📒 Files selected for processing (12)
  • examples/llm_ptq/cast_mxfp4_to_nvfp4.py
  • modelopt/torch/quantization/calib/mse.py
  • modelopt/torch/quantization/config.py
  • modelopt/torch/quantization/model_calib.py
  • modelopt/torch/quantization/nn/modules/tensor_quantizer.py
  • modelopt/torch/quantization/utils/core_utils.py
  • modelopt_recipes/configs/ptq/presets/model/nvfp4_w4a4_weight_mse_fp8_sweep.yaml
  • tests/gpu/torch/quantization/test_gptq.py
  • tests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.py
  • tests/gpu/torch/quantization/test_nvfp4_static_quantizer_cuda.py
  • tests/unit/torch/quantization/plugins/test_fused_experts.py
  • tests/unit/torch/quantization/test_mse_calibrator.py
💤 Files with no reviewable changes (2)
  • tests/gpu/torch/quantization/test_gptq.py
  • modelopt/torch/quantization/utils/core_utils.py

Comment thread modelopt/torch/quantization/model_calib.py
Comment thread tests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.py Outdated
Comment thread tests/unit/torch/quantization/test_mse_calibrator.py Outdated
Comment thread modelopt/torch/quantization/model_calib.py Outdated
@codecov

codecov Bot commented May 22, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 99.07407% with 1 line in your changes missing coverage. Please review.
✅ Project coverage is 76.98%. Comparing base (f99d83e) to head (5aa79a4).
⚠️ Report is 2 commits behind head on main.

Files with missing lines Patch % Lines
modelopt/torch/quantization/utils/core_utils.py 91.66% 1 Missing ⚠️
Additional details and impacted files
@@            Coverage Diff             @@
##             main    #1536      +/-   ##
==========================================
+ Coverage   76.67%   76.98%   +0.30%     
==========================================
  Files         478      478              
  Lines       52393    52408      +15     
==========================================
+ Hits        40174    40347     +173     
+ Misses      12219    12061     -158     
Flag Coverage Δ
examples 41.64% <34.25%> (+8.77%) ⬆️
gpu 59.60% <97.22%> (-0.47%) ⬇️
regression 15.18% <11.11%> (-0.01%) ⬇️
unit 53.60% <77.77%> (+0.08%) ⬆️

Flags with carried forward coverage won't be shown. Click here to find out more.

☔ View full report in Codecov by Sentry.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Regarding CodeRabbit’s stats-lifecycle suggestion at #1536 (comment): I am going to leave this as-is. The helper currently has a simple, linear stats lifecycle, and adding a try/finally here makes the common path harder to read without addressing a demonstrated failure in this PR. We can revisit this if we see stale stats after a failed calibration run.

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Caution

Some comments are outside the diff and can’t be posted inline due to platform limitations.

⚠️ Outside diff range comments (1)
modelopt/torch/quantization/model_calib.py (1)

523-536: ⚠️ Potential issue | 🟠 Major | ⚡ Quick win

Restore the original calibrator after the temporary MSE pass.

This loop leaves weight_quantizer._calibrator pointing at cal, then immediately resets it. For the base MseCalibrator, reset() clears _initial_amax, so a later calibration pass on the same model can hit a broken calibrator state instead of the original one.

Suggested fix
                 if cal is None:
                     continue
-                weight_quantizer._calibrator = cal
-                _run_and_load_max_stats(weight_quantizer, lambda q: q(weight))
-                if hasattr(cal, "reset"):
-                    cal.reset()
+                original_calibrator = weight_quantizer._calibrator
+                weight_quantizer._calibrator = cal
+                try:
+                    _run_and_load_max_stats(weight_quantizer, lambda q: q(weight))
+                finally:
+                    weight_quantizer._calibrator = original_calibrator
+                    if hasattr(cal, "reset"):
+                        cal.reset()
 
                 pbar.update(1)
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@modelopt/torch/quantization/model_calib.py` around lines 523 - 536, The loop
temporarily replaces weight_quantizer._calibrator with a new MSE calibrator from
_make_weight_mse_calibrator, runs _run_and_load_max_stats, then calls
cal.reset(), but never restores the original calibrator; save the original (orig
= weight_quantizer._calibrator) before assigning the temporary cal, run
_run_and_load_max_stats with the temp calibrator, call cal.reset() if needed,
and finally restore weight_quantizer._calibrator = orig so the original
calibrator state is preserved for subsequent passes (references:
weight_quantizer._calibrator, _make_weight_mse_calibrator,
_run_and_load_max_stats, cal.reset()).
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Outside diff comments:
In `@modelopt/torch/quantization/model_calib.py`:
- Around line 523-536: The loop temporarily replaces
weight_quantizer._calibrator with a new MSE calibrator from
_make_weight_mse_calibrator, runs _run_and_load_max_stats, then calls
cal.reset(), but never restores the original calibrator; save the original (orig
= weight_quantizer._calibrator) before assigning the temporary cal, run
_run_and_load_max_stats with the temp calibrator, call cal.reset() if needed,
and finally restore weight_quantizer._calibrator = orig so the original
calibrator state is preserved for subsequent passes (references:
weight_quantizer._calibrator, _make_weight_mse_calibrator,
_run_and_load_max_stats, cal.reset()).

ℹ️ Review info
⚙️ Run configuration

Configuration used: Path: .coderabbit.yaml

Review profile: CHILL

Plan: Enterprise

Run ID: 38c07283-0685-4897-9033-c82711e6c56c

📥 Commits

Reviewing files that changed from the base of the PR and between a263e63 and e980773.

📒 Files selected for processing (5)
  • modelopt/torch/quantization/calib/mse.py
  • modelopt/torch/quantization/model_calib.py
  • modelopt/torch/quantization/nn/modules/tensor_quantizer.py
  • tests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.py
  • tests/unit/torch/quantization/test_mse_calibrator.py

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Auto-replying to realAsma review comments:

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

/claude review

@cjluo-nv cjluo-nv left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Bot review — DM the bot to share feedback.

The actual diff goes well beyond the stated "sync on CUDA + remove _wait_for_cuda_amax" scope: it adds a new public-ish config field apply_mse_nvfp_static_only, removes promote_nvfp4_static_quantizers from utils/core_utils.py, restructures mse_calibrate, narrows _bootstrap_uncalibrated_weight_quantizers to NVFP4-static only (renamed accordingly), and changes load_calib_amax so that an existing _amax buffer is always replaced (inheriting calib_amax's dtype) instead of copy_'d in place (preserving the buffer's dtype). That last change is model-wide, not NVFP4-specific, and the PR body doesn't call it out — worth a human pass to confirm no downstream caller relies on the buffer dtype being preserved (e.g. an fp16 _amax registered earlier silently becomes fp32). Tests look reasonable: new fp32-amax preservation test, fast-path-cache one-shot test, and a static-only-skip test for the registry. Couple of smaller items also worth eyeballing:

  • _promote_nvfp4_static_quantizers_with_global_amax_sync is now called inside max_calibrate, so the bootstrap (formerly only triggered from mse_calibrate) now runs for every max_calibrate consumer (gptq, awq postprocess, smoothquant, etc.). Intended? The _check_moe_calibration_complete path runs after promotion — fine since the check looks at _amax on input quantizers too, but worth confirming nothing now warns where it didn't before.
  • local_hessian_calibrate lost its explicit _sync_grouped_weight_global_amax(model) call. It's now subsumed by max_calibrate → promotion path, but local_hessian also has its own per-module promotion later. Double-promotion is a no-op (covered by from_tensor_quantizer), just noting the redundancy.
  • _compute_candidate_amax now multiplies by torch.ones_like(self._initial_amax, dtype=torch.float32) — the ones_like is purely for broadcasting and a plain (self._global_amax.to(torch.float32) * candidates).expand_as(self._initial_amax) (or simpler, just relying on broadcast) would be cheaper; minor.

Comment thread modelopt_recipes/configs/ptq/presets/model/nvfp4_w4a4_weight_mse_fp8_sweep.yaml Outdated

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude review passed — no blocking issues found. LGTM

Findings: 0 CRITICAL, 0 IMPORTANT, 2 SUGGESTIONS

The refactor is internally consistent:

  • NVFP4MSECalibrator becoming one-shot is fine — it's only constructed for weights in _make_weight_mse_calibrator, never for activations.
  • Bootstrap scope narrowing (_bootstrap_uncalibrated_static_weight_quantizers) is intentional and reflected in the test rename; all in-tree MSE recipes are NVFP4-only.
  • load_calib_amax change to preserve calibrator dtype on the buffer is covered by the new test_load_calib_amax_preserves_fp32_result_dtype regression test.
  • _promote_nvfp4_static_quantizers_with_global_amax_sync is idempotent (from_tensor_quantizer is a no-op for already-promoted modules) and preprocess_linear_fusion correctly unifies grouped global_amax after promotion.

Two non-blocking suggestions left as inline comments:

  1. Likely-unnecessary torch.cuda.synchronize in _run_reference_collect.
  2. load_calib_amax could keep the new-buffer branch going through register_buffer for symmetry with amax.setter.

Comment thread modelopt/torch/quantization/calib/mse.py Outdated
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py Outdated
@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Regarding #1536 (review): leaving this unchanged per branch-owner review. The current mse_calibrate flow intentionally installs the selected MSE calibrator on the weight quantizer before loading the chosen stats, and the extra restore/reset change is not needed here. I reverted the local experimental change.

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Regarding the NVFP4 reference MSE sync comment: updated the source comment to clarify that the CPU-GPU sync is there to prevent reference MSE calibration for another weight from running in parallel. The comment now also documents the retained _losses_sum memory: one fp32 reduced loss per candidate per block, about 126 * num_blocks * 4 bytes for NVFP4.

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Correction to my previous note about the NVFP4 reference MSE sync comment: I updated the memory wording to express _losses_sum relative to the calibrated weight. With 16-element NVFP4 blocks and bf16 weights, _losses_sum is roughly 128 / 16 * (4 / 2) = 16x the calibrated weight size, and the sync prevents reference MSE calibration for another weight from overlapping that allocation.

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Regarding #1536 (review): Thanks for the careful pass. The broader scope is intentional: the fp32 amax preservation, static-NVFP4-only MSE gate, max-calibration promotion, and grouped global-amax sync are tied together so static NVFP4 weights are finalized consistently before downstream calibrators/export. I agree the PR body should make that scope clearer, especially the load_calib_amax dtype behavior and the fact that max_calibrate now performs static NVFP4 promotion. The local-hessian promotion path is redundant but no-op after max_calibrate. I left the ones_like expression as-is because it keeps the shape/dtype broadcasting explicit. I also pushed fc8a95d0 to clarify the NVFP4 reference MSE sync comment and its _losses_sum memory cost.

@realAsma

Copy link
Copy Markdown
Contributor Author

Is fp32 MSE scale preserved after save/restore?

@realAsma

Copy link
Copy Markdown
Contributor Author

🤖 Bot comment.

Regarding #1536 (comment): yes. The MSE amax is stored in _amax as fp32 by load_calib_amax, and the quantizer save/restore path records buffer dtype in _pytorch_state_metadata before recreating buffers on restore. I also verified a local TensorQuantizer round trip: original _amax, saved state, metadata, and restored _amax all stayed torch.float32.

Comment thread modelopt/torch/quantization/config.py Outdated
@Fridah-nv

Copy link
Copy Markdown
Contributor

At first glance, I find the functionality of the two flag fp8_scale_sweep and apply_mse_nvfp_static_only a bit confusing. Talked with claude for some better naming ideas.
Do we agree that for NVFP4 we always recommend using FP8 sweep? If so, the the fp8_scale_sweep=false, apply_mse_nvfp_static_only=true is redundant.

Here's the full suggestion, makes sense on my side:

Suggestion: collapse fp8_scale_sweep + apply_mse_nvfp_static_only into a single sweep_mode enum

The two booleans give four combinations but only three are sensible — the fp8_scale_sweep=false, apply_mse_nvfp_static_only=true cell silently runs
multiplier-MSE on NVFP4-static weights, which nobody would intentionally pick. An enum eliminates that footgun and makes intent readable at the call site.

algorithm:
  method: mse
  sweep_mode: fp8_grid_nvfp4_only   # multiplier | fp8_grid | fp8_grid_nvfp4_only
sweep_mode NVFP4-static weight Other weight Replaces today's flags
multiplier (default) multiplier MSE multiplier MSE false / false
fp8_grid FP8 grid sweep FP8 grid via registry, else multiplier MSE true / false
fp8_grid_nvfp4_only FP8 grid sweep skipped (keeps max-cal amax) true / true

The preset becomes one line:

algorithm:
  method: mse
  sweep_mode: fp8_grid_nvfp4_only

Wins:

  • Three legible modes instead of four boolean combos with one degenerate cell.
  • fp8_grid keeps the backend-generic FP8 sweep registry as a valid extension point (matches how fp8_scale_sweep works today).
  • The _only suffix makes the scope semantics explicit — clearer than apply_mse_nvfp_static_only and fixes the nvfp → nvfp4 typo on the way.
  • Forward-compatible: a future fp8_grid_*_only mode (e.g. fp8_grid_mxfp8_only) is easy to add without another boolean.

If a full enum migration is heavier than worth it for this PR, the minimal-change fallback is just to rename apply_mse_nvfp_static_only →
nvfp4_static_only (typo + drop the redundant apply_mse_ prefix) and leave fp8_scale_sweep alone since it's genuinely backend-generic via the
registry.

Comment thread modelopt/torch/quantization/calib/mse.py
Comment thread modelopt/torch/quantization/calib/mse.py
Comment thread modelopt/torch/quantization/model_calib.py
quant_func=quant_func,
)

return MseCalibrator(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we should skip initializing MseCalibrator for non-NVFP4 static quantizers .. is this code even necessary?

@realAsma realAsma May 27, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should skip initializing MseCalibrator for non-NVFP4 static quantizers
We skip MSE calibration for non-NVFP4 if fp8_sweep is enabled; otherwise we support sweeping MSE scales. This design is general and support broadly.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This needs a if not fp8_scale_sweep: before it like in the original line 533 in main

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

thanks for catching

@jenchen13 jenchen13 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can we identify the minimal amount of FP32 casting needed for TensorQuantizer amax? right now it seems to be repeated in 10+ places for safety reasons which should not be necessary if the quantization algorithms respect the initial dtypes

Comment thread modelopt/torch/quantization/calib/mse.py Outdated
Comment thread modelopt/torch/quantization/calib/mse.py
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py Outdated
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py
Comment thread modelopt/torch/quantization/nn/modules/tensor_quantizer.py Outdated
@realAsma

Copy link
Copy Markdown
Contributor Author

/claude review

@realAsma
realAsma requested a review from jenchen13 May 28, 2026 21:47

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude review passed — no blocking issues found. LGTM

@realAsma
realAsma enabled auto-merge (squash) May 28, 2026 21:52
@realAsma
realAsma force-pushed the asma/mse_cleanups branch from 41e735c to b9dfea7 Compare May 29, 2026 13:32
@realAsma
realAsma disabled auto-merge May 29, 2026 14:01
@realAsma
realAsma force-pushed the asma/mse_cleanups branch from b9dfea7 to 5a281e1 Compare May 29, 2026 14:13
@realAsma
realAsma enabled auto-merge (squash) May 29, 2026 14:13
@realAsma
realAsma disabled auto-merge May 29, 2026 14:13
@realAsma
realAsma enabled auto-merge (squash) May 29, 2026 15:17
@realAsma
realAsma force-pushed the asma/mse_cleanups branch 2 times, most recently from 1db9749 to 6672519 Compare May 29, 2026 17:01
Comment thread modelopt/torch/quantization/calib/mse.py
Comment thread modelopt/torch/quantization/calib/mse.py
quant_func=quant_func,
)

return MseCalibrator(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This needs a if not fp8_scale_sweep: before it like in the original line 533 in main

Comment thread modelopt/torch/quantization/model_calib.py
@realAsma
realAsma requested a review from jenchen13 May 29, 2026 20:00
realAsma and others added 3 commits May 29, 2026 20:01
Signed-off-by: realAsma <akuriparambi@nvidia.com>
Under fp8_scale_sweep=True, only registered backends and static NVFP4
weights are MSE-calibrated; all other quantizers (INT8, plain FP8,
unregistered backends) are skipped instead of falling through to the
multiplier-search MseCalibrator. Fixes the gpu_megatron mixed-precision
test that asserts plain FP8 layers are left untouched under sweep.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Signed-off-by: realAsma <akuriparambi@nvidia.com>
The cached per-block amax is populated by both the Triton fast path and
the reference sweep, so the _fast suffix is misleading. Per review.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Signed-off-by: realAsma <akuriparambi@nvidia.com>
@realAsma
realAsma force-pushed the asma/mse_cleanups branch from e20f6ba to 5aa79a4 Compare May 29, 2026 20:01
@realAsma
realAsma merged commit d7e72f4 into main May 29, 2026
51 checks passed
@realAsma
realAsma deleted the asma/mse_cleanups branch May 29, 2026 21:04
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants