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
35 changes: 7 additions & 28 deletions python/sglang/srt/batch_overlap/two_batch_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,11 @@
)
from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.boundary_layout import Layout
from sglang.srt.layers.communicator import (
CommunicateContext,
CommunicateSummableTensorPairFn,
ScatterMode,
reduce_output,
tbo_split_moves,
)
from sglang.srt.layers.moe import (
get_deepep_mode,
Expand Down Expand Up @@ -935,7 +935,6 @@ def model_forward_maybe_tbo(
positions: torch.Tensor,
forward_batch: ForwardBatch,
hidden_states: torch.Tensor,
input_data_scatter_mode: ScatterMode,
residual: Optional[torch.Tensor],
zero_allocator: Optional[BumpAllocator] = None,
):
Expand All @@ -946,16 +945,14 @@ def model_forward_maybe_tbo(
residual=residual,
zero_allocator=zero_allocator,
)
layer_input_scatter_mode = layers[0].layer_scatter_modes.layer_input_mode
operations_strategy = OperationsStrategy.init_new_tbo(
layers, forward_batch.global_forward_mode
)
if enable_tbo:
return _model_forward_tbo(
inputs=inputs,
operations_strategy=operations_strategy,
input_data_scatter_mode=input_data_scatter_mode,
layer_input_scatter_mode=layer_input_scatter_mode,
layer_input_rows=layers[0].layer_communicator.input_rows,
)
else:
return _model_forward_non_tbo(inputs, operations_strategy)
Expand All @@ -964,14 +961,11 @@ def model_forward_maybe_tbo(
def _model_forward_tbo(
inputs,
operations_strategy: OperationsStrategy,
input_data_scatter_mode: ScatterMode,
layer_input_scatter_mode: ScatterMode,
layer_input_rows: Layout,
):
inputs["hidden_states"] = reduce_output(inputs["hidden_states"])
inputs_arr = _model_forward_tbo_split_inputs(
**inputs,
input_data_scatter_mode=input_data_scatter_mode,
layer_input_scatter_mode=layer_input_scatter_mode,
**inputs, layer_input_rows=layer_input_rows
)
original_hidden_states_len = inputs["hidden_states"].shape[0]
del inputs
Expand Down Expand Up @@ -1005,25 +999,10 @@ def _model_forward_tbo_split_inputs(
positions: torch.Tensor,
forward_batch: ForwardBatch,
zero_allocator: Optional[BumpAllocator],
input_data_scatter_mode: ScatterMode,
layer_input_scatter_mode: ScatterMode,
layer_input_rows: Layout,
) -> List[Dict]:
tbo_splitter_scatter_mode = ScatterMode.TP_ATTN_FULL
context = CommunicateContext.init_new()
# The splitter cuts the attention-TP-full layout; each microbatch then moves
# to the first layer's input layout.
to_splitter = CommunicateSummableTensorPairFn.get_fn(
hidden_states_input_mode=input_data_scatter_mode,
residual_input_mode=input_data_scatter_mode,
output_mode=tbo_splitter_scatter_mode,
context=context,
)
to_layer_input = CommunicateSummableTensorPairFn.get_fn(
hidden_states_input_mode=tbo_splitter_scatter_mode,
residual_input_mode=tbo_splitter_scatter_mode,
output_mode=layer_input_scatter_mode,
context=context,
)
to_splitter, to_layer_input = tbo_split_moves(layer_input_rows)

hidden_states, residual = to_splitter(
hidden_states=hidden_states,
Expand Down
11 changes: 9 additions & 2 deletions python/sglang/srt/layers/boundary_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -229,11 +229,14 @@ def decoder_layer_sides(
leaves_for_next_layer: bool,
leaves_for_reduce_scatter: bool,
leaves_for_reduce_scatterv: bool,
hands_on_attention_rows: bool = False,
) -> DecoderLayerSides:
"""An attention followed by an FFN, derived from the groups each computes
over. The FFN runs either on the TP group (a dense MLP, or a MoE not
dispatched per DP shard) or on this rank's local rows (a MoE dispatched per
DP shard, which completes its own combine, or a dense MLP on every rank)."""
DP shard, which completes its own combine, or a dense MLP on every rank).
A layer on local rows hands on the attention's rows when it is the last, or
with ``hands_on_attention_rows``."""
# Attention computes over the attention-TP ranks of one DP (and CP) shard.
attention = Layout.sharded_over(
TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=axis_sizes
Expand Down Expand Up @@ -277,5 +280,9 @@ def decoder_layer_sides(
)
),
ffn_residual_rows=local if ffn_on_local_rows else attention,
output_rows=local if ffn_on_local_rows and not is_last_layer else attention,
output_rows=(
local
if ffn_on_local_rows and not (is_last_layer or hands_on_attention_rows)
else attention
),
)
59 changes: 51 additions & 8 deletions python/sglang/srt/layers/communicator.py
Original file line number Diff line number Diff line change
Expand Up @@ -467,6 +467,8 @@ class LayerScatterModes:
# Whether the layer before this one has a sparse MLP; None when the modes
# were given directly, not planned from the layer sequence.
is_previous_layer_sparse: Optional[bool] = None
# Whether the layer after this one has a sparse MLP; None likewise.
is_next_layer_sparse: Optional[bool] = None

@classmethod
def init_new(cls, **kwargs):
Expand All @@ -481,6 +483,7 @@ def init_new(cls, **kwargs):
is_first_layer=context.layer_id == 0,
is_last_layer=context.layer_id == context.num_layers - 1,
is_previous_layer_sparse=context.is_previous_layer_sparse,
is_next_layer_sparse=context.is_next_layer_sparse,
)

@classmethod
Expand Down Expand Up @@ -854,6 +857,7 @@ def __init__(
)
# The steps the layer's ordinary batches run.
sides = self._declared_sides()
self._declared = sides
self._steps = (
self._steps_from_declarations(
sides,
Expand Down Expand Up @@ -903,6 +907,19 @@ def __init__(
else None
)

@property
def input_rows(self) -> Layout:
"""The rows the layer's input and residual arrive on in a batch that
runs its ordinary steps."""
if self._declared is not None:
return self._declared.input_rows
# Such a batch holds every token on each CP rank.
return scatter_mode_layouts(
attn_dp_size=self._context.attn_dp_size,
attn_cp_size=1,
attn_tp_size=self._context.attn_tp_size,
)[self.layer_scatter_modes.layer_input_mode]

def _input_can_be_scattered(self) -> bool:
"""Whether a batch may run this layer with input-scattered attention:
configured, on pure TP without an a2a backend or a dense MLP on every
Expand Down Expand Up @@ -938,6 +955,17 @@ def _declared_sides(
def on_local_rows(sparse: bool) -> bool:
return moe_on_local_rows if sparse else dense_on_local_rows

def gathers_for_tbo(sparse: bool, next_sparse: Optional[bool]) -> bool:
# Under two-batch overlap a dense layer on local rows hands the
# sparse layer after it, where the split happens, the attention's
# rows.
return (
dense_on_local_rows
and get_exec().overlap.enable_two_batch_overlap
and not sparse
and bool(next_sparse)
)

if not (
self._takes_declared_boundaries
and (parallel.attn_cp_size == 1 or _cp_on_declarations())
Expand All @@ -949,14 +977,6 @@ def on_local_rows(sparse: bool) -> bool:
and modes.is_layer_sparse
and not _gathers_over_attention_cp()
)
# Under two-batch overlap a dense layer before a sparse one gathers
# its output over attention TP for the split; those layers keep the
# scatter-mode steps.
and not (
dense_on_local_rows
and parallel.attn_tp_size > 1
and get_exec().overlap.enable_two_batch_overlap
)
and (modes.is_first_layer or modes.is_previous_layer_sparse is not None)
):
return None
Expand All @@ -974,8 +994,14 @@ def on_local_rows(sparse: bool) -> bool:
previous_on_local_rows=(
not modes.is_first_layer
and on_local_rows(modes.is_previous_layer_sparse)
and not gathers_for_tbo(
modes.is_previous_layer_sparse, modes.is_layer_sparse
)
),
is_last_layer=modes.is_last_layer,
hands_on_attention_rows=gathers_for_tbo(
modes.is_layer_sparse, modes.is_next_layer_sparse
),
attention_gathers_local_rows=_use_ag_after_qlora,
ffn_group=SumGroup.MOE_OUTPUT if modes.is_layer_sparse else SumGroup.TP,
leaves_for_next_layer=self.allow_deferred_ffn_reduction and may_leave,
Expand Down Expand Up @@ -2153,6 +2179,23 @@ def _token_axis_sizes(*, cp_active: bool = False) -> Dict[TokenAxis, int]:
}


def tbo_split_moves(layer_input_rows: Layout) -> Tuple[Callable, Callable]:
"""The moves around the two-batch-overlap split, which cuts the attention's
rows: from the rows the first overlapped layer takes to the attention's,
and back again for each half."""
attention = Layout.sharded_over(
TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=_token_axis_sizes()
)
pair = CommunicateSummableTensorPairFn
if layer_input_rows == attention:
return pair._trivial, pair._trivial
if layer_input_rows.sharded - attention.sharded == {TokenAxis.ATTN_TP_SCATTER}:
# Each rank's slice: write the residual in and gather over attention
# TP, then take the slice of each half.
return pair._gather, pair._scatter
raise NotImplementedError(f"{layer_input_rows=}")


def _cp_on_declarations() -> bool:
"""Whether attention CP is one the declarations cover: a prefill CP that
shards tokens, with DSA or MLA attention, or with the FFN input gathered
Expand Down
3 changes: 0 additions & 3 deletions python/sglang/srt/models/deepseek_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -3124,9 +3124,6 @@ def forward(
forward_batch=forward_batch,
hidden_states=hidden_states,
residual=residual,
input_data_scatter_mode=self.layers[
normal_end_layer - 1
].layer_scatter_modes.layer_output_mode,
zero_allocator=zero_allocator,
)

Expand Down
3 changes: 0 additions & 3 deletions python/sglang/srt/models/dots3_common/modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -1797,9 +1797,6 @@ def forward(
forward_batch=forward_batch,
hidden_states=hidden_states,
residual=residual,
input_data_scatter_mode=self.layers[
normal_end_layer - 1
].layer_scatter_modes.layer_output_mode,
zero_allocator=zero_allocator,
)

Expand Down
3 changes: 0 additions & 3 deletions python/sglang/srt/models/glm4_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -1124,9 +1124,6 @@ def forward(
forward_batch=forward_batch,
hidden_states=hidden_states,
residual=residual,
input_data_scatter_mode=self.layers[
normal_end_layer - 1
].layer_scatter_modes.layer_output_mode,
)

last_layer = self.layers[self.end_layer - 1]
Expand Down
3 changes: 0 additions & 3 deletions python/sglang/srt/models/glm4_moe_lite.py
Original file line number Diff line number Diff line change
Expand Up @@ -820,9 +820,6 @@ def forward(
forward_batch=forward_batch,
hidden_states=hidden_states,
residual=residual,
input_data_scatter_mode=self.layers[
normal_end_layer - 1
].layer_scatter_modes.layer_output_mode,
zero_allocator=zero_allocator,
)

Expand Down
3 changes: 0 additions & 3 deletions python/sglang/srt/models/glm5_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -1183,9 +1183,6 @@ def forward(
forward_batch=forward_batch,
hidden_states=hidden_states,
residual=residual,
input_data_scatter_mode=self.layers[
normal_end_layer - 1
].layer_scatter_modes.layer_output_mode,
zero_allocator=zero_allocator,
)

Expand Down
8 changes: 0 additions & 8 deletions python/sglang/srt/models/mimo_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,6 @@
from sglang.srt.layers.communicator import (
LayerCommunicator,
LayerScatterModes,
ScatterMode,
enable_moe_dense_fully_dp,
)
from sglang.srt.layers.dp_attention import (
Expand Down Expand Up @@ -1060,13 +1059,6 @@ def forward(
hidden_states, residual = model_forward_maybe_tbo(
layers=self.layers[tbo_start_layer:tbo_end_layer],
enable_tbo=True,
input_data_scatter_mode=(
ScatterMode.model_input_output()
if tbo_start_layer == self.start_layer
else self.layers[
tbo_start_layer - 1
].layer_scatter_modes.layer_output_mode
),
positions=positions,
forward_batch=forward_batch,
hidden_states=hidden_states,
Expand Down
2 changes: 0 additions & 2 deletions python/sglang/srt/models/minimax_m2.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,6 @@
from sglang.srt.layers.communicator import (
LayerCommunicator,
LayerScatterModes,
ScatterMode,
)
from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce,
Expand Down Expand Up @@ -1153,7 +1152,6 @@ def forward(
hidden_states, residual = model_forward_maybe_tbo(
layers=self.layers,
enable_tbo=True,
input_data_scatter_mode=ScatterMode.model_input_output(),
positions=positions,
forward_batch=forward_batch,
hidden_states=hidden_states,
Expand Down
2 changes: 0 additions & 2 deletions python/sglang/srt/models/minimax_m3.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,6 @@
from sglang.srt.layers.communicator import (
LayerCommunicator,
LayerScatterModes,
ScatterMode,
enable_moe_dense_fully_dp,
)
from sglang.srt.layers.dp_attention import (
Expand Down Expand Up @@ -1488,7 +1487,6 @@ def forward(
hidden_states, residual = model_forward_maybe_tbo(
layers=self.layers,
enable_tbo=True,
input_data_scatter_mode=ScatterMode.model_input_output(),
positions=positions,
forward_batch=forward_batch,
hidden_states=hidden_states,
Expand Down
2 changes: 0 additions & 2 deletions python/sglang/srt/models/qwen2_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,6 @@
from sglang.srt.layers.communicator import (
LayerCommunicator,
LayerScatterModes,
ScatterMode,
)
from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled,
Expand Down Expand Up @@ -1184,7 +1183,6 @@ def forward(
hidden_states, residual = model_forward_maybe_tbo(
layers=self.layers,
enable_tbo=True,
input_data_scatter_mode=ScatterMode.model_input_output(),
positions=positions,
forward_batch=forward_batch,
hidden_states=hidden_states,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,8 @@
import torch

from sglang.srt.batch_overlap import two_batch_overlap as tbo
from sglang.srt.layers.communicator import ScatterMode, UnreducedOutput
from sglang.srt.layers.boundary_layout import Layout
from sglang.srt.layers.communicator import UnreducedOutput
from sglang.srt.utils import empty_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
Expand Down Expand Up @@ -58,8 +59,7 @@ def merge(output_a, output_b, original_len):
operations_strategy=SimpleNamespace(
deep_gemm_num_sms=None, operations=[], tbo_delta_stages=0
),
input_data_scatter_mode=ScatterMode.TP_ATTN_FULL,
layer_input_scatter_mode=ScatterMode.TP_ATTN_FULL,
layer_input_rows=Layout(frozenset()),
)

self.assertIs(seen["split"], reduced)
Expand Down
7 changes: 7 additions & 0 deletions test/registered/unit/layers/test_declared_attention_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,12 @@ def as_rank(self, cp, collectives, moe_group=None):
(comm, "post_experts_reduction_group"),
lambda: moe_group or parallel.tp_group,
),
(
(comm, "get_exec"),
lambda: SimpleNamespace(
overlap=SimpleNamespace(enable_two_batch_overlap=False)
),
),
(
(comm, "get_spec"),
lambda: SimpleNamespace(speculative_algorithm=None),
Expand Down Expand Up @@ -159,6 +165,7 @@ def build(self, allow_reduce_scatter):
is_last_layer=False,
is_layer_sparse=True,
is_previous_layer_sparse=False,
is_next_layer_sparse=False,
),
input_layernorm=layernorm,
post_attention_layernorm=layernorm,
Expand Down
Loading
Loading