From 5eb731b5f433a9dc019ba3af72f9c227e1a2a47f Mon Sep 17 00:00:00 2001 From: janEbert Date: Thu, 21 May 2026 11:43:50 +0200 Subject: [PATCH 1/2] Revert "Add Python-side guardrail for DeepEP IB limits (#4719)" This reverts commit ad58411ddb396aeb196f6a08bd9c4000a0f10361. --- megatron/core/transformer/moe/fused_a2a.py | 71 ++-------------------- 1 file changed, 6 insertions(+), 65 deletions(-) diff --git a/megatron/core/transformer/moe/fused_a2a.py b/megatron/core/transformer/moe/fused_a2a.py index defbe996a72..07f33deca6c 100644 --- a/megatron/core/transformer/moe/fused_a2a.py +++ b/megatron/core/transformer/moe/fused_a2a.py @@ -3,7 +3,6 @@ # Copyright (c) 2025 DeepSeek # Licensed under the MIT License - https://github.com/deepseek-ai/DeepEP/blob/main/LICENSE -import os from typing import Optional from megatron.core.utils import internal_api @@ -269,7 +268,6 @@ def set_deepep_num_sms(num_sms): try: - import hybrid_ep_cpp from deep_ep import HybridEPBuffer HAVE_HYBRIDEP = True @@ -277,68 +275,12 @@ def set_deepep_num_sms(num_sms): HAVE_HYBRIDEP = False _hybrid_ep_buffer = None -_HYBRID_EP_TOKEN_ALIGNMENT = 16 -_HYBRID_EP_MIN_BUFFER_TOKENS = 512 -_HYBRID_EP_IB_QP_MAX_DEPTH = 65535 -_HYBRID_EP_IB_DISPATCH_DEPTH_PER_TOKEN = 3 - - -def _round_up_to_multiple(value: int, multiple: int) -> int: - return ((value + multiple - 1) // multiple) * multiple - - -def _hybrid_ep_num_nodes(group: torch.distributed.ProcessGroup) -> int: - """Mirror HybridEP's NVLink-domain detection without constructing the full buffer.""" - ranks_per_nvlink_domain_env = os.getenv("NUM_OF_HYBRID_EP_RANKS_PER_NVLINK_DOMAIN") - if ranks_per_nvlink_domain_env is not None: - ranks_per_nvlink_domain = int(ranks_per_nvlink_domain_env) - else: - allocator = hybrid_ep_cpp.ExtendedMemoryAllocator() - ranks_per_nvlink_domain = allocator.detect_accessible_ranks(group) - - assert group.size() % ranks_per_nvlink_domain == 0, ( - f"The number of ranks {group.size()} should be divisible by the number of ranks per " - f"NVLink domain {ranks_per_nvlink_domain}." - ) - return group.size() // ranks_per_nvlink_domain - - -def _hybrid_ep_uses_internode_rdma(group: torch.distributed.ProcessGroup) -> bool: - if _hybrid_ep_buffer is not None and hasattr(_hybrid_ep_buffer, "num_of_nodes"): - return _hybrid_ep_buffer.num_of_nodes > 1 - return _hybrid_ep_num_nodes(group) > 1 - - -def _validate_hybrid_ep_ib_tx_depth(num_tokens: int, group: torch.distributed.ProcessGroup) -> None: - buffer_tokens = max( - _round_up_to_multiple(num_tokens, _HYBRID_EP_TOKEN_ALIGNMENT), _HYBRID_EP_MIN_BUFFER_TOKENS - ) - tx_depth = _HYBRID_EP_IB_DISPATCH_DEPTH_PER_TOKEN * buffer_tokens + 1 - if tx_depth <= _HYBRID_EP_IB_QP_MAX_DEPTH: - return - - if not _hybrid_ep_uses_internode_rdma(group): - return - - max_supported_tokens = ( - ((_HYBRID_EP_IB_QP_MAX_DEPTH - 1) // _HYBRID_EP_IB_DISPATCH_DEPTH_PER_TOKEN) - // _HYBRID_EP_TOKEN_ALIGNMENT - * _HYBRID_EP_TOKEN_ALIGNMENT - ) - raise ValueError( - f"HybridEP InfiniBand dispatch queue pair depth ({tx_depth}) exceeds the hardware " - f"limit of {_HYBRID_EP_IB_QP_MAX_DEPTH}. DeepEP computes this depth from the " - f"tokens per rank rounded up to a {_HYBRID_EP_TOKEN_ALIGNMENT}-token buffer " - f"alignment ({buffer_tokens}). Reduce sequence length or micro-batch size, or " - f"increase Tensor Parallelism (TP) / Context Parallelism (CP), so tokens per rank " - f"are at most {max_supported_tokens} for multi-node HybridEP." - ) def init_hybrid_ep_buffer( group: torch.distributed.ProcessGroup, hidden_dim: int, - num_tokens: int, + seq_len: int, num_local_experts: int, num_sms_dispatch_api: Optional[int] = None, num_sms_combine_api: Optional[int] = None, @@ -360,8 +302,8 @@ def init_hybrid_ep_buffer( Process group for HybridEP all-to-all communication. hidden_dim (int): Hidden dimension of the input tensor. - num_tokens (int): - Maximum token count of the input tensor. + seq_len (int): + Maximum sequence length of the input tensor. num_local_experts (int): Number of local experts. num_sms_dispatch_api (Optional[int]): @@ -393,7 +335,7 @@ def init_hybrid_ep_buffer( _hybrid_ep_buffer = HybridEPBuffer( group=group, hidden_dim=hidden_dim, - max_num_of_tokens_per_rank=num_tokens, + max_num_of_tokens_per_rank=seq_len, num_local_experts=num_local_experts, use_fp8=fp8_dispatch, **kwargs, @@ -450,14 +392,13 @@ def forward( num_blocks_permute = None num_blocks_unpermute = None - num_tokens, hidden_dim = x.shape[-2:] - _validate_hybrid_ep_ib_tx_depth(num_tokens, group) if _hybrid_ep_buffer is None: + seq_len, hidden_dim = x.shape[-2:] fp8_dispatch = False # Currently, we do not support fp8 dispatch init_hybrid_ep_buffer( group, hidden_dim, - num_tokens, + seq_len, num_local_experts, num_sms_dispatch_api, num_sms_combine_api, From e007ee710fff76631e3f34c3de0a0fb2456a68f5 Mon Sep 17 00:00:00 2001 From: janEbert Date: Thu, 21 May 2026 11:45:14 +0200 Subject: [PATCH 2/2] Improve variable name Rename `seq_len` -> `num_tokens`. --- megatron/core/transformer/moe/fused_a2a.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/megatron/core/transformer/moe/fused_a2a.py b/megatron/core/transformer/moe/fused_a2a.py index 07f33deca6c..5e97d7a6786 100644 --- a/megatron/core/transformer/moe/fused_a2a.py +++ b/megatron/core/transformer/moe/fused_a2a.py @@ -280,7 +280,7 @@ def set_deepep_num_sms(num_sms): def init_hybrid_ep_buffer( group: torch.distributed.ProcessGroup, hidden_dim: int, - seq_len: int, + num_tokens: int, num_local_experts: int, num_sms_dispatch_api: Optional[int] = None, num_sms_combine_api: Optional[int] = None, @@ -302,8 +302,8 @@ def init_hybrid_ep_buffer( Process group for HybridEP all-to-all communication. hidden_dim (int): Hidden dimension of the input tensor. - seq_len (int): - Maximum sequence length of the input tensor. + num_tokens (int): + Maximum token count of the input tensor. num_local_experts (int): Number of local experts. num_sms_dispatch_api (Optional[int]): @@ -335,7 +335,7 @@ def init_hybrid_ep_buffer( _hybrid_ep_buffer = HybridEPBuffer( group=group, hidden_dim=hidden_dim, - max_num_of_tokens_per_rank=seq_len, + max_num_of_tokens_per_rank=num_tokens, num_local_experts=num_local_experts, use_fp8=fp8_dispatch, **kwargs, @@ -393,12 +393,12 @@ def forward( num_blocks_unpermute = None if _hybrid_ep_buffer is None: - seq_len, hidden_dim = x.shape[-2:] + num_tokens, hidden_dim = x.shape[-2:] fp8_dispatch = False # Currently, we do not support fp8 dispatch init_hybrid_ep_buffer( group, hidden_dim, - seq_len, + num_tokens, num_local_experts, num_sms_dispatch_api, num_sms_combine_api,