[NVIDIA][comm] Merge EP+MoE-TP post-experts all-reduces into one _TP reduction (#32963)
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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():
|
||||
"""
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user