Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions python/sglang/srt/models/bailing_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/bailing_moe_v3.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/glm4_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/glm4_moe_lite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/glm5_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/gpt_oss.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/laguna.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/models/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
)
Expand Down
2 changes: 2 additions & 0 deletions python/sglang/srt/models/qwen3_vl_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down Expand Up @@ -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]
)
Expand Down
149 changes: 149 additions & 0 deletions test/registered/unit/models/test_aux_capture_deferred_allreduce.py
Original file line number Diff line number Diff line change
@@ -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()
139 changes: 139 additions & 0 deletions test/registered/unit/models/test_deepstack_deferred_allreduce.py
Original file line number Diff line number Diff line change
@@ -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()
Loading