From 636b474beccfc0c3b0384876e6f6b7e53651cfc0 Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Mon, 21 Sep 2026 16:19:12 -0700 Subject: [PATCH 1/7] Fix grouped expert quantizer checkpoint replicas Signed-off-by: Hung-Yueh Chiang --- .../torch/quantization/plugins/megatron.py | 14 +++++++---- .../quantization/plugins/test_megatron.py | 24 +++++++++++++++++++ 2 files changed, 33 insertions(+), 5 deletions(-) diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index ec1958649dd..3d0ab56278a 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -946,9 +946,11 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): # Per-expert amax: assign the same global expert identity the weights use. ep_group, expt_dp_group = self._expert_parallel_groups() + parallel_state = self.parallel_state + assert parallel_state is not None + expt_tp_group = parallel_state.tensor_parallel_group.group num_global_experts = get_pg_size(ep_group) * self.num_gemms local_expert_indices_offset = get_pg_rank(ep_group) * self.num_gemms - edp_replica_id = get_pg_rank(expt_dp_group) ep_axis = len(sharded_offsets) for gemm_idx, subs in enumerate(per_expert_subs): if not subs: @@ -970,16 +972,18 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): if axis is not None } sub_sd = make_sharded_tensors_for_checkpoint( - expert_state, "", expert_axis, new_sharded_offsets + expert_state, + "", + expert_axis, + new_sharded_offsets, + tp_group=expt_tp_group, + dp_cp_group=expt_dp_group, ) # Rewrite each ShardedTensor.key to carry the global expert identity (dict keys, # which map to the local buffers on restore, are left untouched). replace_prefix_for_sharding(sub_sd, f"{gemm_idx}.", expert_prefix) for sub, _, _ in subs: sh_ten = sub_sd[f"{gemm_idx}.weight_quantizer.{sub}"] - replica_id = sh_ten.replica_id - if len(replica_id) == 3: - sh_ten.replica_id = (*replica_id[:2], edp_replica_id) sharded_state_dict[f"{prefix}weight_quantizer.{gemm_idx}.{sub}"] = sh_ten return sharded_state_dict diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py index 041a0bd3fca..dea6fe90f18 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py @@ -1184,6 +1184,8 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( checkpoint_path, rank, size, + save_etp_size=None, + load_etp_size=None, ): """Round-trip TEGroupedMLP amax through a topology change.""" num_experts = 4 @@ -1191,12 +1193,14 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( initialize_for_megatron( tensor_model_parallel_size=save_tp_size, expert_model_parallel_size=save_ep_size, + expert_tensor_parallel_size=save_etp_size, seed=SEED, ) source = _gpt_model_provider( tp_size=save_tp_size, ep_size=save_ep_size, + etp_size=save_etp_size, hidden_size=32, moe_grouped_gemm=True, transformer_impl="transformer_engine", @@ -1219,11 +1223,13 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( initialize_for_megatron( tensor_model_parallel_size=load_tp_size, expert_model_parallel_size=load_ep_size, + expert_tensor_parallel_size=load_etp_size, seed=SEED, ) target = _gpt_model_provider( tp_size=load_tp_size, ep_size=load_ep_size, + etp_size=load_etp_size, hidden_size=32, moe_grouped_gemm=True, transformer_impl="transformer_engine", @@ -1305,6 +1311,24 @@ def test_te_grouped_sharded_state_dict_reshard( ) +def test_te_grouped_sharded_state_dict_combined_tp_ep(dist_workers_size_4, tmp_path): + """Round-trip grouped expert quantizer state with TP and EP both greater than one.""" + dist_workers_size_4.run( + partial( + _test_te_grouped_sharded_state_dict_reshard_helper, + 2, + 2, + 2, + 2, + mtq.NVFP4_DEFAULT_CFG, + False, + tmp_path, + save_etp_size=1, + load_etp_size=1, + ) + ) + + def _test_te_grouped_vs_sequential_default_loss_helper(tp_size, ep_size, quant_cfg, rank, size): """TEGrouped quantized output should diverge from BF16 more than SequentialMLP under default sync=False.""" initialize_for_megatron( From 89c1c9d98c13bb0d417ea3310cbc1de943325d56 Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Tue, 22 Sep 2026 09:12:23 -0700 Subject: [PATCH 2/7] Document grouped expert checkpoint fix Signed-off-by: Hung-Yueh Chiang --- CHANGELOG.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 22169f974c0..29530e50e40 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -71,6 +71,7 @@ Changelog **Bug Fixes** +- Fix Megatron-Core checkpoint saving for quantized grouped MoE experts when tensor and expert parallelism are both enabled. - Fix shared ONNX export metadata and Diffusers attention policy: every ``NVFP4QuantExporter`` post-process now upgrades the default-domain opset to at least 23, all FP8 custom-op exports re-run ONNX shape/type inference after setting output metadata, and quantized SDPA derives FP8 MHA enablement from the live Q/K/V quantizers instead of honoring a caller-set ``_disable_fp8_mha`` attribute. - Fix ONNX FP16 conversion failing to preserve public output types when type inference changes a graph output declaration before output casts are inserted. - Fix ``examples/hf_ptq/hf_ptq.py`` discarding a completed PTQ run (no checkpoint exported) when the optional post-quantization sanity-check ``generate()`` call raised, for example because ``device_map="auto"`` placed part of the model on CPU. That failure is now caught and only skips the sanity check; export proceeds regardless. From dbfa1cdb5059d4fa156f0509bc850feb28cefcb0 Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Wed, 23 Sep 2026 11:31:24 -0700 Subject: [PATCH 3/7] Address grouped expert process group review Signed-off-by: Hung-Yueh Chiang --- modelopt/torch/quantization/plugins/megatron.py | 17 ++++++++++------- .../torch/quantization/plugins/test_megatron.py | 17 +++++++++++++++++ 2 files changed, 27 insertions(+), 7 deletions(-) diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index 3d0ab56278a..e6d7e131ac1 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -876,12 +876,13 @@ def _process_quantizer_amax(self, k, v, quantizer_state_dict): quantizer_state_dict[k] = v.view(-1) if v.numel() == 1 else v def _expert_parallel_groups(self): - """Return the (ep, expt_dp) process groups used to place fused experts globally.""" + """Return the process groups used to place fused experts globally.""" pg_collection = getattr(self, "_pg_collection", None) if pg_collection is not None: - return pg_collection.ep, pg_collection.expt_dp + return pg_collection.ep, pg_collection.expt_tp, pg_collection.expt_dp return ( mcore_parallel.get_expert_model_parallel_group(), + mcore_parallel.get_expert_tensor_parallel_group(), mcore_parallel.get_expert_data_parallel_group(), ) @@ -924,6 +925,7 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): # Channel shard axes (per real key); _global_amax stays un-sharded along channels but # still rides with the expert identity below. shard_axis_dict = self._get_shard_axis_dict(quantizer_state_dict) + ep_group, expt_tp_group, expt_dp_group = self._expert_parallel_groups() # Split per-expert weight_quantizer.{i}.* from shared (input/output) quantizer buffers. expert_re = re.compile(r"^weight_quantizer\.(\d+)\.(.+)$") @@ -940,15 +942,16 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): shared_axis_dict = {k: shard_axis_dict[k] for k in shared_state if k in shard_axis_dict} sharded_state_dict.update( make_sharded_tensors_for_checkpoint( - shared_state, prefix, shared_axis_dict, sharded_offsets + shared_state, + prefix, + shared_axis_dict, + sharded_offsets, + tp_group=expt_tp_group, + dp_cp_group=expt_dp_group, ) ) # Per-expert amax: assign the same global expert identity the weights use. - ep_group, expt_dp_group = self._expert_parallel_groups() - parallel_state = self.parallel_state - assert parallel_state is not None - expt_tp_group = parallel_state.tensor_parallel_group.group num_global_experts = get_pg_size(ep_group) * self.num_gemms local_expert_indices_offset = get_pg_rank(ep_group) * self.num_gemms ep_axis = len(sharded_offsets) diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py index dea6fe90f18..84286cc7374 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py @@ -1145,6 +1145,15 @@ def _assert_te_grouped_weight_quantizer_state(model, expected_amax, expect_globa assert checked > 0, "no TEGrouped per-expert weight quantizer amax was checked" +def _override_te_grouped_modelopt_tp_group(model, tp_group): + grouped_linears = [ + module for module in model.modules() if isinstance(module, _QuantMegatronTEGroupedLinear) + ] + assert grouped_linears, "no quantized TEGroupedLinear found" + for linear in grouped_linears: + linear.parallel_state.tensor_parallel_group.group = tp_group + + def test_initialize_grouped_weight_quantizer_state_for_restore(): """Missing grouped state inherits the shape and dtype of a populated sibling.""" source = mtq.nn.StaticBlockScaleQuantizer.from_tensor_quantizer( @@ -1186,6 +1195,7 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( size, save_etp_size=None, load_etp_size=None, + override_modelopt_tp_group=False, ): """Round-trip TEGroupedMLP amax through a topology change.""" num_experts = 4 @@ -1211,6 +1221,8 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( if isinstance(module, TopKRouter): module.topk = module.num_experts mtq.quantize(source, copy.deepcopy(quant_cfg), forward) + if override_modelopt_tp_group: + _override_te_grouped_modelopt_tp_group(source, get_tensor_model_parallel_group()) _set_te_grouped_weight_quantizer_state( source, get_expert_model_parallel_rank(), save_num_local_experts ) @@ -1238,6 +1250,8 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( target_models = [target] restore_sharded_modelopt_state(target_models, checkpoint_path) target = target_models[0] + if override_modelopt_tp_group: + _override_te_grouped_modelopt_tp_group(target, get_tensor_model_parallel_group()) load_distributed_checkpoint(checkpoint_path, target) load_num_local_experts = num_experts // load_ep_size expected_amax = tuple( @@ -1325,6 +1339,9 @@ def test_te_grouped_sharded_state_dict_combined_tp_ep(dist_workers_size_4, tmp_p tmp_path, save_etp_size=1, load_etp_size=1, + # Simulate child conversion without the parent MLP setup: checkpoint groups must still + # come from the grouped linear's MCore process-group collection. + override_modelopt_tp_group=True, ) ) From 6a5a123f8b240b0896b0b63496f507179dab16d4 Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Wed, 23 Sep 2026 12:20:01 -0700 Subject: [PATCH 4/7] Fix shared TE grouped quantizer replica groups Signed-off-by: Hung-Yueh Chiang --- modelopt/torch/quantization/plugins/megatron.py | 7 +------ .../torch/quantization/plugins/test_megatron.py | 1 + 2 files changed, 2 insertions(+), 6 deletions(-) diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index e6d7e131ac1..ffd43bb6cac 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -942,12 +942,7 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): shared_axis_dict = {k: shard_axis_dict[k] for k in shared_state if k in shard_axis_dict} sharded_state_dict.update( make_sharded_tensors_for_checkpoint( - shared_state, - prefix, - shared_axis_dict, - sharded_offsets, - tp_group=expt_tp_group, - dp_cp_group=expt_dp_group, + shared_state, prefix, shared_axis_dict, sharded_offsets ) ) diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py index 84286cc7374..18482f7f87b 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py @@ -1151,6 +1151,7 @@ def _override_te_grouped_modelopt_tp_group(model, tp_group): ] assert grouped_linears, "no quantized TEGroupedLinear found" for linear in grouped_linears: + assert getattr(linear, "_pg_collection", None) is not None linear.parallel_state.tensor_parallel_group.group = tp_group From 87cb850f74acaa2919d555872baf1ebbc062eb8c Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Wed, 23 Sep 2026 13:13:56 -0700 Subject: [PATCH 5/7] Document shared quantizer replica invariant Signed-off-by: Hung-Yueh Chiang --- modelopt/torch/quantization/plugins/megatron.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index ffd43bb6cac..d7a408ce957 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -938,7 +938,10 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): else: shared_state[k] = v - # Shared quantizer buffers: replicated across experts, plain base offsets. + # Shared quantizer buffers have no expert identity in their keys or offsets. Keep the + # dense TP/DP defaults so replica IDs distinguish EP ranks; using expt_tp/expt_dp here + # would collide across EP ranks. Expert-axis sharding would require an EP-aware + # replica ID in addition to the expert process groups. shared_axis_dict = {k: shard_axis_dict[k] for k in shared_state if k in shard_axis_dict} sharded_state_dict.update( make_sharded_tensors_for_checkpoint( From c2be0f2c68be1f9656cd0d2bfc5ec094e99a366f Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Wed, 23 Sep 2026 16:27:33 -0700 Subject: [PATCH 6/7] Test grouped checkpoint with unmodified TP state Signed-off-by: Hung-Yueh Chiang --- .../torch/quantization/plugins/test_megatron.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py index 18482f7f87b..cdc4a69c16f 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py @@ -1326,7 +1326,10 @@ def test_te_grouped_sharded_state_dict_reshard( ) -def test_te_grouped_sharded_state_dict_combined_tp_ep(dist_workers_size_4, tmp_path): +@pytest.mark.parametrize("override_modelopt_tp_group", [False, True]) +def test_te_grouped_sharded_state_dict_combined_tp_ep( + dist_workers_size_4, tmp_path, override_modelopt_tp_group +): """Round-trip grouped expert quantizer state with TP and EP both greater than one.""" dist_workers_size_4.run( partial( @@ -1340,9 +1343,9 @@ def test_te_grouped_sharded_state_dict_combined_tp_ep(dist_workers_size_4, tmp_p tmp_path, save_etp_size=1, load_etp_size=1, - # Simulate child conversion without the parent MLP setup: checkpoint groups must still - # come from the grouped linear's MCore process-group collection. - override_modelopt_tp_group=True, + # The True case simulates child conversion without the parent MLP setup: checkpoint + # groups must still come from the grouped linear's MCore process-group collection. + override_modelopt_tp_group=override_modelopt_tp_group, ) ) From 40419504c76d5685197bc5323f8a7df27621c467 Mon Sep 17 00:00:00 2001 From: Hung-Yueh Chiang Date: Wed, 23 Sep 2026 16:45:30 -0700 Subject: [PATCH 7/7] Document expert process group tuple order Signed-off-by: Hung-Yueh Chiang --- modelopt/torch/quantization/plugins/megatron.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index d7a408ce957..2980d9e5d30 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -876,7 +876,7 @@ def _process_quantizer_amax(self, k, v, quantizer_state_dict): quantizer_state_dict[k] = v.view(-1) if v.numel() == 1 else v def _expert_parallel_groups(self): - """Return the process groups used to place fused experts globally.""" + """Return ``(ep_group, expt_tp_group, expt_dp_group)`` for fused experts.""" pg_collection = getattr(self, "_pg_collection", None) if pg_collection is not None: return pg_collection.ep, pg_collection.expt_tp, pg_collection.expt_dp