From ecaa324724be2fae3808f7fc60aa0de9e2f581f6 Mon Sep 17 00:00:00 2001 From: qiyuw Date: Mon, 11 May 2026 17:49:01 -0700 Subject: [PATCH 1/4] fix mxfp8 without overlap Signed-off-by: qiyuw --- megatron/core/optimizer/distrib_optimizer.py | 10 ++-- megatron/core/optimizer/optimizer.py | 55 +++++++++++++++++--- 2 files changed, 55 insertions(+), 10 deletions(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index a9768d5a49b..6683cdaf917 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -2987,10 +2987,12 @@ def step_with_ready_grads(self) -> bool: self._state_offloader.sync_before_step() update_successful = super().step_with_ready_grads() + should_sync_params = self.ddp_config.use_megatron_fsdp or ( + not self.ddp_config.overlap_param_gather and not getattr(self, '_defer_param_sync', False) + ) timers = self.config.timers - if timers is not None: + if timers is not None and should_sync_params: timers('params-all-gather', log_level=1).start(barrier=self.config.barrier_with_L1_time) - if self.ddp_config.use_megatron_fsdp: # Optionally all-gather Megatron-FSDP sharded main weights # early in preparation for the subsequent forward pass. @@ -3001,10 +3003,10 @@ def step_with_ready_grads(self) -> bool: # communication calls here. If overlapping all-gather for parameters, the following # the first all-gather is launched asynchronously in the next optimizer.zero_grad() # call and subsequent all-gathers are launched in the forward pre-hook. - if not self.ddp_config.overlap_param_gather: + if should_sync_params: for model_chunk in self.model_chunks: model_chunk.start_param_sync() - if timers is not None: + if timers is not None and should_sync_params: timers('params-all-gather').stop() if self._state_offloader is not None: diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index 9ae23bb4b7f..705bc1cab3e 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -1284,12 +1284,55 @@ def prepare_grads(self) -> bool: def step_with_ready_grads(self) -> bool: """Step the optimizer with ready gradients, return successful.""" success = True - for optimizer_idx, optimizer in enumerate(self.chained_optimizers): - success &= optimizer.step_with_ready_grads() - if self.config.overlap_param_gather_with_optimizer_step and optimizer_idx == 0: - assert success - assert len(optimizer.model_chunks) == 1 - optimizer.model_chunks[0].start_param_sync(force_dispatch=True) + # With MXFP8 grad-buffer reuse and non-overlap param gather, each DistOpt stages + # its own updated main-param shards into its param buffers during step. However, + # param sync is a DDP model-chunk operation: model_chunk.start_param_sync() gathers + # both dense and expert bucket groups, copies gathered values into model weights, + # and zeros the shared MXFP8 param/grad buffers. For MoE, dense and expert DistOpts + # may share the same model chunk, so defer param sync until all chained optimizers + # have staged their params, then sync each model chunk once. + defer_param_sync = ( + self.config.reuse_grad_buf_for_mxfp8_param_ag + and not self.config.overlap_param_gather + ) + deferred_model_chunks = [] + deferred_model_chunk_ids = set() + + if defer_param_sync: + from .distrib_optimizer import DistributedOptimizer + + for optimizer in self.chained_optimizers: + if isinstance(optimizer, DistributedOptimizer): + optimizer._defer_param_sync = True + for model_chunk in optimizer.model_chunks: + model_chunk_id = id(model_chunk) + if model_chunk_id not in deferred_model_chunk_ids: + deferred_model_chunk_ids.add(model_chunk_id) + deferred_model_chunks.append(model_chunk) + + try: + for optimizer_idx, optimizer in enumerate(self.chained_optimizers): + success &= optimizer.step_with_ready_grads() + if self.config.overlap_param_gather_with_optimizer_step and optimizer_idx == 0: + assert success + assert len(optimizer.model_chunks) == 1 + optimizer.model_chunks[0].start_param_sync(force_dispatch=True) + finally: + if defer_param_sync: + for optimizer in self.chained_optimizers: + if hasattr(optimizer, '_defer_param_sync'): + optimizer._defer_param_sync = False + + if defer_param_sync and success: + timers = self.config.timers + if timers is not None: + timers('params-all-gather', log_level=1).start( + barrier=self.config.barrier_with_L1_time + ) + for model_chunk in deferred_model_chunks: + model_chunk.start_param_sync() + if timers is not None: + timers('params-all-gather').stop() return success From e022e637a925b600f9b7092a68af3ab05d9d4cc6 Mon Sep 17 00:00:00 2001 From: qiyuw Date: Mon, 11 May 2026 17:53:12 -0700 Subject: [PATCH 2/4] minor Signed-off-by: qiyuw --- megatron/core/optimizer/distrib_optimizer.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 6683cdaf917..b8083e00594 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -2987,11 +2987,11 @@ def step_with_ready_grads(self) -> bool: self._state_offloader.sync_before_step() update_successful = super().step_with_ready_grads() - should_sync_params = self.ddp_config.use_megatron_fsdp or ( - not self.ddp_config.overlap_param_gather and not getattr(self, '_defer_param_sync', False) + should_sync_params = not self.ddp_config.overlap_param_gather and not getattr( + self, '_defer_param_sync', False ) timers = self.config.timers - if timers is not None and should_sync_params: + if timers is not None and (self.ddp_config.use_megatron_fsdp or should_sync_params): timers('params-all-gather', log_level=1).start(barrier=self.config.barrier_with_L1_time) if self.ddp_config.use_megatron_fsdp: # Optionally all-gather Megatron-FSDP sharded main weights @@ -3006,7 +3006,7 @@ def step_with_ready_grads(self) -> bool: if should_sync_params: for model_chunk in self.model_chunks: model_chunk.start_param_sync() - if timers is not None and should_sync_params: + if timers is not None and (self.ddp_config.use_megatron_fsdp or should_sync_params): timers('params-all-gather').stop() if self._state_offloader is not None: From 19266cd2937cb28a763851d55eb0ebe8051e2d6e Mon Sep 17 00:00:00 2001 From: qiyuw Date: Tue, 12 May 2026 17:13:37 -0700 Subject: [PATCH 3/4] unit test for mxfp8 on moe model Signed-off-by: qiyuw --- tests/unit_tests/test_fp8_param.py | 59 ++++++++++++++++++++++++++---- 1 file changed, 51 insertions(+), 8 deletions(-) diff --git a/tests/unit_tests/test_fp8_param.py b/tests/unit_tests/test_fp8_param.py index 9cc77a2c397..113ac7224b8 100644 --- a/tests/unit_tests/test_fp8_param.py +++ b/tests/unit_tests/test_fp8_param.py @@ -91,7 +91,10 @@ def model_provider( model_parallel_cuda_manual_seed(_SEED) args = get_args() config = core_transformer_config_from_args(args) - transformer_layer_spec = layer_spec_fn() + transformer_layer_spec = layer_spec_fn( + num_experts=args.num_experts, + moe_grouped_gemm=args.moe_grouped_gemm, + ) return GPTModel( config=config, transformer_layer_spec=transformer_layer_spec, @@ -219,7 +222,7 @@ def _run_test_helper( eval_transition: bool = False, **kwargs, ): - """Test fp8_param with gpt_model.""" + """Test fp8_param with a small GPT model.""" args = self.create_test_args( tp_size, recipe, @@ -238,7 +241,10 @@ def _run_test_helper( set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp_size) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp_size, + expert_model_parallel_size=args.expert_model_parallel_size, + ) input_ids, labels, position_ids, attention_mask, loss_mask = self.get_batch( self.seq_length, self.micro_batch_size @@ -276,14 +282,21 @@ def _run_test_helper( if is_float8tensor(param): num_fp8_params += 1 - # Verify the number of fp8 params. fp8_layers = args.num_layers if kwargs.get("first_last_layers_bf16", False): fp8_layers -= kwargs["num_layers_at_start_in_bf16"] fp8_layers -= kwargs["num_layers_at_end_in_bf16"] - # Each layer has 4 GEMM weights: qkv, proj, fc1, fc2. - if fp8_param_gather: - assert num_fp8_params == 4 * fp8_layers + if fp8_param_gather and fp8_layers > 0: + if args.num_experts is None: + # Each dense layer has 4 GEMM weights: qkv, proj, fc1, fc2. + assert num_fp8_params == 4 * fp8_layers + else: + assert num_fp8_params > 0 + assert any( + not getattr(param, 'allreduce', True) for param in gpt_model[0].parameters() + ) + if not inference: + assert len(optimizer.chained_optimizers) >= 2 # Verify that bf16 params (embedding, LN, etc.) in the MXFP8 model are mapped # to the param buffer (shared with grad buffer) rather than allocated separately. @@ -382,7 +395,7 @@ def _run_test_helper( return torch.tensor(loss_list) def run_test(self, tp_size, recipe, inference: bool = False, **kwargs): - """Test fp8_param with gpt_model.""" + """Test fp8_param with a small GPT model.""" if inference: with torch.inference_mode(): self._run_test_helper(tp_size, recipe, inference=True, **kwargs) @@ -502,6 +515,36 @@ def test_mxfp8(self, tp_size, dp_overlap): kwargs = {"overlap_param_gather": dp_overlap[0], "overlap_grad_reduce": dp_overlap[1]} self.run_test(tp_size=tp_size, recipe="mxfp8", **kwargs) + @pytest.mark.skipif( + get_device_arch_version() < 10, reason="MXFP8 is supported since Blackwell architecture" + ) + @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) + @pytest.mark.skipif(not is_te_min_version("2.3.0.dev0"), reason="TE 2.3.0.dev0 is required") + @pytest.mark.parametrize("tp_size", [1]) + @pytest.mark.parametrize("dp_overlap", [(False, False), (False, True), (True, True)]) + def test_mxfp8_moe(self, tp_size, dp_overlap): + """ + dp_overlap: (overlap_param_gather, overlap_grad_reduce) + """ + kwargs = { + "overlap_param_gather": dp_overlap[0], + "overlap_grad_reduce": dp_overlap[1], + "num_layers": 4, + "vocal_size": 128800, + "hidden_size": 128, + "num_attention_heads": 8, + "expert_model_parallel_size": 2, + "num_experts": 2, + "moe_grouped_gemm": True, + "moe_token_dispatcher_type": "alltoall", + "moe_router_topk": 1, + "moe_router_pre_softmax": True, + "moe_router_load_balancing_type": "none", + "moe_aux_loss_coeff": 0.0, + "moe_ffn_hidden_size": 128, + } + self.run_test(tp_size=tp_size, recipe="mxfp8", **kwargs) + @pytest.mark.skipif( get_device_arch_version() < 10, reason="MXFP8 is supported since Blackwell architecture" ) From d68550d7ad00289676cc69cd1ca50f26fa72b529 Mon Sep 17 00:00:00 2001 From: qiyuw Date: Tue, 12 May 2026 17:19:26 -0700 Subject: [PATCH 4/4] lint Signed-off-by: qiyuw --- megatron/core/optimizer/optimizer.py | 3 +-- tests/unit_tests/test_fp8_param.py | 3 +-- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index 705bc1cab3e..eb163bc5622 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -1292,8 +1292,7 @@ def step_with_ready_grads(self) -> bool: # may share the same model chunk, so defer param sync until all chained optimizers # have staged their params, then sync each model chunk once. defer_param_sync = ( - self.config.reuse_grad_buf_for_mxfp8_param_ag - and not self.config.overlap_param_gather + self.config.reuse_grad_buf_for_mxfp8_param_ag and not self.config.overlap_param_gather ) deferred_model_chunks = [] deferred_model_chunk_ids = set() diff --git a/tests/unit_tests/test_fp8_param.py b/tests/unit_tests/test_fp8_param.py index 113ac7224b8..6c23b3e26ca 100644 --- a/tests/unit_tests/test_fp8_param.py +++ b/tests/unit_tests/test_fp8_param.py @@ -92,8 +92,7 @@ def model_provider( args = get_args() config = core_transformer_config_from_args(args) transformer_layer_spec = layer_spec_fn( - num_experts=args.num_experts, - moe_grouped_gemm=args.moe_grouped_gemm, + num_experts=args.num_experts, moe_grouped_gemm=args.moe_grouped_gemm ) return GPTModel( config=config,