diff --git a/megatron/core/transformer/moe/fused_a2a.py b/megatron/core/transformer/moe/fused_a2a.py index c281c51b4fb..09e50c9f4ad 100644 --- a/megatron/core/transformer/moe/fused_a2a.py +++ b/megatron/core/transformer/moe/fused_a2a.py @@ -277,6 +277,10 @@ def set_deepep_num_sms(num_sms): _hybrid_ep_buffer = None +# HybridEP dispatch/combine kernels use 64-token chunks for their public APIs. +HYBRIDEP_TOKEN_ALIGNMENT = 64 + + def init_hybrid_ep_buffer( group: torch.distributed.ProcessGroup, hidden_dim: int, diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 4c4b65679c3..683450a4a28 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -19,6 +19,7 @@ ) from megatron.core.transformer.enums import CudaGraphModule from megatron.core.transformer.moe.fused_a2a import ( + HYBRIDEP_TOKEN_ALIGNMENT, ensure_nccl_ep_bootstrapped, fused_combine, fused_dispatch, @@ -1039,11 +1040,40 @@ def __init__( self.moe_expert_rank_capacity_factor = self.config.moe_expert_rank_capacity_factor self.over_budget = torch.zeros(1, dtype=torch.bool, device='cuda') + # HybridEP dispatch expects equal per-rank input sizes. When requested, + # variable token counts are padded to the group-wide max and trimmed in combine. + self._original_num_tokens: Optional[int] = None + self._padded_num_tokens: Optional[int] = None def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): num_tokens = routing_map.shape[0] - self.routing_map = routing_map.reshape(num_tokens, self.num_experts) - self.token_probs = probs.reshape(num_tokens, self.num_experts) + self._original_num_tokens = num_tokens + + padded_num_tokens = num_tokens + if self.config.moe_hybridep_pad_uneven_dispatch_inputs: + # 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( + [num_tokens], device=routing_map.device, dtype=torch.long + ) + torch.distributed.all_reduce( + max_num_tokens_across_ep, op=torch.distributed.ReduceOp.MAX, group=self.group + ) + padded_num_tokens = int(max_num_tokens_across_ep.item()) + padded_num_tokens += -padded_num_tokens % HYBRIDEP_TOKEN_ALIGNMENT + self._padded_num_tokens = padded_num_tokens + + routing_map = routing_map.reshape(num_tokens, self.num_experts) + probs = probs.reshape(num_tokens, self.num_experts) + if 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 + ) + probs = torch.cat([probs, probs.new_zeros((pad_rows, self.num_experts))], dim=0) + + self.routing_map = routing_map + self.token_probs = probs if self.moe_expert_rank_capacity_factor is not None: pad_multiple = get_align_size_for_quantization(self.config) @@ -1051,7 +1081,7 @@ def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): # budget). Tokens above this budget are dropped inside HybridEP; dispatch then # sets overflow_flag on the handle (accumulated in over_budget in dispatch()). budget = int( - routing_map.shape[0] + padded_num_tokens * self.config.moe_router_topk * self.moe_expert_rank_capacity_factor ) @@ -1062,7 +1092,7 @@ def setup_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor): # in dispatch) and does not drop tokens or report overflow. # Compute the capacity for each expert at the drop_and_pad mode if self.drop_and_pad: - num_out_tokens = num_tokens * self.config.moe_router_topk + num_out_tokens = padded_num_tokens * self.config.moe_router_topk # Drop and pad the input to capacity. self.capacity = get_capacity( num_tokens=num_out_tokens, @@ -1091,6 +1121,11 @@ def dispatch( self.token_probs = self.token_probs.float() # downcast or upcast if self.config.fp8 or self.config.fp4: self.pad_multiple = get_align_size_for_quantization(self.config) + if self._padded_num_tokens is not None and hidden_states.shape[0] < self._padded_num_tokens: + pad_rows = self._padded_num_tokens - hidden_states.shape[0] + hidden_states = torch.cat( + [hidden_states, hidden_states.new_zeros((pad_rows, hidden_states.shape[-1]))], dim=0 + ) dispatched_hidden, self.dispatched_probs, _, tokens_per_expert, self.handle = ( hybrid_ep_dispatch( x=hidden_states, @@ -1137,12 +1172,20 @@ def combine( pad_multiple=self.pad_multiple, fused=self.config.moe_permute_fusion_into_hybridep, ) + if ( + self._padded_num_tokens is not None + and self._original_num_tokens is not None + and hidden_states.shape[0] > self._original_num_tokens + ): + hidden_states = hidden_states[: self._original_num_tokens] # Release the used handle/num_permuted_tokens which could change in each iteration. # For drop_and_pad mode, we don't need to reset the num_permuted_tokens and # num_dispatched_tokens, because their values never change. self.handle = None if not self.drop_and_pad: self.num_permuted_tokens = None + self._original_num_tokens = None + self._padded_num_tokens = None return hidden_states def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 0c9ce022db7..9bf67fc8854 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -865,6 +865,13 @@ class TransformerConfig(ModelParallelConfig): moe_permute_fusion_into_hybridep: bool = False """Fuse token rearrangement ops during token dispatching for HybridEP.""" + moe_hybridep_pad_uneven_dispatch_inputs: bool = False + """Pad uneven HybridEP dispatch inputs to the group maximum before dispatch. + Enable when local HybridEP input token counts can differ across ranks, for example + with dynamically packed THD inputs. Leave disabled when dispatcher inputs are + already padded to equal token counts. + """ + 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 f7dc78ce9a2..ef8dcd94c5d 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -337,6 +337,7 @@ "use_transformer_engine_op_fuser": False, "moe_single_grouped_weight": False, "moe_single_grouped_bias": False, + "moe_hybridep_pad_uneven_dispatch_inputs": False, } # Fields to ignore entirely (ephemeral, environment-specific, very large). SKIP_FIELDS = set() diff --git a/tests/unit_tests/transformer/moe/test_token_dispatcher.py b/tests/unit_tests/transformer/moe/test_token_dispatcher.py index 20839310ce2..a558cd2dc2d 100644 --- a/tests/unit_tests/transformer/moe/test_token_dispatcher.py +++ b/tests/unit_tests/transformer/moe/test_token_dispatcher.py @@ -7,8 +7,10 @@ from megatron.core import config, parallel_state from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules +from megatron.core.transformer.moe.fused_a2a import HYBRIDEP_TOKEN_ALIGNMENT, reset_hybrid_ep_buffer from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules from megatron.core.transformer.moe.moe_utils import get_capacity +from megatron.core.transformer.moe.token_dispatcher import _HybridEPManager from megatron.core.transformer.spec_utils import get_submodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module @@ -426,6 +428,49 @@ def is_nccl_ep_available(): return HAVE_TE_EP +def test_hybridep_pad_uneven_dispatch_inputs_metadata(monkeypatch): + manager = _HybridEPManager.__new__(_HybridEPManager) + manager.group = object() + manager.num_local_experts = 2 + manager.num_experts = 4 + manager.config = TransformerConfig( + num_layers=1, + hidden_size=16, + num_attention_heads=4, + num_moe_experts=4, + moe_router_topk=2, + moe_hybridep_pad_uneven_dispatch_inputs=True, + ) + manager.moe_expert_rank_capacity_factor = None + manager.drop_and_pad = False + + local_num_tokens = 17 + max_num_tokens_across_ep = 70 + padded_num_tokens = ( + max_num_tokens_across_ep + -max_num_tokens_across_ep % HYBRIDEP_TOKEN_ALIGNMENT + ) + routing_map = torch.ones((local_num_tokens, manager.num_experts), dtype=torch.bool) + probs = torch.ones((local_num_tokens, manager.num_experts), dtype=torch.float32) + + def fake_all_reduce(tensor, op=None, group=None): + assert op == torch.distributed.ReduceOp.MAX + assert group is manager.group + tensor.fill_(max_num_tokens_across_ep) + + monkeypatch.setattr(torch.distributed, "all_reduce", fake_all_reduce) + + manager.setup_metadata(routing_map, probs) + + assert manager._original_num_tokens == local_num_tokens + assert manager._padded_num_tokens == padded_num_tokens + assert manager.routing_map.shape == (padded_num_tokens, manager.num_experts) + assert manager.token_probs.shape == (padded_num_tokens, manager.num_experts) + torch.testing.assert_close(manager.routing_map[:local_num_tokens], routing_map) + torch.testing.assert_close(manager.token_probs[:local_num_tokens], probs) + assert not manager.routing_map[local_num_tokens:].any() + assert not manager.token_probs[local_num_tokens:].any() + + @pytest.mark.skipif( not is_deep_ep_available() and not is_hybrid_ep_available(), reason="Deep EP and Hybrid EP are not available", @@ -435,6 +480,7 @@ def setup_method(self, method): pass def teardown_method(self, method): + reset_hybrid_ep_buffer() Utils.destroy_model_parallel() @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")