diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 84a7a4f42282..2057a33ceb0f 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -833,6 +833,7 @@ def forward( for i in range(self.start_layer, self.end_layer): with get_global_expert_distribution_recorder().with_current_layer(i): if i in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append( hidden_states if residual is None else hidden_states + residual ) diff --git a/python/sglang/srt/models/bailing_moe_v3.py b/python/sglang/srt/models/bailing_moe_v3.py index c5143e8f5c89..d5c5b81d525a 100644 --- a/python/sglang/srt/models/bailing_moe_v3.py +++ b/python/sglang/srt/models/bailing_moe_v3.py @@ -1387,6 +1387,7 @@ def forward( and i in self.layers_to_capture and hidden_states.shape[0] != 0 ): + hidden_states = complete_deferred_allreduce(hidden_states) if residual is None: dspark_aux_hidden_states.append(hidden_states) else: diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 6dd084c48c35..99efde41bbdd 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -1116,6 +1116,7 @@ def forward( for i in range(normal_start_layer, normal_end_layer): with get_global_expert_distribution_recorder().with_current_layer(i): if i in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append(hidden_states + residual) layer = self.layers[i] hidden_states, residual = layer( diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index fb01414ad5ad..7eb44eb1e59f 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -813,6 +813,7 @@ def forward( for i in range(normal_start_layer, normal_end_layer): with get_global_expert_distribution_recorder().with_current_layer(i): if i in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append(hidden_states + residual) layer = self.layers[i] hidden_states, residual = layer( diff --git a/python/sglang/srt/models/glm5_next.py b/python/sglang/srt/models/glm5_next.py index 90428a02e2b1..3f1879cc2934 100644 --- a/python/sglang/srt/models/glm5_next.py +++ b/python/sglang/srt/models/glm5_next.py @@ -1182,6 +1182,7 @@ def forward( ) with ctx: if i in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_state = self._prepare_aux_hidden_state( hidden_states, residual ) diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 01e87aadd8fb..b876a1fe0d3d 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -731,6 +731,7 @@ def forward( positions, hidden_states, forward_batch, residual ) if i + 1 in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append( hidden_states + residual if residual is not None diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index c0c4464f9054..3d45cf055726 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -597,6 +597,7 @@ def forward( aux_hidden_states = [] for i in range(self.start_layer, self.end_layer): if i in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append( hidden_states + residual if residual is not None else hidden_states ) diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index 7659a2087b4d..4b3ca9498141 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -1938,6 +1938,7 @@ def forward( and layer_idx < 3 ): sep = self.hidden_size * layer_idx + hidden_states = complete_deferred_allreduce(hidden_states) hidden_states.add_( input_deepstack_embeds[:, sep : sep + self.hidden_size] ) diff --git a/python/sglang/srt/models/qwen3_vl_moe.py b/python/sglang/srt/models/qwen3_vl_moe.py index e4147c0f0849..e319d8b57cff 100644 --- a/python/sglang/srt/models/qwen3_vl_moe.py +++ b/python/sglang/srt/models/qwen3_vl_moe.py @@ -107,6 +107,7 @@ def forward( ): layer_idx += self.start_layer if layer_idx in self.layers_to_capture: + hidden_states = complete_deferred_allreduce(hidden_states) aux_hidden_states.append( hidden_states + residual if residual is not None else hidden_states ) @@ -139,6 +140,7 @@ def forward( and layer_idx in self.deepstack_embed_to_decoder_layer ): sep = self.hidden_size * layer_idx + hidden_states = complete_deferred_allreduce(hidden_states) hidden_states.add_( input_deepstack_embeds[:, sep : sep + self.hidden_size] ) diff --git a/test/registered/unit/models/test_aux_capture_deferred_allreduce.py b/test/registered/unit/models/test_aux_capture_deferred_allreduce.py new file mode 100644 index 000000000000..122527c5ec4a --- /dev/null +++ b/test/registered/unit/models/test_aux_capture_deferred_allreduce.py @@ -0,0 +1,149 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch +from torch import nn + +from sglang.srt.layers import communicator as comm +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.models.bailing_moe import BailingMoEModel +from sglang.srt.models.bailing_moe_v3 import BailingMoELinearModel +from sglang.srt.models.glm4_moe import Glm4MoeModel +from sglang.srt.models.glm4_moe_lite import Glm4MoeLiteModel +from sglang.srt.models.glm5_next import Glm5NextModel +from sglang.srt.models.gpt_oss import GptOssModel +from sglang.srt.models.laguna import LagunaModel +from sglang.srt.models.qwen3_vl_moe import Qwen3MoeLLMModel +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + +NUM_LAYERS = 4 +MODELS = ( + BailingMoEModel, + BailingMoELinearModel, + Glm4MoeModel, + Glm4MoeLiteModel, + Glm5NextModel, + GptOssModel, + LagunaModel, + Qwen3MoeLLMModel, +) + + +def all_reduce(hidden_states): + return hidden_states * 2 + + +class DeferringLayer(nn.Module): + def __init__(self, defer, return_topk=False): + super().__init__() + self.return_topk = return_topk + self.layer_communicator = comm.LayerCommunicator.__new__(comm.LayerCommunicator) + self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer = ( + lambda batch: defer + ) + self.layer_communicator.should_use_reduce_scatter = lambda batch: False + self.layer_communicator.postprocess_layer = lambda hidden, residual, batch: ( + all_reduce(hidden), + residual, + ) + + def forward( + self, + positions=None, + hidden_states=None, + forward_batch=None, + residual=None, + *args, + **kwargs, + ): + hidden_states = comm.complete_deferred_allreduce(hidden_states) + if residual is None: + residual = hidden_states.clone() + else: + # later norms mutate the residual; captured snapshots must stay intact + residual.add_(hidden_states) + with self.layer_communicator.ffn_exit(forward_batch) as ffn_exit: + partial = torch.full_like(hidden_states, 0.5) + hidden_states, residual = ffn_exit.finish(partial, residual) + if self.return_topk: + return hidden_states, residual, None + return hidden_states, residual + + +class SumNorm(nn.Module): + def forward(self, hidden_states, residual=None, **kwargs): + if residual is None: + return hidden_states + return hidden_states + residual, residual + + +def build_model(model_cls, *, defer, capture): + model = model_cls.__new__(model_cls) + nn.Module.__init__(model) + model.pp_group = SimpleNamespace(is_first_rank=True, is_last_rank=True) + model.start_layer = 0 + model.end_layer = NUM_LAYERS + model.first_k_dense_replace = 0 + model.layers_to_capture = [1, 2, 3] if capture else [] + model.capture_aux_hidden_states = capture + model.use_hf_deepstack_order = False + model.dflash_capture = capture + model.enable_a2a_moe = False + model.config = SimpleNamespace(mhc=False) + if model_cls is BailingMoELinearModel: + # Bailing-v3 captures after the layer; other models use boundary indices + model.layers_to_capture = [0, 1, 2] if capture else [] + model.layers = nn.ModuleList( + DeferringLayer( + defer and i < NUM_LAYERS - 1, return_topk=model_cls is Glm5NextModel + ) + for i in range(NUM_LAYERS) + ) + model.norm = SumNorm() + return model + + +class TestAuxCaptureDeferredAllreduce(CustomTestCase): + def test_capture_matches_eager_reduction(self): + inputs = torch.tensor([[0.25, -0.5, 0.75, 1.0], [1.5, 2.0, -1.0, 0.0]]) + batch = SimpleNamespace( + can_run_tbo=False, + forward_mode=ForwardMode.DECODE, + capture_hidden_mode=SimpleNamespace(need_capture=lambda: True), + ) + for model_cls in MODELS: + for defer in (False, True): + for capture in (False, True): + with ( + self.subTest( + model=model_cls.__name__, defer=defer, capture=capture + ), + patch.object( + comm, + "deferred_post_experts_all_reduce", + side_effect=all_reduce, + ) as reduce, + ): + model = build_model(model_cls, defer=defer, capture=capture) + result = model( + input_ids=None, + positions=None, + forward_batch=batch, + input_embeds=inputs.clone(), + ) + output, snapshots = result if capture else (result, []) + torch.testing.assert_close(output, inputs + NUM_LAYERS) + self.assertEqual(len(snapshots), 3 if capture else 0) + for boundary, snapshot in enumerate(snapshots, 1): + torch.testing.assert_close(snapshot, inputs + boundary) + self.assertEqual( + reduce.call_count, NUM_LAYERS - 1 if defer else 0 + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/models/test_deepstack_deferred_allreduce.py b/test/registered/unit/models/test_deepstack_deferred_allreduce.py new file mode 100644 index 000000000000..4172e17bdbdb --- /dev/null +++ b/test/registered/unit/models/test_deepstack_deferred_allreduce.py @@ -0,0 +1,139 @@ +"""Deepstack visual embeddings are added to a decoder layer's output. When the +layer left its FFN all-reduce to the next layer, that output is one rank's +partial sum, and an embedding added to it is counted once per rank.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch +from torch import nn + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + +TP_SIZE = 2 +HIDDEN = 4 +TOKENS = 3 +NUM_LAYERS = 4 +MARKER = "_sglang_needs_allreduce_fusion" + + +def all_reduce(hidden_states): + # Every rank holds the same partial sum in this single-process stand-in. + return hidden_states * TP_SIZE + + +class DeferringLayer(nn.Module): + """One TP rank's view of a decoder layer whose FFN output sums to one. It + completes a reduction left by the previous layer, and leaves its own to the + next layer unless it is the last.""" + + def __init__(self, is_last_layer): + super().__init__() + self.is_last_layer = is_last_layer + + def forward( + self, positions=None, hidden_states=None, forward_batch=None, residual=None, **_ + ): + if getattr(hidden_states, MARKER, False): + hidden_states = all_reduce(hidden_states) + residual = hidden_states if residual is None else hidden_states + residual + if self.is_last_layer: + return torch.ones_like(residual), residual + partial = torch.full_like(residual, 1 / TP_SIZE) + setattr(partial, MARKER, True) + return partial, residual + + +class SumNorm(nn.Module): + def forward(self, hidden_states, residual=None, post_residual_addition=None): + if residual is None: + return hidden_states + return hidden_states + residual, residual + + +def stub_model(cls, **attrs): + model = cls.__new__(cls) + nn.Module.__init__(model) + layers = [DeferringLayer(i == NUM_LAYERS - 1) for i in range(NUM_LAYERS)] + common = dict( + pp_group=SimpleNamespace(is_first_rank=True, is_last_rank=True), + layers=layers, + layers_to_capture=[], + hidden_size=HIDDEN, + norm=SumNorm(), + ) + for name, value in {**common, **attrs}.items(): + object.__setattr__(model, name, value) + return model + + +def qwen3_vl_moe(): + from sglang.srt.models.qwen3_vl_moe import Qwen3MoeLLMModel + + return stub_model( + Qwen3MoeLLMModel, + start_layer=0, + end_layer=NUM_LAYERS, + use_hf_deepstack_order=False, + deepstack_embed_to_decoder_layer=range(3), + ) + + +def qwen3_5(): + from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM + + return stub_model( + Qwen3_5ForCausalLM, + _start_layer=0, + _end_layer=NUM_LAYERS, + flashinfer_mnnvl_cutedsl_fusion=None, + ) + + +MODELS = (qwen3_vl_moe, qwen3_5) + + +class TestDeepstackOnDeferredReduction(CustomTestCase): + def setUp(self): + patcher = patch( + "sglang.srt.layers.communicator.deferred_post_experts_all_reduce", + all_reduce, + ) + patcher.start() + self.addCleanup(patcher.stop) + self.embeds = torch.zeros(TOKENS, HIDDEN) + + def run_model(self, build, deepstack): + model = build() + return model.forward( + input_ids=None, + positions=None, + forward_batch=None, + input_embeds=self.embeds.clone(), + input_deepstack_embeds=deepstack, + ) + + def test_each_deepstack_embedding_is_added_once(self): + values = (0.125, 0.25, 0.5) + deepstack = torch.cat([torch.full((TOKENS, HIDDEN), v) for v in values], 1) + for build in MODELS: + with self.subTest(model=build.__name__): + hidden_states = self.run_model(build, deepstack) + expected = torch.full((TOKENS, HIDDEN), NUM_LAYERS + sum(values)) + torch.testing.assert_close(hidden_states, expected) + + def test_without_deepstack(self): + for build in MODELS: + with self.subTest(model=build.__name__): + hidden_states = self.run_model(build, None) + torch.testing.assert_close( + hidden_states, torch.full((TOKENS, HIDDEN), float(NUM_LAYERS)) + ) + + +if __name__ == "__main__": + unittest.main()