[NVIDIA][comm] Merge EP+MoE-TP post-experts all-reduces into one _TP reduction (#32963)

This commit is contained in:
Shu Wang
2026-09-18 01:35:10 -07:00
committed by GitHub
parent 8ac39c66d8
commit 1e8699fda3
9 changed files with 376 additions and 130 deletions
+14 -11
View File
@@ -24,7 +24,6 @@ from sglang.srt.distributed import (
attention_tensor_model_parallel_all_reduce,
attention_tensor_model_parallel_quant_all_reduce,
get_tp_group,
moe_tensor_model_parallel_all_reduce,
tensor_model_parallel_all_reduce,
)
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
@@ -60,6 +59,8 @@ from sglang.srt.layers.dp_attention import (
)
from sglang.srt.layers.flashinfer_comm_fusion import is_flashinfer_allreduce_unavailable
from sglang.srt.layers.moe import (
can_merge_post_experts_all_reduce,
deferred_post_experts_all_reduce,
get_moe_a2a_backend,
should_use_dp_reduce_scatterv,
should_use_flashinfer_cutlass_moe_fp4_allgather,
@@ -724,7 +725,9 @@ class LayerCommunicator:
)
)
else:
hidden_states = moe_tensor_model_parallel_all_reduce(hidden_states)
# Fusion was published but this shape can't use the kernel,
# so run the deferred reduction inline.
hidden_states = deferred_post_experts_all_reduce(hidden_states)
hidden_states, residual = self.input_layernorm(
hidden_states, residual
)
@@ -938,16 +941,16 @@ class LayerCommunicator:
):
return False
# Fusing makes the next layer's residual+LN absorb the post-experts
# all-reduce, and that fused kernel reduces over a single group. Under
# hybrid EP+TP the post-experts reduction spans two disjoint groups
# (moe_expert_parallel_all_reduce over _MOE_EP, then
# moe_tensor_model_parallel_all_reduce over _MOE_TP), and
# should_skip_post_experts_all_reduce() skips *both* once fusion is
# published -- so the fused reduce would cover only half the peers and
# silently return under-reduced activations.
# The fused residual+LN reduces over a single group. Hybrid EP+TP spans
# two disjoint groups; post_experts_all_reduce() merges them into one
# _TP reduction when moe_dp_size == 1, which the fused kernel can absorb.
# When merging is blocked, no single group covers both, so fusion stays off.
parallel = get_parallel()
if parallel.moe_ep_size > 1 and parallel.moe_tp_size > 1:
if (
parallel.moe_ep_size > 1
and parallel.moe_tp_size > 1
and not can_merge_post_experts_all_reduce()
):
return False
if (
@@ -647,6 +647,37 @@ def _get_workspace_manager(use_attn_tp_group: bool) -> FlashInferWorkspaceManage
return manager
def resolve_fusion_world_size(*, use_attn_tp_group: bool) -> int:
"""Peer count of the fusion group. Reads sizes only -- deliberately does not
reach for a group coordinator, which some callers hit before one exists."""
from sglang.srt.layers.moe.utils import can_merge_post_experts_all_reduce
parallel = get_parallel()
if use_attn_tp_group:
return parallel.attn_tp_size
if can_merge_post_experts_all_reduce():
return parallel.tp_size
return parallel.moe_ep_size if parallel.moe_ep_size > 1 else parallel.moe_tp_size
def resolve_fusion_group(*, use_attn_tp_group: bool):
"""Return (world_size, rank, coordinator) for the fusion workspace.
Must match the group the fused residual+LN kernel reduces over; a mismatch
silently reduces across the wrong peers.
"""
from sglang.srt.layers.moe.utils import can_merge_post_experts_all_reduce
parallel = get_parallel()
if use_attn_tp_group:
return parallel.attn_tp_size, parallel.attn_tp_rank, get_attn_tp_group()
if can_merge_post_experts_all_reduce():
return parallel.tp_size, parallel.tp_rank, get_tp_group()
if parallel.moe_ep_size > 1:
return parallel.moe_ep_size, parallel.moe_ep_rank, get_moe_ep_group()
return parallel.moe_tp_size, parallel.moe_tp_rank, get_moe_tp_group()
def _sync_allreduce_unavailable_across_tp():
"""Synchronize _flashinfer_allreduce_unavailable across all TP ranks.
@@ -694,19 +725,9 @@ def ensure_workspace_initialized(
if not is_flashinfer_available() or _flashinfer_comm is None:
return False
if use_attn_tp_group:
world_size = get_parallel().attn_tp_size
rank = get_parallel().attn_tp_rank
coordinator = get_attn_tp_group()
else:
if get_parallel().moe_ep_size > 1:
world_size = get_parallel().moe_ep_size
rank = get_parallel().moe_ep_rank
coordinator = get_moe_ep_group()
else:
world_size = get_parallel().moe_tp_size
rank = get_parallel().moe_tp_rank
coordinator = get_moe_tp_group()
world_size, rank, coordinator = resolve_fusion_group(
use_attn_tp_group=use_attn_tp_group
)
# Always pass the coordinator's groups: flashinfer >=0.6.10 reads the
# rendezvous group from `group=...` (falling back to WORLD when None),
@@ -814,13 +835,7 @@ def flashinfer_allreduce_residual_rmsnorm(
)
return None, None
if use_attn_tp_group:
world_size = get_parallel().attn_tp_size
else:
if get_parallel().moe_ep_size > 1:
world_size = get_parallel().moe_ep_size
else:
world_size = get_parallel().moe_tp_size
world_size = resolve_fusion_world_size(use_attn_tp_group=use_attn_tp_group)
if world_size <= 1:
logger.debug("Single GPU, no need for allreduce fusion")
+6
View File
@@ -3,6 +3,8 @@ from sglang.srt.layers.moe.utils import (
DeepEPMode,
MoeA2ABackend,
MoeRunnerBackend,
can_merge_post_experts_all_reduce,
deferred_post_experts_all_reduce,
get_deepep_config,
get_deepep_mode,
get_moe_a2a_backend,
@@ -11,6 +13,7 @@ from sglang.srt.layers.moe.utils import (
initialize_moe_config,
is_moe_input_scattered_across_dp_ranks,
is_tbo_enabled,
post_experts_all_reduce,
should_skip_mlp_all_reduce,
should_skip_post_experts_all_reduce,
should_use_dp_reduce_scatterv,
@@ -28,6 +31,9 @@ __all__ = [
"get_moe_runner_backend",
"get_deepep_mode",
"should_skip_mlp_all_reduce",
"can_merge_post_experts_all_reduce",
"deferred_post_experts_all_reduce",
"post_experts_all_reduce",
"should_skip_post_experts_all_reduce",
"should_use_dp_reduce_scatterv",
"should_use_flashinfer_cutlass_moe_fp4_allgather",
+66 -23
View File
@@ -731,30 +731,9 @@ def should_skip_mlp_all_reduce() -> bool:
def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool:
"""Whether to skip the post-experts all-reduce (EP or TP) because a
downstream component will fuse, replace, or absorb it.
"""Whether a downstream component will fuse, replace, or absorb the post-experts all-reduce.
Skip reasons, in order:
- ``get_forward().fuse_mlp_allreduce``: LayerCommunicator will fuse the
all-reduce with the next layer's residual all-reduce.
- ``get_forward().mlp_reduce_scatter``: LayerCommunicator's post-attention
scatter will do reduce-scatter, which would double-reduce on top of
an all-reduce.
- ``should_use_dp_reduce_scatterv()``: the standard dispatcher's combine
path replaces the all-reduce with a reduce-scatterv.
- ``should_use_flashinfer_cutlass_moe_fp4_allgather()`` (TP path only):
the flashinfer cutlass FP4 kernel performs an all-gather that absorbs
the post-experts TP all-reduce. Not relevant to the EP all-reduce.
- ``get_moe_a2a_backend().is_flashinfer()``: the flashinfer A2A
dispatcher's ``MoeAlltoAll.combine`` already alltoall-reduces partial
MoE outputs back to the source rank, so any further EP/TP all-reduce
would double-count and overflow BF16. Mirrors TRTLLM's
``not enable_alltoall`` gate
(``tensorrt_llm/_torch/modules/fused_moe/interface.py:879``).
The first two reasons come from per-layer ``ForwardFlags`` published by
the decoder via ``get_forward().scoped(...)``. Pass ``is_tp_path=True``
for the post-experts TP all-reduce, ``False`` for the EP all-reduce.
Pass ``is_tp_path=True`` for the TP all-reduce, ``False`` for the EP one.
"""
if should_skip_mlp_all_reduce():
return True
@@ -778,6 +757,70 @@ def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool:
return False
def can_merge_post_experts_all_reduce() -> bool:
"""Whether the EP and MoE-TP reductions can collapse into one _TP all-reduce.
True when moe_dp_size == 1: the two groups are an orthogonal decomposition
of _TP, so reducing over each in turn equals one _TP reduction.
"""
parallel = get_parallel()
return (
parallel.moe_ep_size > 1
and parallel.moe_tp_size > 1
and parallel.moe_dp_size == 1
)
def post_experts_all_reduce(hidden_states: torch.Tensor) -> torch.Tensor:
"""Reduce the post-experts MoE output across the EP and MoE-TP groups.
When both are live and mergeable, issues one _TP all-reduce instead of two
sequential ones, which also restores the invariant the fused residual+LN path
depends on.
"""
from sglang.srt.distributed.communication_op import (
moe_expert_parallel_all_reduce,
moe_tensor_model_parallel_all_reduce,
tensor_model_parallel_all_reduce,
)
parallel = get_parallel()
reduce_ep = parallel.moe_ep_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=False
)
reduce_tp = parallel.moe_tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True
)
if reduce_ep and reduce_tp and can_merge_post_experts_all_reduce():
return tensor_model_parallel_all_reduce(hidden_states)
if reduce_ep:
hidden_states = moe_expert_parallel_all_reduce(hidden_states)
if reduce_tp:
hidden_states = moe_tensor_model_parallel_all_reduce(hidden_states)
return hidden_states
def deferred_post_experts_all_reduce(hidden_states: torch.Tensor) -> torch.Tensor:
"""Run the post-experts reduction that was deferred to allreduce fusion.
Called when the fused residual+LN kernel cannot service the shape. Reduces
over the same group ``resolve_fusion_group`` builds the workspace on.
"""
from sglang.srt.distributed.communication_op import (
moe_expert_parallel_all_reduce,
moe_tensor_model_parallel_all_reduce,
tensor_model_parallel_all_reduce,
)
if can_merge_post_experts_all_reduce():
return tensor_model_parallel_all_reduce(hidden_states)
if get_parallel().moe_ep_size > 1:
return moe_expert_parallel_all_reduce(hidden_states)
return moe_tensor_model_parallel_all_reduce(hidden_states)
@contextmanager
def speculative_moe_backend_context():
"""
+5 -18
View File
@@ -51,11 +51,7 @@ from sglang.srt.configs.model_config import (
is_deepseek_dsa,
is_glm_moe_dsa,
)
from sglang.srt.distributed import (
divide,
get_pp_group,
tensor_model_parallel_all_reduce,
)
from sglang.srt.distributed import divide, get_pp_group
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
@@ -95,7 +91,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import (
get_moe_a2a_backend,
get_moe_runner_backend,
should_skip_post_experts_all_reduce,
post_experts_all_reduce,
should_use_flashinfer_cutlass_moe_fp4_allgather,
)
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
@@ -1037,10 +1033,7 @@ class DeepseekV2MoE(nn.Module):
self.routed_scaling_factor,
)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
):
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
# TP1 shared experts are replicated, so add them after all-reduce to
# avoid summing the same shared output once per TP rank.
if self._shared_expert_tp1:
@@ -1179,10 +1172,7 @@ class DeepseekV2MoE(nn.Module):
self.routed_scaling_factor,
)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
):
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
# TP1 shared experts are replicated, so add them after all-reduce to
# avoid summing the same shared output once per TP rank.
if shared_output is not None and self._shared_expert_tp1:
@@ -1240,10 +1230,7 @@ class DeepseekV2MoE(nn.Module):
), # block_size
True, # is_vnni
)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
):
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
return final_hidden_states
def forward_deepep(
+3 -27
View File
@@ -18,10 +18,6 @@ import torch
from torch import nn
from transformers import PretrainedConfig
from sglang.srt.distributed import (
moe_expert_parallel_all_reduce,
moe_tensor_model_parallel_all_reduce,
)
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
@@ -31,7 +27,7 @@ from sglang.srt.layers.linear import (
RowParallelLinear,
)
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import should_skip_post_experts_all_reduce
from sglang.srt.layers.moe import post_experts_all_reduce
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -191,17 +187,7 @@ class HYV3MoEFused(nn.Module):
hidden_states=hidden_states, topk_output=topk_output
)
if self.ep_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=False,
):
final_hidden_states = moe_expert_parallel_all_reduce(final_hidden_states)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
):
final_hidden_states = moe_tensor_model_parallel_all_reduce(
final_hidden_states
)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
return final_hidden_states.view(orig_shape)
@@ -226,17 +212,7 @@ class HYV3MoEFused(nn.Module):
current_stream.wait_stream(self.alt_stream)
final_hidden_states = final_hidden_states + shared_output
if self.ep_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=False,
):
final_hidden_states = moe_expert_parallel_all_reduce(final_hidden_states)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
):
final_hidden_states = moe_tensor_model_parallel_all_reduce(
final_hidden_states
)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
return final_hidden_states.view(orig_shape)
+15 -4
View File
@@ -60,6 +60,7 @@ from sglang.srt.layers.linear import (
)
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import (
can_merge_post_experts_all_reduce,
get_moe_a2a_backend,
should_skip_post_experts_all_reduce,
)
@@ -1161,10 +1162,20 @@ class Qwen2MoeModel(nn.Module):
and hasattr(hidden_states, "_sglang_needs_allreduce_fusion")
and hidden_states._sglang_needs_allreduce_fusion
):
if get_parallel().moe_ep_size > 1:
hidden_states = moe_expert_parallel_all_reduce(hidden_states)
if get_parallel().moe_tp_size > 1:
hidden_states = moe_tensor_model_parallel_all_reduce(hidden_states)
# The deferred reduction the next layer would have fused; no
# layer follows on this rank, so run it here. Unconditional --
# the skip flags that deferred it are what got us into this
# branch -- so it bypasses post_experts_all_reduce()'s guards
# while reusing its merge rule.
if can_merge_post_experts_all_reduce():
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
else:
if get_parallel().moe_ep_size > 1:
hidden_states = moe_expert_parallel_all_reduce(hidden_states)
if get_parallel().moe_tp_size > 1:
hidden_states = moe_tensor_model_parallel_all_reduce(
hidden_states
)
hidden_states._sglang_needs_allreduce_fusion = False
return PPProxyTensors(
{
+2 -14
View File
@@ -28,8 +28,6 @@ from transformers import PretrainedConfig
from sglang.srt.distributed import (
get_pp_group,
moe_expert_parallel_all_reduce,
moe_tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
@@ -44,7 +42,7 @@ from sglang.srt.layers.linear import (
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import (
get_moe_a2a_backend,
should_skip_post_experts_all_reduce,
post_experts_all_reduce,
)
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
@@ -335,17 +333,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
topk_output = self.topk.empty_topk_output(hidden_states.device)
final_hidden_states = self.experts(hidden_states, topk_output)
if self.ep_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=False
):
final_hidden_states = moe_expert_parallel_all_reduce(final_hidden_states)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True
):
final_hidden_states = moe_tensor_model_parallel_all_reduce(
final_hidden_states
)
final_hidden_states = post_experts_all_reduce(final_hidden_states)
return final_hidden_states.view(num_tokens, hidden_dim)