From 629b6c6a85b9381695070e161dd49eb2a0dd8fd8 Mon Sep 17 00:00:00 2001 From: weireweire Date: Sat, 23 May 2026 10:18:25 +0800 Subject: [PATCH] correct allreduce fusion and dummy_run alignment in SCATTERED MLP mode (moe_dense_tp_size=1) (#19918) --- python/sglang/srt/layers/communicator.py | 5 +++++ python/sglang/srt/model_executor/model_runner.py | 6 ++++-- test/registered/moe/test_hybrid_dp_ep_tp_mtp.py | 1 + 3 files changed, 10 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 7c40b0c10..09f38265f 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -742,6 +742,11 @@ class LayerCommunicator: else 0 ) + # When mlp_mode is SCATTERED, the MLP runs on scattered data with no TP + # all-reduce, so there is nothing to fuse with the next layer. + if self.layer_scatter_modes.mlp_mode == ScatterMode.SCATTERED: + return False + return ( ( apply_flashinfer_allreduce_fusion(batch_size) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index efe66a10a..fd83f5e64 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -209,6 +209,7 @@ from sglang.srt.utils import ( set_cuda_arch, slow_rank_detector, ) +from sglang.srt.utils.common import ceil_align, require_mlp_sync from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.offloader import ( @@ -2454,10 +2455,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): num_tokens = batch_size * num_tokens_per_bs - if require_gathered_buffer(self.server_args): + # Keep warmup aligned with scheduler MLP-sync padding. + if require_mlp_sync(self.server_args): attn_tp_size = get_attention_tp_size() if attn_tp_size > 1 and num_tokens % attn_tp_size != 0: - num_tokens = num_tokens // attn_tp_size * attn_tp_size + num_tokens = ceil_align(num_tokens, attn_tp_size) batch_size = num_tokens // num_tokens_per_bs seq_len_fill_value = self.attn_backend.get_cuda_graph_seq_len_fill_value() diff --git a/test/registered/moe/test_hybrid_dp_ep_tp_mtp.py b/test/registered/moe/test_hybrid_dp_ep_tp_mtp.py index 06e73443d..9690fc6ec 100644 --- a/test/registered/moe/test_hybrid_dp_ep_tp_mtp.py +++ b/test/registered/moe/test_hybrid_dp_ep_tp_mtp.py @@ -143,6 +143,7 @@ class Test03(CustomTestCase): "8", "--moe-dense-tp-size", "1", + "--enable-flashinfer-allreduce-fusion", ], )