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. diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index ec1958649dd..2980d9e5d30 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 ``(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_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+)\.(.+)$") @@ -936,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( @@ -945,10 +950,8 @@ 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() 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 +973,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..cdc4a69c16f 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron.py @@ -1145,6 +1145,16 @@ 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: + assert getattr(linear, "_pg_collection", None) is not None + 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( @@ -1184,6 +1194,9 @@ def _test_te_grouped_sharded_state_dict_reshard_helper( checkpoint_path, rank, 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 @@ -1191,12 +1204,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", @@ -1207,6 +1222,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 ) @@ -1219,11 +1236,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", @@ -1232,6 +1251,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( @@ -1305,6 +1326,30 @@ def test_te_grouped_sharded_state_dict_reshard( ) +@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( + _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, + # 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, + ) + ) + + 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(