Repository navigation
Refine static NVFP4 MSE calibration - #1536
Conversation
|
Note Reviews pausedIt 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 Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughThis 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. ChangesNVFP4 MSE Calibration and Static Promotion Refactor
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Suggested reviewers
🚥 Pre-merge checks | ✅ 6✅ Passed checks (6 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
|
There was a problem hiding this comment.
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.
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
📒 Files selected for processing (12)
examples/llm_ptq/cast_mxfp4_to_nvfp4.pymodelopt/torch/quantization/calib/mse.pymodelopt/torch/quantization/config.pymodelopt/torch/quantization/model_calib.pymodelopt/torch/quantization/nn/modules/tensor_quantizer.pymodelopt/torch/quantization/utils/core_utils.pymodelopt_recipes/configs/ptq/presets/model/nvfp4_w4a4_weight_mse_fp8_sweep.yamltests/gpu/torch/quantization/test_gptq.pytests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.pytests/gpu/torch/quantization/test_nvfp4_static_quantizer_cuda.pytests/unit/torch/quantization/plugins/test_fused_experts.pytests/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
Codecov Report❌ Patch coverage is
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
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Sentry. 🚀 New features to boost your workflow:
|
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 |
There was a problem hiding this comment.
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 winRestore the original calibrator after the temporary MSE pass.
This loop leaves
weight_quantizer._calibratorpointing atcal, then immediately resets it. For the baseMseCalibrator,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
📒 Files selected for processing (5)
modelopt/torch/quantization/calib/mse.pymodelopt/torch/quantization/model_calib.pymodelopt/torch/quantization/nn/modules/tensor_quantizer.pytests/gpu/torch/quantization/test_nvfp4_fp8_sweep_kernel.pytests/unit/torch/quantization/test_mse_calibrator.py
Auto-replying to realAsma review comments:
|
/claude review |
cjluo-nv
left a comment
There was a problem hiding this comment.
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_syncis now called insidemax_calibrate, so the bootstrap (formerly only triggered frommse_calibrate) now runs for everymax_calibrateconsumer (gptq, awq postprocess, smoothquant, etc.). Intended? The_check_moe_calibration_completepath runs after promotion — fine since the check looks at_amaxon input quantizers too, but worth confirming nothing now warns where it didn't before.local_hessian_calibratelost its explicit_sync_grouped_weight_global_amax(model)call. It's now subsumed bymax_calibrate→ promotion path, but local_hessian also has its own per-module promotion later. Double-promotion is a no-op (covered byfrom_tensor_quantizer), just noting the redundancy._compute_candidate_amaxnow multiplies bytorch.ones_like(self._initial_amax, dtype=torch.float32)— theones_likeis 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.
There was a problem hiding this comment.
Claude review passed — no blocking issues found. LGTM
Findings: 0 CRITICAL, 0 IMPORTANT, 2 SUGGESTIONS
The refactor is internally consistent:
NVFP4MSECalibratorbecoming 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_amaxchange to preserve calibrator dtype on the buffer is covered by the newtest_load_calib_amax_preserves_fp32_result_dtyperegression test._promote_nvfp4_static_quantizers_with_global_amax_syncis idempotent (from_tensor_quantizeris a no-op for already-promoted modules) andpreprocess_linear_fusioncorrectly unifies groupedglobal_amaxafter promotion.
Two non-blocking suggestions left as inline comments:
- Likely-unnecessary
torch.cuda.synchronizein_run_reference_collect. load_calib_amaxcould keep the new-buffer branch going throughregister_bufferfor symmetry withamax.setter.
Regarding #1536 (review): leaving this unchanged per branch-owner review. The current |
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 |
Correction to my previous note about the NVFP4 reference MSE sync comment: I updated the memory wording to express |
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 |
|
Is fp32 MSE scale preserved after save/restore? |
Regarding #1536 (comment): yes. The MSE amax is stored in |
|
At first glance, I find the functionality of the two flag Here's the full suggestion, makes sense on my side: Suggestion: collapse The two booleans give four combinations but only three are sensible — the algorithm:
method: mse
sweep_mode: fp8_grid_nvfp4_only # multiplier | fp8_grid | fp8_grid_nvfp4_only
The preset becomes one line: algorithm:
method: mse
sweep_mode: fp8_grid_nvfp4_onlyWins:
If a full enum migration is heavier than worth it for this PR, the minimal-change fallback is just to rename |
| quant_func=quant_func, | ||
| ) | ||
|
|
||
| return MseCalibrator( |
There was a problem hiding this comment.
we should skip initializing MseCalibrator for non-NVFP4 static quantizers .. is this code even necessary?
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
This needs a if not fp8_scale_sweep: before it like in the original line 533 in main
There was a problem hiding this comment.
thanks for catching
4bc83af to
afca4b0
Compare
afca4b0 to
49d5073
Compare
jenchen13
left a comment
There was a problem hiding this comment.
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
|
/claude review |
41e735c to
b9dfea7
Compare
b9dfea7 to
5a281e1
Compare
1db9749 to
6672519
Compare
| quant_func=quant_func, | ||
| ) | ||
|
|
||
| return MseCalibrator( |
There was a problem hiding this comment.
This needs a if not fp8_scale_sweep: before it like in the original line 533 in main
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>
e20f6ba to
5aa79a4
Compare
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:
_amaxand_global_amaxin FP32.Usage
Testing
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.
CONTRIBUTING.md: N/AAdditional Information
N/A