Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions megatron/core/transformer/moe/fused_a2a.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
51 changes: 47 additions & 4 deletions megatron/core/transformer/moe/token_dispatcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -1039,19 +1040,48 @@ 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)
# Static upper bound on permuted tokens passed to HybridEP (dropless EP rank
# 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
)
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
7 changes: 7 additions & 0 deletions megatron/core/transformer/transformer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
1 change: 1 addition & 0 deletions tests/unit_tests/models/test_hybrid_moe_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
46 changes: 46 additions & 0 deletions tests/unit_tests/transformer/moe/test_token_dispatcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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")
Expand Down
Loading