diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 6ea0d97160f..a9d1d5f50cc 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1065,7 +1065,11 @@ def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): self._original_num_tokens = num_tokens padded_num_tokens = num_tokens - if self.config.sequence_packing_scheduler is not None: + equalize_thd_token_counts = ( + self.config.sequence_packing_scheduler is not None + or self.config.moe_hybridep_pad_variable_tokens + ) + if equalize_thd_token_counts: # Use the actual tp_ep max so all ranks in the MoE communication # group pass the same token count to HybridEP. max_num_tokens_across_ep = torch.tensor( @@ -1080,7 +1084,7 @@ def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): routing_map = routing_map.reshape(num_tokens, self.num_experts) probs = probs.reshape(num_tokens, self.num_experts) - if self.config.sequence_packing_scheduler is not None and padded_num_tokens > num_tokens: + if equalize_thd_token_counts and padded_num_tokens > num_tokens: pad_rows = padded_num_tokens - num_tokens routing_map = torch.cat( [routing_map, routing_map.new_zeros((pad_rows, self.num_experts))], dim=0 diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 10d9562f0bb..8c663e8bcbf 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -875,6 +875,13 @@ class TransformerConfig(ModelParallelConfig): moe_permute_fusion_into_hybridep: bool = False """Fuse token rearrangement ops during token dispatching for HybridEP.""" + moe_hybridep_pad_variable_tokens: bool = False + """Pad uneven local token counts to the HybridEP group maximum before dispatch. + + This is needed when the frontend supplies locally packed THD inputs whose token counts + can differ across ranks, without using Megatron Core's sequence_packing_scheduler. + """ + moe_per_layer_logging: bool = False """Enable per-layer logging for MoE, currently supports auxiliary loss and z loss.""" diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index 6f59ec6b0b0..736c56ba034 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -341,6 +341,7 @@ "moe_single_grouped_weight": False, "moe_single_grouped_bias": False, "head_wise_attn_gate": False, + "moe_hybridep_pad_variable_tokens": False, } # Fields to ignore entirely (ephemeral, environment-specific, very large). SKIP_FIELDS = set()