From 1e8699fda39fdf7b78fa39b8e14962f34b83bd90 Mon Sep 17 00:00:00 2001 From: Shu Wang Date: Fri, 18 Sep 2026 03:35:10 -0500 Subject: [PATCH] [NVIDIA][comm] Merge EP+MoE-TP post-experts all-reduces into one _TP reduction (#32963) --- python/sglang/srt/layers/communicator.py | 25 +- .../srt/layers/flashinfer_comm_fusion.py | 55 ++-- python/sglang/srt/layers/moe/__init__.py | 6 + python/sglang/srt/layers/moe/utils.py | 89 +++++-- python/sglang/srt/models/deepseek_v2.py | 23 +- python/sglang/srt/models/hunyuan_v3.py | 30 +-- python/sglang/srt/models/qwen2_moe.py | 19 +- python/sglang/srt/models/qwen3_moe.py | 16 +- .../test_layer_communicator_fusion_gate.py | 243 +++++++++++++++++- 9 files changed, 376 insertions(+), 130 deletions(-) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 05af063fb..a60b70f22 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -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 ( diff --git a/python/sglang/srt/layers/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index 3d6e8770f..48cf1bc0c 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -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") diff --git a/python/sglang/srt/layers/moe/__init__.py b/python/sglang/srt/layers/moe/__init__.py index 4dc3992f2..f3b05fa1e 100644 --- a/python/sglang/srt/layers/moe/__init__.py +++ b/python/sglang/srt/layers/moe/__init__.py @@ -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", diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 1876e158d..d6c3ea032 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -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(): """ diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 1f12cea55..7c2cd05a0 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -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( diff --git a/python/sglang/srt/models/hunyuan_v3.py b/python/sglang/srt/models/hunyuan_v3.py index 6dedbe220..f95840dd9 100644 --- a/python/sglang/srt/models/hunyuan_v3.py +++ b/python/sglang/srt/models/hunyuan_v3.py @@ -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) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 83aad13eb..6e7b6363c 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -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( { diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index c0e93c371..5e204ecf3 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -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) diff --git a/test/registered/unit/layers/test_layer_communicator_fusion_gate.py b/test/registered/unit/layers/test_layer_communicator_fusion_gate.py index aa52f9903..1732571fe 100644 --- a/test/registered/unit/layers/test_layer_communicator_fusion_gate.py +++ b/test/registered/unit/layers/test_layer_communicator_fusion_gate.py @@ -1,9 +1,18 @@ +import contextlib import types import unittest -from unittest.mock import patch +from unittest.mock import MagicMock, patch + +import torch from sglang.srt.layers import communicator as comm from sglang.srt.layers.communicator import LayerCommunicator, ScatterMode +from sglang.srt.layers.moe import ( + can_merge_post_experts_all_reduce, + deferred_post_experts_all_reduce, + post_experts_all_reduce, +) +from sglang.srt.layers.moe import utils as moe_utils from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -20,19 +29,218 @@ def _fake_communicator(mlp_mode=ScatterMode.TP_ATTN_FULL): ) -class TestFuseMlpAllReduceGate(CustomTestCase): - """Hybrid EP+TP must not fuse the post-experts all-reduce away. +@contextlib.contextmanager +def _recorded_all_reduces(called, *, moe_ep_size, moe_tp_size, moe_dp_size): + """Log which group each all-reduce helper reduces over, under a fixed topo.""" - The fused residual+LN reduces over a single group, but with moe_ep_size > 1 - and moe_tp_size > 1 the post-experts reduction spans two disjoint groups - (_MOE_EP then _MOE_TP) and should_skip_post_experts_all_reduce() drops both - once fusion is published. The result is activations reduced over only half - the peers -- wrong output, no crash. Observed as garbage completions on - Qwen3-30B-A3B with --tp-size 4 --ep-size 2. + def record(name): + return lambda x: called.append(name) or x + + with ( + patch( + "sglang.srt.distributed.communication_op.tensor_model_parallel_all_reduce", + side_effect=record("tp"), + ), + patch( + "sglang.srt.distributed.communication_op.moe_expert_parallel_all_reduce", + side_effect=record("ep"), + ), + patch( + "sglang.srt.distributed.communication_op.moe_tensor_model_parallel_all_reduce", + side_effect=record("moe_tp"), + ), + get_parallel().override( + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + tp_size=moe_ep_size * moe_tp_size * moe_dp_size, + ), + ): + yield + + +class TestPostExpertsAllReduceMerge(CustomTestCase): + """The two post-experts reductions collapse into one _TP reduction. + + _MOE_EP and _MOE_TP are orthogonal subgroups of _TP, so with + moe_dp_size == 1 reducing over each in turn equals one _TP reduction -- + one collective instead of two. With moe_dp_size > 1 they cover only part of + _TP and merging would sum across DP replicas, which hold different tokens. + """ + + def _calls(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1, skip=False): + called = [] + with ( + patch.object( + moe_utils, "should_skip_post_experts_all_reduce", return_value=skip + ), + _recorded_all_reduces( + called, + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + ), + ): + post_experts_all_reduce(torch.zeros(2, 2)) + return called + + def test_hybrid_issues_one_tp_reduction(self): + self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2), ["tp"]) + + def test_moe_dp_keeps_the_two_step_form(self): + # Server args reject moe_ep_size > 1 together with moe_tp_size > 1 and + # moe_dp_size > 1 (they force ep_size * moe_dp_size == tp_size), so this + # pins the guard rather than a topology that can be launched today. + self.assertEqual( + self._calls(moe_ep_size=2, moe_tp_size=2, moe_dp_size=2), ["ep", "moe_tp"] + ) + + def test_single_dimension_issues_one_reduction(self): + self.assertEqual(self._calls(moe_ep_size=1, moe_tp_size=4), ["moe_tp"]) + self.assertEqual(self._calls(moe_ep_size=4, moe_tp_size=1), ["ep"]) + + def test_skipped_when_deferred_to_fusion(self): + self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2, skip=True), []) + + +class TestDeferredPostExpertsAllReduce(CustomTestCase): + """The inline fallback must reduce over the same peers the fused kernel would. + + _MOE_TP holds a single rank under pure EP, so reducing over it there is a + no-op that drops the deferred reduction instead of performing it. + """ + + def _calls(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1): + called = [] + with _recorded_all_reduces( + called, + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + ): + deferred_post_experts_all_reduce(torch.zeros(2, 2)) + return called + + def test_hybrid_reduces_over_tp(self): + self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2), ["tp"]) + + def test_pure_ep_reduces_over_ep(self): + self.assertEqual(self._calls(moe_ep_size=4, moe_tp_size=1), ["ep"]) + + def test_pure_tp_reduces_over_moe_tp(self): + self.assertEqual(self._calls(moe_ep_size=1, moe_tp_size=4), ["moe_tp"]) + + def test_moe_dp_reduces_over_moe_tp(self): + self.assertEqual( + self._calls(moe_ep_size=1, moe_tp_size=2, moe_dp_size=2), ["moe_tp"] + ) + + +class TestCanMergePostExpertsAllReduce(CustomTestCase): + def _can_merge(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1): + with get_parallel().override( + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + tp_size=moe_ep_size * moe_tp_size * moe_dp_size, + ): + return can_merge_post_experts_all_reduce() + + def test_hybrid_ep_tp_merges(self): + self.assertTrue(self._can_merge(moe_ep_size=2, moe_tp_size=2)) + + def test_single_dimension_does_not_merge(self): + self.assertFalse(self._can_merge(moe_ep_size=1, moe_tp_size=4)) + self.assertFalse(self._can_merge(moe_ep_size=4, moe_tp_size=1)) + + +class TestResolveFusionGroup(CustomTestCase): + """EP2/MoE-TP2/DP1 (e.g. DeepSeek-V4-Flash with --tp-size 4 --ep-size 2) must + resolve to the _TP group with world_size=4 and the TP rank.""" + + def _resolve(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1, tp_rank=0): + from sglang.srt.layers.flashinfer_comm_fusion import ( + resolve_fusion_group, + resolve_fusion_world_size, + ) + + fake_tp_group = MagicMock(name="tp_group") + fake_ep_group = MagicMock(name="ep_group") + fake_moe_tp_group = MagicMock(name="moe_tp_group") + tp_size = moe_ep_size * moe_tp_size * moe_dp_size + + with ( + get_parallel().override( + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + tp_size=tp_size, + tp_rank=tp_rank, + moe_ep_rank=tp_rank % moe_ep_size, + moe_tp_rank=tp_rank % moe_tp_size, + ), + patch( + "sglang.srt.layers.flashinfer_comm_fusion.get_tp_group", + return_value=fake_tp_group, + ), + patch( + "sglang.srt.layers.flashinfer_comm_fusion.get_moe_ep_group", + return_value=fake_ep_group, + ), + patch( + "sglang.srt.layers.flashinfer_comm_fusion.get_moe_tp_group", + return_value=fake_moe_tp_group, + ), + ): + ws = resolve_fusion_world_size(use_attn_tp_group=False) + group_tuple = resolve_fusion_group(use_attn_tp_group=False) + return ws, group_tuple, (fake_tp_group, fake_ep_group, fake_moe_tp_group) + + def test_hybrid_ep2_tp2_dp1_resolves_to_tp_ws4(self): + # EP2/MoE-TP2/DP1 (DeepSeek-V4-Flash on 4 GPUs): workspace must sit on + # _TP (ws=4) so the fused kernel reduces over all 4 peers. + ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve( + moe_ep_size=2, moe_tp_size=2, moe_dp_size=1, tp_rank=3 + ) + self.assertEqual(ws, 4) + self.assertEqual(size, 4) + self.assertEqual(rank, 3) + self.assertIs(group, tp_grp) + + def test_pure_ep_resolves_to_ep_group(self): + ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve( + moe_ep_size=4, moe_tp_size=1, moe_dp_size=1, tp_rank=2 + ) + self.assertEqual(ws, 4) + self.assertEqual(size, 4) + self.assertIs(group, ep_grp) + + def test_pure_tp_resolves_to_moe_tp_group(self): + ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve( + moe_ep_size=1, moe_tp_size=4, moe_dp_size=1, tp_rank=1 + ) + self.assertEqual(ws, 4) + self.assertEqual(size, 4) + self.assertIs(group, moe_tp_grp) + + +class TestFuseMlpAllReduceGate(CustomTestCase): + """Fusion is allowed only when one group covers the whole reduction. + + The fused residual+LN reduces over a single group. Hybrid EP+TP produces two + reductions over disjoint groups; merging collapses them to one _TP reduction + that the fused kernel can absorb. When merging does not apply + (moe_dp_size > 1) there is no such group and fusion must stay off -- + otherwise the fused reduce covers half the peers and silently under-reduces. """ def _should_fuse( - self, *, moe_ep_size, moe_tp_size, mlp_mode=ScatterMode.TP_ATTN_FULL + self, + *, + moe_ep_size, + moe_tp_size, + moe_dp_size=1, + mlp_mode=ScatterMode.TP_ATTN_FULL, ): forward_batch = types.SimpleNamespace( input_ids=types.SimpleNamespace(shape=(8,)) @@ -46,15 +254,24 @@ class TestFuseMlpAllReduceGate(CustomTestCase): return_value=types.SimpleNamespace(input_scattered=False), ), get_parallel().override( - moe_ep_size=moe_ep_size, moe_tp_size=moe_tp_size, tp_size=4 + moe_ep_size=moe_ep_size, + moe_tp_size=moe_tp_size, + moe_dp_size=moe_dp_size, + tp_size=moe_ep_size * moe_tp_size * moe_dp_size, ), ): return LayerCommunicator.should_fuse_mlp_allreduce_with_next_layer( _fake_communicator(mlp_mode), forward_batch ) - def test_hybrid_ep_tp_does_not_fuse(self): - self.assertFalse(self._should_fuse(moe_ep_size=2, moe_tp_size=2)) + def test_hybrid_ep_tp_fuses_when_mergeable(self): + self.assertTrue(self._should_fuse(moe_ep_size=2, moe_tp_size=2)) + + def test_hybrid_ep_tp_does_not_fuse_when_moe_dp_blocks_the_merge(self): + # Same caveat as test_moe_dp_keeps_the_two_step_form: unreachable today, + # kept so a future relaxation cannot silently re-enable fusion over a + # reduction that no single group covers. + self.assertFalse(self._should_fuse(moe_ep_size=2, moe_tp_size=2, moe_dp_size=2)) def test_pure_tp_still_fuses(self): self.assertTrue(self._should_fuse(moe_ep_size=1, moe_tp_size=4))