From f3667cc813496944d85ccf97da56179937c074b0 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Fri, 25 Sep 2026 17:58:32 -0700 Subject: [PATCH 1/3] Choose a dense layer's steps from its declarations when every rank runs the MLP With moe_dense_tp_size 1 each rank runs the dense MLP on its own slice of the tokens and owes no sum, as a MoE dispatched per DP shard does. Such a layer now takes the declarations with its FFN on the local rows, and the layer after it takes those rows as its input; the steps are the ones the scatter modes chose. Under two-batch overlap a dense layer before a sparse one gathers its output for the split, so with it these layers keep the scatter-mode steps. Input-scattered attention, which never runs with such an MLP, gets no steps of its own for these layers. --- python/sglang/srt/layers/boundary_layout.py | 2 +- python/sglang/srt/layers/communicator.py | 25 ++++++--- .../layers/test_declared_decoder_boundary.py | 51 ++++++++++++++++++- 3 files changed, 67 insertions(+), 11 deletions(-) diff --git a/python/sglang/srt/layers/boundary_layout.py b/python/sglang/srt/layers/boundary_layout.py index f1779824f8f4..9dc3c9565673 100644 --- a/python/sglang/srt/layers/boundary_layout.py +++ b/python/sglang/srt/layers/boundary_layout.py @@ -187,7 +187,7 @@ def decoder_layer_sides( """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).""" + DP shard, which completes its own combine, or a dense MLP on every rank).""" # 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 diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index ca052523f2da..585811f7dba9 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -814,9 +814,9 @@ def __init__( 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. The rest of what - ``AttnTpContext.init_context`` requires is only known once the model is - built.""" + configured, on pure TP without an a2a backend or a dense MLP on every + rank. The rest of what ``AttnTpContext.init_context`` requires is only + known once the model is built.""" parallel = get_parallel() return ( parallel.enable_attn_tp_input_scattered @@ -824,6 +824,7 @@ def _input_can_be_scattered(self) -> bool: and parallel.attn_dp_size == 1 and parallel.attn_cp_size == 1 and get_moe_a2a_backend().is_none() + and not enable_moe_dense_fully_dp() ) def _declared_sides( @@ -839,8 +840,13 @@ def _declared_sides( modes = self.layer_scatter_modes parallel = get_parallel() # A MoE dispatched per DP shard computes on this rank's local rows and - # hands its layer's output on there. + # hands its layer's output on there; so does a dense MLP on every rank. moe_on_local_rows = is_moe_input_scattered_across_dp_ranks() + dense_on_local_rows = enable_moe_dense_fully_dp() + + def on_local_rows(sparse: bool) -> bool: + return moe_on_local_rows if sparse else dense_on_local_rows + if not ( self._takes_declared_boundaries and (parallel.attn_cp_size == 1 or _cp_on_declarations()) @@ -850,7 +856,11 @@ def _declared_sides( and parallel.attn_dp_size > 1 and modes.is_layer_sparse ) - and not enable_moe_dense_fully_dp() + # Under two-batch overlap a dense layer before a sparse one gathers + # its output for the split; those layers keep the scatter-mode steps. + and not ( + dense_on_local_rows and get_exec().overlap.enable_two_batch_overlap + ) and (modes.is_first_layer or modes.is_previous_layer_sparse is not None) ): return None @@ -858,11 +868,10 @@ def _declared_sides( may_leave = parallel.attn_cp_size == 1 return decoder_layer_sides( axis_sizes=_token_axis_sizes(cp_active=cp_active), - ffn_on_local_rows=modes.is_layer_sparse and moe_on_local_rows, + ffn_on_local_rows=on_local_rows(modes.is_layer_sparse), previous_on_local_rows=( not modes.is_first_layer - and modes.is_previous_layer_sparse - and moe_on_local_rows + and on_local_rows(modes.is_previous_layer_sparse) ), is_last_layer=modes.is_last_layer, attention_gathers_local_rows=_use_ag_after_qlora, diff --git a/test/registered/unit/layers/test_declared_decoder_boundary.py b/test/registered/unit/layers/test_declared_decoder_boundary.py index 8e181d953b3c..8dfb91996bb1 100644 --- a/test/registered/unit/layers/test_declared_decoder_boundary.py +++ b/test/registered/unit/layers/test_declared_decoder_boundary.py @@ -401,6 +401,46 @@ def test_a_dense_first_layer_before_sparse_ones(self): ) self.assertTrue(layers[3]._steps.ffn_input.keywords["gathers_residual"]) + def test_a_dense_mlp_on_every_rank(self): + # moe_dense_tp_size 1: each rank runs the dense MLP on its own slice, as + # an a2a MoE does, and owes no sum. + parallel = parallel_of(attn_dp=2, attn_tp=2, moe_dense_tp_size=1) + no_overlap = SimpleNamespace( + overlap=SimpleNamespace(enable_two_batch_overlap=False) + ) + for a2a in (False, True): + with ( + self.subTest(a2a=a2a), + patch.object(comm, "get_exec", lambda: no_overlap), + ): + layers = [ + build( + planned_modes( + i, + 4, + sparse=sparse, + previous_sparse=previous, + parallel=parallel, + a2a=a2a, + ), + parallel, + a2a=a2a, + allow_reduce_scatter=True, + ) + for i, (sparse, previous) in enumerate( + ((False, False), (False, False), (True, False), (True, True)) + ) + ] + self.assertEqual([self.declared(layer) for layer in layers], [True] * 4) + for dense in layers[:2]: + self.assertIs(dense._steps.ffn_input.func, comm._mlp_input_scatter) + self.assertIsNone(dense._steps.ffn_output.group) + for after_dense in layers[1:3]: + self.assertIs( + after_dense._steps.attention_input, + comm.CommunicateSimpleFn._scattered_to_tp_attn_full, + ) + def test_the_last_a2a_layer_folds_the_residual_back(self): parallel = parallel_of(attn_dp=2, attn_tp=2) last = build( @@ -431,14 +471,21 @@ def test_layers_that_keep_the_scatter_modes(self): layer_facts(1, 3), parallel_of(attn_dp=2, attn_tp=1, attn_cp=2), ), - "dense MLP fully DP": ( + # The layer before a sparse one gathers its output for the split. + "dense MLP fully DP under two-batch overlap": ( layer_facts(1, 3), parallel_of(attn_dp=2, attn_tp=2, moe_dense_tp_size=1), ), } + two_batch_overlap = SimpleNamespace( + overlap=SimpleNamespace(enable_two_batch_overlap=True) + ) for name, (facts, parallel) in cases.items(): with self.subTest(name): - with planning(parallel): + with ( + planning(parallel), + patch.object(comm, "get_exec", lambda: two_batch_overlap), + ): communicator = LayerCommunicator.__new__(LayerCommunicator) communicator.layer_scatter_modes = facts communicator.allow_deferred_ffn_reduction = True From 187a323e40154b36bbb1731e5dd28a81fd86d9a3 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Fri, 25 Sep 2026 18:13:48 -0700 Subject: [PATCH 2/3] Choose DSA and MLA prefill CP layers' steps from their declarations DSA and MLA prefill CP ran a communicator subclass with its own tables. Under them attention TP is one and the dense MLP runs on every rank, so most layers only normalize. A MoE on the TP group (DSA's interleave split without an a2a backend) gathers a CP extend's shards over the attention-CP group and reduce-scatters its output there, which also completes the sum the MoE leaves. These layers now take the declarations, as GQA prefill CP does: a CP extend runs the layer's CP steps, other batches its ordinary ones. Under DSA and MLA CP the CP steps gather over the attention-CP group in equal shards, and the FFN output declares that it leaves its sum to the reduce-scatter that takes each rank's shard back. GQA prefill CP keeps its padded MoE-CP gather and completes the sum within the layer. Under DSA and MLA CP a MoE under attention DP is declared too (it dispatches with a2a), and so is a dense layer under two-batch overlap: with attention TP one its gather for the split moves nothing. The subclass is removed; its collectives and the KV prefetch stay in communicator_dsa_cp. A layer under DSA or MLA prefill CP that the declarations do not cover fails at construction, since the scatter-mode steps have no attention-CP gather. On a CP extend, a layer whose FFN runs on the rank's own rows no longer publishes mlp_reduce_scatter. It owes no sum, so no reader of the flag behaves differently. --- .../arg_groups/model_overrides/deepseek_v2.py | 8 +- python/sglang/srt/layers/communicator.py | 129 +++++++++-- .../sglang/srt/layers/communicator_dsa_cp.py | 153 -------------- .../sglang/srt/model_executor/model_runner.py | 2 +- .../runner/decode_cuda_graph_runner.py | 11 +- python/sglang/srt/models/deepseek_v2.py | 9 +- .../layers/test_declared_decoder_boundary.py | 200 +++++++++++++++--- 7 files changed, 298 insertions(+), 214 deletions(-) diff --git a/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py b/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py index 869363332117..54e2cd2a2e5d 100644 --- a/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py +++ b/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py @@ -102,8 +102,8 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues." ) # Note(kpham-sgl): Keep attn_tp_size == 1 under DSA CP. - # DSACPLayerCommunicator does not all-reduce attention-TP - # partial o_proj outputs before replicated dense FFNs. + # The DSA / MLA CP gather and reduce-scatter + # (communicator_dsa_cp) assume it. attn_cp_size = cfg.tp_size // cfg.dp_size overrides["attn_cp_size"] = attn_cp_size logger.warning( @@ -167,8 +167,8 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict: "For MLA CP, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, batch_size == 1" ) # FIXME(kpham-sgl): Keep attn_tp_size == 1 under MLA CP. - # DSACPLayerCommunicator does not all-reduce attention-TP - # partial o_proj outputs before replicated dense FFNs. + # The DSA / MLA CP gather and reduce-scatter + # (communicator_dsa_cp) assume it. attn_cp_size = cfg.tp_size // cfg.dp_size overrides["attn_cp_size"] = attn_cp_size logger.warning( diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 585811f7dba9..d0af80b26443 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -48,6 +48,10 @@ input_scattered_layer_sides, sequence_parallel_layer_sides, ) +from sglang.srt.layers.communicator_dsa_cp import ( + dsa_cp_gather_hidden_states, + dsa_cp_reduce_scatter_hidden_states, +) from sglang.srt.layers.cp.utils import ( is_mla_cp_active, is_mla_cp_enabled, @@ -494,7 +498,7 @@ def _compute_mlp_mode(cls, context: _LayerModeComputationContext): # A TP-sharded dense MLP reduces over the whole TP group, which spans # every CP rank; a CP-sharded prefill must gather tokens across CP # first or the all-reduce sums different tokens' partial outputs. - # MLA/DSA CP models do this in DSACPLayerCommunicator instead. + # MLA/DSA CP models gather over the attention-CP group instead. if _generic_prefill_cp_shards_tokens() and not ( is_dsa_enable_prefill_cp() or is_mla_cp_enabled() ): @@ -758,6 +762,15 @@ def __init__( ) # The steps the layer's ordinary batches run. sides = self._declared_sides() + if ( + sides is None + and self._takes_declared_boundaries + and (_generic_prefill_cp_shards_tokens() and _gathers_over_attention_cp()) + ): + # The scatter-mode steps have no attention-CP gather. + raise NotImplementedError( + "a DSA or MLA prefill CP layer outside the declarations" + ) self._steps = ( _select_boundary_steps( sides, @@ -850,22 +863,35 @@ def on_local_rows(sparse: bool) -> bool: if not ( self._takes_declared_boundaries and (parallel.attn_cp_size == 1 or _cp_on_declarations()) - # MoE layers under attention DP and CP keep the scatter-mode steps. + # MoE layers under attention DP and GQA prefill CP keep the + # scatter-mode steps. and not ( parallel.attn_cp_size > 1 and parallel.attn_dp_size > 1 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 for the split; those layers keep the scatter-mode steps. + # its output over attention TP for the split; those layers keep the + # scatter-mode steps. and not ( - dense_on_local_rows and get_exec().overlap.enable_two_batch_overlap + 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 - # Under attention CP the FFN completes its own sum. - may_leave = parallel.attn_cp_size == 1 + if _gathers_over_attention_cp(): + # DSA and MLA CP: a CP extend's FFN leaves its sum to the + # reduce-scatter that takes each rank's shard back. + may_leave = not cp_active + may_leave_to_reduce_scatter = True + else: + # Under GQA prefill CP the FFN completes its own sum: the next + # layer's input holds only this rank's chunk, and a reduce-scatter + # back over attention DP would split across the CP ranks. + may_leave = may_leave_to_reduce_scatter = parallel.attn_cp_size == 1 return decoder_layer_sides( axis_sizes=_token_axis_sizes(cp_active=cp_active), ffn_on_local_rows=on_local_rows(modes.is_layer_sparse), @@ -877,7 +903,8 @@ def on_local_rows(sparse: bool) -> bool: 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, - leaves_for_reduce_scatter=self.allow_reduce_scatter and may_leave, + leaves_for_reduce_scatter=self.allow_reduce_scatter + and may_leave_to_reduce_scatter, # A MoE block leaves its sum to reduce_scatterv whenever that combine # applies (should_skip_post_experts_all_reduce). leaves_for_reduce_scatterv=( @@ -1090,7 +1117,7 @@ def _batch_steps(self, forward_batch: ForwardBatch) -> "BoundarySteps": and get_attn_tp_context().input_scattered ): return self._input_scattered_steps - if self._cp_steps is not None and moe_cp_gathered_rows(forward_batch): + if self._cp_steps is not None and _batch_shards_over_cp(forward_batch): return self._cp_steps return self._steps @@ -1301,13 +1328,18 @@ def _ffn_leaves_sum_to_reduce_scatter( """Whether the FFN leaves its sum out because a reduce-scatter completes it: the attention-DP one ``dp_step`` names, or the CP / input-scattered one.""" - if not self._batch_steps(forward_batch).ffn_output.leaves_for_reduce_scatter: + steps = self._batch_steps(forward_batch) + if not steps.ffn_output.leaves_for_reduce_scatter: return False - if dp_step is not None: + if dp_step is not None or steps is self._cp_steps: return True - # Prefill CP predicates must stay out of decode graph capture. - if forward_batch.forward_mode.is_context_parallel_extend() and ( - dsa_use_prefill_cp(forward_batch) or is_mla_cp_active(forward_batch) + # The scatter-mode steps of a DSA or MLA CP extend (the subclasses that + # pick their own steps). Prefill CP predicates must stay out of decode + # graph capture. + if ( + self._cp_steps is None + and forward_batch.forward_mode.is_context_parallel_extend() + and (dsa_use_prefill_cp(forward_batch) or is_mla_cp_active(forward_batch)) ): return True return get_attn_tp_context().input_scattered and not self.is_last_layer @@ -1958,15 +1990,31 @@ def _token_axis_sizes(*, cp_active: bool = False) -> Dict[TokenAxis, int]: def _cp_on_declarations() -> bool: """Whether attention CP is one the declarations cover: a prefill CP that - shards tokens, whose FFN input gathers over a MoE-CP group that is the - whole CP group. DSA and MLA CP pick their own steps.""" - return ( - _generic_prefill_cp_shards_tokens() - and not (is_dsa_enable_prefill_cp() or is_mla_cp_enabled()) - and get_parallel().moe_dp_size == 1 + shards tokens, with DSA or MLA attention, or with the FFN input gathered + over a MoE-CP group that is the whole CP group.""" + return _generic_prefill_cp_shards_tokens() and ( + _gathers_over_attention_cp() or get_parallel().moe_dp_size == 1 ) +def _gathers_over_attention_cp() -> bool: + """Whether a CP extend gathers the FFN input over the attention-CP group in + equal shards and takes the output back with a reduce-scatter there: DSA and + MLA CP. GQA prefill CP gathers over the MoE-CP group instead.""" + return is_dsa_enable_prefill_cp() or is_mla_cp_enabled() + + +def _batch_shards_over_cp(forward_batch: ForwardBatch) -> bool: + """Whether this batch's tokens are split across the CP ranks. Only a context + parallel extend is, so other batches, decode graph capture among them, + never read the CP predicates.""" + if not forward_batch.forward_mode.is_context_parallel_extend(): + return False + if _gathers_over_attention_cp(): + return dsa_use_prefill_cp(forward_batch) or is_mla_cp_active(forward_batch) + return moe_cp_gathered_rows(forward_batch) is not None + + class BoundarySteps(msgspec.Struct, frozen=True): """The steps a batch runs at a layer's boundaries: into the attention, from the attention output to the FFN input, and the FFN output on to the @@ -2105,8 +2153,9 @@ def _select_ffn_input( (), ) if gathered == {TokenAxis.ATTN_CP}: - # Each CP rank completes its own chunk, then the chunks are gathered - # over the MoE-CP group, each padded to the longest. + # Each CP rank completes its own chunk, then the chunks are gathered: + # under DSA and MLA CP over the attention-CP group in equal shards, + # otherwise over the MoE-CP group, each padded to the longest. on_chunk, fused = _select_ffn_input( produced, residual=residual, @@ -2116,6 +2165,8 @@ def _select_ffn_input( fusions=fusions, residual_joins_sum=residual_joins_sum, ) + if _gathers_over_attention_cp(): + return partial(_mlp_input_gather_attention_cp, gather=on_chunk), fused return partial(_mlp_input_gather_moe_cp, gather=on_chunk), fused if ( residual_to != produced.layout @@ -2192,6 +2243,12 @@ def _select_ffn_output_move( if to != residual or not produced.layout.sharded <= residual.sharded: raise NotImplementedError(f"{produced=} {residual=} {to=}") if returned == {TokenAxis.ATTN_CP}: + if _gathers_over_attention_cp(): + # The reduce-scatter over attention CP completes the sum the FFN + # left and takes this rank's shard back. + if not produced.leaves_for_reduce_scatter: + raise NotImplementedError(f"{produced=} {residual=} {to=}") + return False, CommunicateSummableTensorPairFn._reduce_scatter_over_cp # This rank's chunk of the rows gathered over CP; no collective. return False, CommunicateSummableTensorPairFn._scatter_hidden_states_moe if returned == {TokenAxis.ATTN_DP, TokenAxis.ATTN_CP}: @@ -2354,6 +2411,24 @@ def _mlp_input_gather( return order(hidden_states, residual, forward_batch, layernorm, context) +def _mlp_input_gather_attention_cp( + hidden_states: torch.Tensor, + residual: torch.Tensor, + forward_batch: ForwardBatch, + layernorm: torch.nn.Module, + context: CommunicateContext, + *, + gather: Callable, +): + """DSA and MLA CP: complete this rank's shard, then gather the shards, of + equal length, over the attention-CP group. The residual stays on the + shard.""" + hidden_states, residual = gather( + hidden_states, residual, forward_batch, layernorm, context + ) + return dsa_cp_gather_hidden_states(hidden_states), residual + + def _mlp_input_gather_moe_cp( hidden_states: torch.Tensor, residual: torch.Tensor, @@ -2570,6 +2645,18 @@ def _scatter( assert residual is None, "not yet handled residual!=None" return _redistribute_to_attn_tp_shards(hidden_states, context), None + @staticmethod + def _reduce_scatter_over_cp( + hidden_states: torch.Tensor, + residual: torch.Tensor, + forward_batch: ForwardBatch, + context: CommunicateContext, + **kwargs, + ): + """DSA and MLA CP: sum the FFN output over the attention-CP group and + keep this rank's shard.""" + return dsa_cp_reduce_scatter_hidden_states(hidden_states), residual + @staticmethod def _take_back_cp_shard( hidden_states: torch.Tensor, diff --git a/python/sglang/srt/layers/communicator_dsa_cp.py b/python/sglang/srt/layers/communicator_dsa_cp.py index 7a0dfd0600c3..0b74f4290b52 100644 --- a/python/sglang/srt/layers/communicator_dsa_cp.py +++ b/python/sglang/srt/layers/communicator_dsa_cp.py @@ -13,7 +13,6 @@ # ============================================================================== -from functools import partial from typing import Optional import torch @@ -21,15 +20,6 @@ from sglang.srt.layers.attention.dsa.utils import ( dsa_use_prefill_cp, ) -from sglang.srt.layers.communicator import ( - CommunicateContext, - CommunicateSimpleFn, - CommunicateSummableTensorPairFn, - LayerCommunicator, - ScatterMode, - _mlp_input_norm, -) -from sglang.srt.layers.cp.utils import is_mla_cp_active from sglang.srt.layers.dp_attention import ( attn_cp_all_gather_into_tensor, attn_cp_reduce_scatter_tensor, @@ -81,146 +71,3 @@ def dsa_cp_reduce_scatter_hidden_states(hidden_states: torch.Tensor): hidden_states = hidden_states.tensor_split(cp_size)[cp_rank] attn_cp_reduce_scatter_tensor(hidden_states, input_hidden_states) return hidden_states - - -class DSACPLayerCommunicator(LayerCommunicator): - # Chooses its own boundary steps, not from the declarations. - _takes_declared_boundaries = False - - def _post_init_communicate(self): - # SCATTERED in attn tp is different from SCATTERED in global tp when dp_size > 1 - if self.layer_scatter_modes.mlp_mode != ScatterMode.SCATTERED: - assert self._context.attn_dp_size == 1, ( - f"dp_size should be 1 when moe_runner_backend is none" - ) - return ( - DSACPCommunicateSimpleFn.get_fn( - input_mode=ScatterMode.SCATTERED, - output_mode=ScatterMode.SCATTERED, - context=self._context, - ), - DSACPCommunicateSummableTensorPairFn.get_fn( - hidden_states_input_mode=self.layer_scatter_modes.mlp_mode, # SCATTERED, FULL - residual_input_mode=ScatterMode.SCATTERED, - output_mode=ScatterMode.SCATTERED, - context=self._context, - ), - ) - - def _select_mlp_input(self): - fn = DSACPCommunicateWithAllReduceAndLayerNormFn.get_fn( - hidden_states_input_mode=ScatterMode.SCATTERED, - residual_input_mode=ScatterMode.SCATTERED, - hidden_states_output_mode=self.layer_scatter_modes.mlp_mode, # SCATTERED, FULL - residual_output_mode=ScatterMode.SCATTERED, - context=self._context, - ) - return fn, () - - -class DSACPCommunicateSimpleFn(CommunicateSimpleFn): - @staticmethod - def get_fn( - input_mode: ScatterMode, - output_mode: ScatterMode, - context: CommunicateContext, - ): - if context.is_same_group_size(input_mode, output_mode): - return DSACPCommunicateSimpleFn._trivial - - raise NotImplementedError(f"{input_mode=} {output_mode=}") - - -class DSACPCommunicateWithAllReduceAndLayerNormFn: - """Besides communication, needs to - 1. All reduce in tp_attn_group on hidden_states - 2. Apply layer norm - """ - - @staticmethod - def get_fn( - hidden_states_input_mode: ScatterMode, - residual_input_mode: ScatterMode, - hidden_states_output_mode: ScatterMode, - residual_output_mode: ScatterMode, - context: CommunicateContext, - ): - assert hidden_states_input_mode == ScatterMode.SCATTERED - assert residual_input_mode == ScatterMode.SCATTERED - assert residual_output_mode == ScatterMode.SCATTERED - if hidden_states_output_mode == ScatterMode.SCATTERED: - return _mlp_input_norm - - if hidden_states_output_mode == ScatterMode.FULL: - return partial( - DSACPCommunicateWithAllReduceAndLayerNormFn._gather_hidden_states_and_residual, - residual_input_mode=residual_input_mode, - ) - - raise NotImplementedError( - f"{hidden_states_input_mode=} {residual_input_mode=} {hidden_states_output_mode=} {residual_output_mode=}" - ) - - @staticmethod - def _gather_hidden_states_and_residual( - hidden_states: torch.Tensor, - residual: torch.Tensor, - forward_batch: ForwardBatch, - layernorm: torch.nn.Module, - context: CommunicateContext, - *, - residual_input_mode, - ): - if hidden_states.shape[0] != 0: - hidden_states, residual = layernorm(hidden_states, residual) - # for prefill: attn tp scattered -> full - # for decode: attn tp full -> full - if dsa_use_prefill_cp(forward_batch) or is_mla_cp_active(forward_batch): - hidden_states = dsa_cp_gather_hidden_states(hidden_states) - return hidden_states, residual - - -class DSACPCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn): - """It is allowed to make (hidden_states, residual) := (hidden_states + residual, None) if needed.""" - - @staticmethod - def get_fn( - hidden_states_input_mode: ScatterMode, - residual_input_mode: ScatterMode, - output_mode: ScatterMode, - context: CommunicateContext, - ): - # Check exact enum match first: even if group sizes happen to be equal - # (e.g. tp_size == attn_cp_size makes FULL and SCATTERED both size 1), - # FULL and SCATTERED have different data layouts under CP and require - # an explicit scatter operation. - if ( - (hidden_states_input_mode == ScatterMode.FULL) - and (residual_input_mode == ScatterMode.SCATTERED) - and (output_mode == ScatterMode.SCATTERED) - ): - return DSACPCommunicateSummableTensorPairFn._scatter_hidden_states - - if context.is_same_group_size( - hidden_states_input_mode, output_mode - ) and context.is_same_group_size(residual_input_mode, output_mode): - return DSACPCommunicateSummableTensorPairFn._trivial - - raise NotImplementedError( - f"{hidden_states_input_mode=} {residual_input_mode=} {output_mode=}" - ) - - @staticmethod - def _scatter_hidden_states( - hidden_states: torch.Tensor, - residual: torch.Tensor, - forward_batch: ForwardBatch, - context: CommunicateContext, - allow_reduce_scatter: bool = False, - **kwargs, - ): - # for prefill: full -> attn tp scattered - # for decode: full -> attn tp full - if dsa_use_prefill_cp(forward_batch) or is_mla_cp_active(forward_batch): - hidden_states = dsa_cp_reduce_scatter_hidden_states(hidden_states) - return hidden_states, residual diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 8171b691cce3..604a72f4ed7f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1630,7 +1630,7 @@ def _prepare_eager_forward_batch(self, forward_batch: ForwardBatch) -> None: self.lora_manager.prepare_lora_batch(forward_batch) # Derive the LOCAL num_token_non_padded from the GLOBAL scalar. sharded is - # cleared for DSACPLayerCommunicator-style CP (DSA, MLA): those flavors + # cleared for DSA and MLA prefill CP: those flavors # already feed a zigzag-split rank-local layout whose token count should # not be further divided by attn_tp_size, so they keep the full count. # MHA-arch prefill CP (Qwen3/Qwen2 MoE) keeps the attn_tp-replicated diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 307c197e82bc..a7b73754114c 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -268,12 +268,11 @@ def __init__( self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_rank = get_parallel().attn_tp_rank - # True if a DSACPLayerCommunicator-style prefill-CP flavor is active - # (DSA or MLA). These flavors feed a zigzag-split rank-local layout - # into the runner; MHA-arch prefill CP (Qwen3/Qwen2 MoE via PR - # #18233) uses the plain LayerCommunicator with an attn_tp-replicated - # layout and is intentionally excluded so the attn_tp-local - # num_token_non_padded adjustment still runs for it. + # True if the DSA or MLA prefill-CP flavor is active. These flavors + # feed a zigzag-split rank-local layout into the runner; MHA-arch + # prefill CP (Qwen3/Qwen2 MoE via PR #18233) keeps an + # attn_tp-replicated layout and is intentionally excluded so the + # attn_tp-local num_token_non_padded adjustment still runs for it. self.enable_prefill_cp = is_dsa_enable_prefill_cp() or is_mla_cp_enabled() self.deepep_adapter = DeepEPCudaGraphRunnerAdapter() diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 9542a8bf5005..3d935baceae5 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -75,7 +75,6 @@ layer_input_buffer, ) from sglang.srt.layers.communicator_dsa_cp import ( - DSACPLayerCommunicator, maybe_prefetch_next_full_attention_kv, ) from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx @@ -2632,9 +2631,11 @@ def _build_layer_communicator( ): """The communicator for this layer's norms; it chooses its boundary steps from them at construction.""" - if get_parallel().enable_prefill_cp: - communicator_cls = DSACPLayerCommunicator - elif not self.is_nextn and _use_mnnvl_cutedsl_fusion(): + if ( + not get_parallel().enable_prefill_cp + and not self.is_nextn + and _use_mnnvl_cutedsl_fusion() + ): # Dense layers too: selecting cutedsl turns the legacy fusion off. from sglang.srt.layers.moe.cutedsl_ar_fusion import ( CuteDSLFusionLayerCommunicator, diff --git a/test/registered/unit/layers/test_declared_decoder_boundary.py b/test/registered/unit/layers/test_declared_decoder_boundary.py index 8dfb91996bb1..017cfe9645b9 100644 --- a/test/registered/unit/layers/test_declared_decoder_boundary.py +++ b/test/registered/unit/layers/test_declared_decoder_boundary.py @@ -34,7 +34,6 @@ UnreducedOutput, scatter_mode_layouts, ) -from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator from sglang.srt.layers.communicator_mhc import MHCLayerCommunicator from sglang.srt.layers.moe.cutedsl_ar_fusion import CuteDSLFusionLayerCommunicator from sglang.test.ci.ci_register import register_cpu_ci @@ -70,10 +69,11 @@ def parallel_of(*, attn_dp, attn_tp, attn_cp=1, **overrides): @contextmanager -def planning(parallel, *, sp=False, a2a=False): +def planning(parallel, *, sp=False, a2a=False, dsa_cp=False): """What layer planning and communicator construction read, without the process-wide parallel state. ``parallel`` may be a callable, for a - per-thread parallel state.""" + per-thread parallel state. ``dsa_cp``: the prefill CP is DSA's (MLA's is + the same to the communicator).""" get_parallel = parallel if callable(parallel) else (lambda: parallel) def moe_cp_gathers(): @@ -81,7 +81,7 @@ def moe_cp_gathers(): with ( patch.object(comm, "get_parallel", get_parallel), - patch.object(comm, "is_dsa_enable_prefill_cp", lambda: False), + patch.object(comm, "is_dsa_enable_prefill_cp", lambda: dsa_cp), patch.object(comm, "is_mla_cp_enabled", lambda: False), patch.object( comm, @@ -103,6 +103,14 @@ def moe_cp_gathers(): ), patch.object(comm, "is_enable_moe_cp_allgather", moe_cp_gathers), patch.object(comm, "get_lora", lambda: SimpleNamespace(enable_lora=False)), + # Planning asks whether a dense layer gathers for two-batch overlap. + patch.object( + comm, + "get_exec", + lambda: SimpleNamespace( + overlap=SimpleNamespace(enable_two_batch_overlap=False) + ), + ), ): yield @@ -168,8 +176,17 @@ def sides_of( ) -def build(modes, parallel, *, cls=LayerCommunicator, sp=False, a2a=False, **kwargs): - with planning(parallel, sp=sp, a2a=a2a): +def build( + modes, + parallel, + *, + cls=LayerCommunicator, + sp=False, + a2a=False, + dsa_cp=False, + **kwargs, +): + with planning(parallel, sp=sp, a2a=a2a, dsa_cp=dsa_cp): return cls( layer_scatter_modes=modes, input_layernorm=Norm(), @@ -179,9 +196,16 @@ def build(modes, parallel, *, cls=LayerCommunicator, sp=False, a2a=False, **kwar def planned_modes( - layer_id, num_layers, *, sparse, previous_sparse, parallel, a2a=False + layer_id, + num_layers, + *, + sparse, + previous_sparse, + parallel, + a2a=False, + dsa_cp=False, ): - with planning(parallel, a2a=a2a): + with planning(parallel, a2a=a2a, dsa_cp=dsa_cp): return LayerScatterModes.init_new( layer_id=layer_id, num_layers=num_layers, @@ -503,11 +527,7 @@ def test_layers_that_keep_the_scatter_modes(self): self.assertFalse(self.declared(build(direct, dp))) def test_subclasses_that_pick_their_own_steps(self): - for cls in ( - MHCLayerCommunicator, - DSACPLayerCommunicator, - CuteDSLFusionLayerCommunicator, - ): + for cls in (MHCLayerCommunicator, CuteDSLFusionLayerCommunicator): with self.subTest(cls.__name__): self.assertFalse(cls._takes_declared_boundaries) self.assertTrue(LayerCommunicator._takes_declared_boundaries) @@ -995,6 +1015,99 @@ def cp_parallel(self, **overrides): attn_dp=1, attn_tp=2, attn_cp=2, enable_prefill_cp=True, **overrides ) + def dsa_parallel(self, **overrides): + # DSA and MLA CP run attention TP 1 and the dense MLP on every rank. + return parallel_of( + attn_dp=1, + attn_tp=1, + attn_cp=2, + enable_prefill_cp=True, + moe_dense_tp_size=1, + **overrides, + ) + + def test_a_dsa_cp_extend_leaves_its_sum_to_the_reduce_scatter(self): + # A MoE on the TP group, as under DSA interleave CP without a2a. + parallel = self.dsa_parallel() + modes = planned_modes( + 1, 3, sparse=True, previous_sparse=False, parallel=parallel, dsa_cp=True + ) + communicator = build(modes, parallel, dsa_cp=True, allow_reduce_scatter=True) + cp, ordinary = communicator._cp_steps, communicator._steps + self.assertIs(cp.ffn_input.func, comm._mlp_input_gather_attention_cp) + self.assertIs(cp.ffn_input.keywords["gather"], comm._mlp_input_norm) + self.assertIs( + cp.ffn_output_move, + comm.CommunicateSummableTensorPairFn._reduce_scatter_over_cp, + ) + self.assertTrue(cp.ffn_output.leaves_for_reduce_scatter) + self.assertFalse(cp.ffn_output.leaves_for_next_layer) + # The dense layer before it ran on the same shard: nothing to move. + self.assertIs(cp.attention_input, comm.CommunicateSimpleFn._trivial) + # Other batches hold every token on each CP rank: the MoE sums itself. + self.assertIs(ordinary.ffn_input, comm._mlp_input_norm) + self.assertIs( + ordinary.ffn_output_move, comm.CommunicateSummableTensorPairFn._trivial + ) + with ( + patch.object(comm, "get_forward", lambda: SimpleNamespace(sp_active=False)), + patch.object( + comm, + "get_attn_tp_context", + lambda: SimpleNamespace(input_scattered=False), + ), + ): + for shards, leaves in ((True, True), (False, False)): + with ( + self.subTest(cp_extend=shards), + patch.object(comm, "_batch_shards_over_cp", lambda fb: shards), + ): + self.assertIs( + communicator._ffn_leaves_sum_to_reduce_scatter( + SimpleNamespace( + forward_mode=SimpleNamespace( + is_context_parallel_extend=lambda: False + ) + ), + None, + ), + leaves, + ) + + def test_dsa_cp_asks_its_own_predicates_for_a_cp_extend(self): + def batch(cp_extend): + return SimpleNamespace( + forward_mode=SimpleNamespace( + is_context_parallel_extend=lambda: cp_extend + ) + ) + + with planning(self.dsa_parallel(), dsa_cp=True): + for cp_extend, active, shards in ( + (True, True, True), + (True, False, False), + (False, True, False), + ): + with ( + self.subTest(cp_extend=cp_extend, active=active), + patch.object(comm, "dsa_use_prefill_cp", lambda fb: active), + patch.object(comm, "is_mla_cp_active", lambda fb: False), + ): + self.assertIs(comm._batch_shards_over_cp(batch(cp_extend)), shards) + + def test_dsa_cp_dense_layers_run_on_their_shard(self): + parallel = self.dsa_parallel() + modes = planned_modes( + 1, 3, sparse=False, previous_sparse=False, parallel=parallel, dsa_cp=True + ) + communicator = build(modes, parallel, dsa_cp=True, allow_reduce_scatter=True) + for steps in (communicator._steps, communicator._cp_steps): + self.assertIs(steps.ffn_input, comm._mlp_input_norm) + self.assertIsNone(steps.ffn_output.group) + self.assertIs( + steps.ffn_output_move, comm.CommunicateSummableTensorPairFn._trivial + ) + def test_a_cp_extend_gathers_over_cp_and_takes_its_chunk_back(self): communicator = build(layer_facts(1, 3), self.cp_parallel()) cp = communicator._cp_steps @@ -1029,16 +1142,30 @@ def test_the_fused_kernels_run_on_each_chunk(self): def test_which_cp_the_declarations_cover(self): dp_cp = parallel_of(attn_dp=2, attn_tp=1, attn_cp=2, enable_prefill_cp=True) - for name, parallel, sparse, declared in ( - ("prefill CP", self.cp_parallel(), False, True), + dsa_dp_cp = parallel_of( + attn_dp=2, attn_tp=1, attn_cp=2, enable_prefill_cp=True, moe_dense_tp_size=1 + ) + for name, parallel, sparse, declared, dsa_cp, a2a in ( + ("DSA or MLA CP", self.dsa_parallel(), True, True, True, False), + ( + "DSA or MLA CP, a MoE under attention DP", + dsa_dp_cp, + True, + True, + True, + True, + ), + ("prefill CP", self.cp_parallel(), False, True, False, False), ( "CP without prefill CP", parallel_of(attn_dp=1, attn_tp=2, attn_cp=2), False, False, + False, + False, ), - ("CP under attention DP", dp_cp, False, True), - ("a MoE under attention DP and CP", dp_cp, True, False), + ("CP under attention DP", dp_cp, False, True, False, False), + ("a MoE under attention DP and GQA CP", dp_cp, True, False, False, False), ( "a MoE-CP group narrower than CP", parallel_of( @@ -1050,13 +1177,23 @@ def test_which_cp_the_declarations_cover(self): ), False, False, + False, + False, ), ): with self.subTest(name): modes = planned_modes( - 1, 3, sparse=sparse, previous_sparse=sparse, parallel=parallel + 1, + 3, + sparse=sparse, + previous_sparse=sparse, + parallel=parallel, + a2a=a2a, + dsa_cp=dsa_cp, + ) + communicator = build( + modes, parallel, a2a=a2a, dsa_cp=dsa_cp, allow_reduce_scatter=True ) - communicator = build(modes, parallel) self.assertIs(communicator._cp_steps is not None, declared) def test_under_attention_dp_one_dp_sum_gathers_both_axes(self): @@ -1078,13 +1215,26 @@ def test_under_attention_dp_one_dp_sum_gathers_both_axes(self): def test_a_batch_runs_them_only_on_a_cp_extend(self): communicator = build(layer_facts(1, 3), self.cp_parallel()) - for rows, expected in ( - (None, communicator._steps), - ([2, 1], communicator._cp_steps), + + def batch(cp_extend): + return SimpleNamespace( + forward_mode=SimpleNamespace( + is_context_parallel_extend=lambda: cp_extend + ) + ) + + def unread(fb): + raise AssertionError("read for a batch that is not a CP extend") + + for fb, rows, expected in ( + (batch(False), unread, communicator._steps), + (batch(True), lambda fb: None, communicator._steps), + (batch(True), lambda fb: [2, 1], communicator._cp_steps), ): with ( - self.subTest(gathered_rows=rows), - patch.object(comm, "moe_cp_gathered_rows", lambda fb: rows), + self.subTest(expected=expected is communicator._cp_steps), + planning(self.cp_parallel()), + patch.object(comm, "moe_cp_gathered_rows", rows), patch.object( comm, "get_forward", lambda: SimpleNamespace(sp_active=False) ), @@ -1094,7 +1244,7 @@ def test_a_batch_runs_them_only_on_a_cp_extend(self): lambda: SimpleNamespace(input_scattered=False), ), ): - self.assertIs(communicator._batch_steps(None), expected) + self.assertIs(communicator._batch_steps(fb), expected) # --------------------------------------------------------------------------- From d90efbd8e013424740a357cea21d6c4674b00a55 Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Fri, 25 Sep 2026 19:21:43 -0700 Subject: [PATCH 3/3] Bind a prefill CP's moves and the sum they complete in one place The FFN input selector and the output selector each asked which kind of prefill CP ran to pick the CP gather and the take-back, and the FFN exit let a CP extend's MoE skip its sum because the batch ran the CP steps. The reduce-scatter under DSA and MLA CP was chosen without checking which group the MoE owes its sum over. A layer's CP steps now take one CpMoves, chosen at construction for the kind of prefill CP: the gather after each rank completes its block, the take-back of a complete output, and, under DSA and MLA CP, the reduce-scatter over the attention-CP group, which completes a sum left over the same ranks. The output selector hands the FFN's sum to the reduce-scatter only when the sum is owed over those ranks and rejects the layer otherwise; a complete output is only taken back, as its equal-length shard. BoundarySteps records whether its output move completes the sum, and the FFN exit publishes the reduce-scatter from that. The functions the reachable configurations run are the ones they ran before. --- python/sglang/srt/layers/communicator.py | 127 ++++++-- .../unit/layers/test_declared_attention_cp.py | 277 ++++++++++++++++++ .../layers/test_declared_decoder_boundary.py | 11 +- 3 files changed, 382 insertions(+), 33 deletions(-) create mode 100644 test/registered/unit/layers/test_declared_attention_cp.py diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index d0af80b26443..22020e90230d 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -787,6 +787,7 @@ def __init__( self._declared_sides(cp_active=True), fusions=self._select_mlp_input_fusions(), force_layernorm_before_gather=force_layernorm_before_dp_gather, + cp_moves=_cp_moves(), ) if sides is not None and get_parallel().attn_cp_size > 1 else None @@ -882,15 +883,13 @@ def on_local_rows(sparse: bool) -> bool: and (modes.is_first_layer or modes.is_previous_layer_sparse is not None) ): return None - if _gathers_over_attention_cp(): - # DSA and MLA CP: a CP extend's FFN leaves its sum to the - # reduce-scatter that takes each rank's shard back. + if parallel.attn_cp_size > 1 and _cp_moves().reduce_scatter is not None: + # A CP extend's FFN may leave its sum to the reduce-scatter that + # takes each rank's shard back (DSA and MLA CP). may_leave = not cp_active may_leave_to_reduce_scatter = True else: - # Under GQA prefill CP the FFN completes its own sum: the next - # layer's input holds only this rank's chunk, and a reduce-scatter - # back over attention DP would split across the CP ranks. + # Otherwise under CP the FFN completes its own sum. may_leave = may_leave_to_reduce_scatter = parallel.attn_cp_size == 1 return decoder_layer_sides( axis_sizes=_token_axis_sizes(cp_active=cp_active), @@ -1331,7 +1330,7 @@ def _ffn_leaves_sum_to_reduce_scatter( steps = self._batch_steps(forward_batch) if not steps.ffn_output.leaves_for_reduce_scatter: return False - if dp_step is not None or steps is self._cp_steps: + if dp_step is not None or steps.ffn_output_move_completes_sum: return True # The scatter-mode steps of a DSA or MLA CP extend (the subclasses that # pick their own steps). Prefill CP predicates must stay out of decode @@ -2015,6 +2014,41 @@ def _batch_shards_over_cp(forward_batch: ForwardBatch) -> bool: return moe_cp_gathered_rows(forward_batch) is not None +class CpMoves(msgspec.Struct, frozen=True): + """How a CP extend's rows reach an FFN that needs all of them and come back, + chosen once for the kind of prefill CP: ``gather`` gathers the FFN input + after each rank has completed its own block, ``take_back`` returns this + rank's block of a complete output, and ``reduce_scatter``, where there is + one, completes a sum left over the ranks of ``reduce_scatter_group()`` and + returns the block in the same collective.""" + + gather: Callable + take_back: Callable + reduce_scatter: Optional[Callable] = None + reduce_scatter_group: Optional[Callable[[], GroupCoordinator]] = None + + +def _cp_moves() -> CpMoves: + """DSA and MLA CP gather equal shards over the attention-CP group and can + complete a sum over it. GQA prefill CP gathers blocks padded to the longest + over the MoE-CP group and takes back only a complete output.""" + if _gathers_over_attention_cp(): + return CpMoves( + gather=_mlp_input_gather_attention_cp, + take_back=CommunicateSummableTensorPairFn._take_back_attention_cp_shard, + reduce_scatter=CommunicateSummableTensorPairFn._reduce_scatter_over_cp, + reduce_scatter_group=lambda: get_parallel().attn_cp_group, + ) + return CpMoves( + gather=_mlp_input_gather_moe_cp, + take_back=CommunicateSummableTensorPairFn._scatter_hidden_states_moe, + ) + + +def _same_ranks(a: GroupCoordinator, b: GroupCoordinator) -> bool: + return sorted(a.ranks) == sorted(b.ranks) + + class BoundarySteps(msgspec.Struct, frozen=True): """The steps a batch runs at a layer's boundaries: into the attention, from the attention output to the FFN input, and the FFN output on to the @@ -2029,6 +2063,8 @@ class BoundarySteps(msgspec.Struct, frozen=True): ffn_output_move: Optional[Callable] # Whether the next layer's input can take the FFN's sum. ffn_sum_is_movable: bool + # Whether ffn_output_move also completes the sum the FFN leaves. + ffn_output_move_completes_sum: bool = False # The fused kernels ffn_input tries first. fused: Tuple["FusedMlpInput", ...] = () # Completes what the layer's input owes before the input norm; None when it @@ -2045,10 +2081,15 @@ def _select_boundary_steps( *, fusions: Tuple["FusedMlpInput", ...] = (), force_layernorm_before_gather: bool = False, + cp_moves: Optional[CpMoves] = None, ) -> BoundarySteps: - """The steps a set of declarations chooses.""" - returns_over_dp, ffn_output_move = _select_ffn_output_move( - sides.ffn_output, residual=sides.ffn_residual_rows, to=sides.output_rows + """The steps a set of declarations chooses; ``cp_moves`` for the ones + that gather over attention CP.""" + returns_over_dp, ffn_output_move, completes_sum = _select_ffn_output_move( + sides.ffn_output, + residual=sides.ffn_residual_rows, + to=sides.output_rows, + cp_moves=cp_moves, ) rows = sides.input_rows layer_input = None @@ -2071,12 +2112,14 @@ def _select_boundary_steps( force_layernorm_before_gather=force_layernorm_before_gather, fusions=fusions, residual_joins_sum=sides.residual_joins_attention_sum, + cp_moves=cp_moves, ) return BoundarySteps( attention_input=_select_attention_input_move(rows, sides.attention), ffn_input=ffn_input, ffn_output=sides.ffn_output, ffn_output_move=None if returns_over_dp else ffn_output_move, + ffn_output_move_completes_sum=completes_sum, ffn_sum_is_movable=sides.ffn_output.group is not None, fused=fused, layer_input=layer_input, @@ -2120,6 +2163,7 @@ def _select_ffn_input( force_layernorm_before_gather: bool, fusions: Tuple[FusedMlpInput, ...], residual_joins_sum: bool = False, + cp_moves: Optional[CpMoves] = None, ) -> Tuple[Callable, Tuple[FusedMlpInput, ...]]: """The steps from the attention output to the FFN input, and the fused kernels they try first: complete the attention-TP sum, move the residual to @@ -2153,9 +2197,9 @@ def _select_ffn_input( (), ) if gathered == {TokenAxis.ATTN_CP}: - # Each CP rank completes its own chunk, then the chunks are gathered: - # under DSA and MLA CP over the attention-CP group in equal shards, - # otherwise over the MoE-CP group, each padded to the longest. + # Each CP rank completes its own chunk, then the CP moves gather them. + if cp_moves is None: + raise NotImplementedError(f"{produced=} {need=}") on_chunk, fused = _select_ffn_input( produced, residual=residual, @@ -2165,9 +2209,7 @@ def _select_ffn_input( fusions=fusions, residual_joins_sum=residual_joins_sum, ) - if _gathers_over_attention_cp(): - return partial(_mlp_input_gather_attention_cp, gather=on_chunk), fused - return partial(_mlp_input_gather_moe_cp, gather=on_chunk), fused + return partial(cp_moves.gather, gather=on_chunk), fused if ( residual_to != produced.layout or gathered @@ -2225,38 +2267,47 @@ def _select_ffn_input( def _select_ffn_output_move( - produced: StageOutput, *, residual: Layout, to: Layout -) -> Tuple[bool, Optional[Callable]]: + produced: StageOutput, + *, + residual: Layout, + to: Layout, + cp_moves: Optional[CpMoves] = None, +) -> Tuple[bool, Optional[Callable], bool]: """How the FFN output reaches the rows the layer hands on: whether it goes back by undoing the attention-DP gather (the FFN exit and postprocess run that step), or else the postprocess that moves it, None when there is none - to choose here.""" + to choose here; and whether that move also completes the sum the FFN + leaves.""" if produced.layout == residual: if to == residual: - return False, CommunicateSummableTensorPairFn._trivial + return False, CommunicateSummableTensorPairFn._trivial, False if to.sharded == residual.sharded - {TokenAxis.ATTN_TP_SCATTER}: # Each rank's slice back to the attention's rows: fold the residual # into the output, then gather over attention TP. - return False, CommunicateSummableTensorPairFn._gather + return False, CommunicateSummableTensorPairFn._gather, False raise NotImplementedError(f"{produced=} {residual=} {to=}") returned = residual.sharded - produced.layout.sharded if to != residual or not produced.layout.sharded <= residual.sharded: raise NotImplementedError(f"{produced=} {residual=} {to=}") if returned == {TokenAxis.ATTN_CP}: - if _gathers_over_attention_cp(): - # The reduce-scatter over attention CP completes the sum the FFN - # left and takes this rank's shard back. - if not produced.leaves_for_reduce_scatter: - raise NotImplementedError(f"{produced=} {residual=} {to=}") - return False, CommunicateSummableTensorPairFn._reduce_scatter_over_cp - # This rank's chunk of the rows gathered over CP; no collective. - return False, CommunicateSummableTensorPairFn._scatter_hidden_states_moe + if cp_moves is None: + raise NotImplementedError(f"{produced=} {residual=} {to=}") + if not produced.leaves_for_reduce_scatter: + # A complete output: this rank's block of it, nothing summed. + return False, cp_moves.take_back, False + # The FFN leaves its sum: only a take-back that sums over the same + # ranks completes it. + if cp_moves.reduce_scatter is None or not _same_ranks( + _sum_group(produced.group), cp_moves.reduce_scatter_group() + ): + raise NotImplementedError(f"{produced=} {residual=} {to=}") + return False, cp_moves.reduce_scatter, True if returned == {TokenAxis.ATTN_DP, TokenAxis.ATTN_CP}: # This rank's CP shard, from where the DP gather put it. - return False, CommunicateSummableTensorPairFn._take_back_cp_shard + return False, CommunicateSummableTensorPairFn._take_back_cp_shard, False if returned != {TokenAxis.ATTN_DP}: raise NotImplementedError(f"{produced=} {residual=} {to=}") - return True, None + return True, None, False class MlpInputKind(Enum): @@ -2645,6 +2696,20 @@ def _scatter( assert residual is None, "not yet handled residual!=None" return _redistribute_to_attn_tp_shards(hidden_states, context), None + @staticmethod + def _take_back_attention_cp_shard( + hidden_states: torch.Tensor, + residual: torch.Tensor, + forward_batch: ForwardBatch, + context: CommunicateContext, + **kwargs, + ): + """DSA and MLA CP: this rank's shard of a complete output gathered in + equal shards over the attention-CP group.""" + parallel = get_parallel() + shard = hidden_states.tensor_split(parallel.attn_cp_size)[parallel.attn_cp_rank] + return shard, residual + @staticmethod def _reduce_scatter_over_cp( hidden_states: torch.Tensor, diff --git a/test/registered/unit/layers/test_declared_attention_cp.py b/test/registered/unit/layers/test_declared_attention_cp.py new file mode 100644 index 000000000000..44efc840bd43 --- /dev/null +++ b/test/registered/unit/layers/test_declared_attention_cp.py @@ -0,0 +1,277 @@ +"""A MoE layer on the TP group under DSA (and MLA) prefill CP, rank by rank. + +Two CP ranks build the layer's communicator for real and run prepare_mlp and +the FFN exit on a CP extend on CPU: the FFN input is gathered over attention CP +in equal shards, and the output comes back to each rank's shard. The +attention-CP all-gather and reduce-scatter are replaced by what the ranks hand +to them. +""" + +import contextlib +import unittest +from contextlib import ExitStack, contextmanager, nullcontext +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.layers import communicator as comm +from sglang.srt.layers import communicator_dsa_cp as dsa_cp +from sglang.srt.layers import layernorm_sp +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +HIDDEN = 4 +CP_SIZE = 2 +# Rows per CP shard: DSA and MLA CP shard a CP extend in equal lengths. +ROWS = 3 +# Binary fractions, so the partial sums add back to the value exactly. +PARTIAL_WEIGHTS = [0.25, 0.75] + + +def layernorm(hidden_states, residual=None): + """Norm as the identity, keeping the fused add of the two-argument form.""" + if residual is None: + return hidden_states.clone() + summed = hidden_states + residual + return summed, summed.clone() + + +class Flags: + def __init__(self): + self.fuse_mlp_allreduce = False + self.mlp_reduce_scatter = False + self.defer_moe_finalize = False + self.sp_active = False + + @contextlib.contextmanager + def scoped(self, **flags): + saved = {k: getattr(self, k) for k in flags} + self.__dict__.update(flags) + try: + yield + finally: + self.__dict__.update(saved) + + +def group(name, ranks): + return SimpleNamespace(name=name, ranks=list(ranks)) + + +class TestAttentionCpBoundary(CustomTestCase): + def setUp(self): + generator = torch.Generator().manual_seed(0) + self.values = [ + torch.randint(-8, 8, (ROWS, HIDDEN), generator=generator).double() + for _ in range(CP_SIZE) + ] + self.residuals = [ + torch.randint(-8, 8, (ROWS, HIDDEN), generator=generator).double() + for _ in range(CP_SIZE) + ] + + @contextmanager + def as_rank(self, cp, collectives, moe_group=None): + """Rank ``cp`` of attention CP 2 (attention DP and TP 1, TP 2), with + the attention-CP collectives the case gives.""" + parallel = SimpleNamespace( + tp_size=CP_SIZE, + tp_rank=cp, + attn_dp_size=1, + attn_dp_rank=0, + enable_dp_attention=False, + attn_tp_size=1, + attn_tp_rank=0, + attn_cp_size=CP_SIZE, + attn_cp_rank=cp, + enable_prefill_cp=True, + moe_dense_tp_size=1, + moe_dp_size=1, + moe_ep_size=1, + moe_tp_size=CP_SIZE, + dwdp_size=1, + enable_attn_tp_input_scattered=False, + tp_group=group("tp", range(CP_SIZE)), + attn_tp_group=group("attn_tp", [cp]), + attn_cp_group=group("attn_cp", range(CP_SIZE)), + ) + flags = Flags() + with ExitStack() as stack: + for target, value in [ + ((comm, "get_parallel"), lambda: parallel), + ((dsa_cp, "get_parallel"), lambda: parallel), + ((comm, "is_dsa_enable_prefill_cp"), lambda: True), + ((comm, "is_mla_cp_enabled"), lambda: False), + ( + (comm, "dsa_use_prefill_cp"), + lambda fb: fb.forward_mode.is_context_parallel_extend(), + ), + ((comm, "is_mla_cp_active"), lambda fb: False), + ((comm, "is_moe_input_scattered_across_dp_ranks"), lambda: False), + ((comm, "is_enable_moe_cp_allgather"), lambda: False), + ((comm, "get_moe_cp_size"), lambda: CP_SIZE), + ((comm, "get_moe_cp_rank"), lambda: cp), + ((comm, "should_use_dp_reduce_scatterv"), lambda: False), + ( + (comm, "get_moe_a2a_backend"), + lambda: SimpleNamespace(is_none=lambda: True), + ), + ( + (comm, "should_use_flashinfer_cutlass_moe_fp4_allgather"), + lambda: False, + ), + ( + (comm, "post_experts_reduction_group"), + lambda: moe_group or parallel.tp_group, + ), + ( + (comm, "get_spec"), + lambda: SimpleNamespace(speculative_algorithm=None), + ), + ((comm, "get_forward"), lambda: flags), + ( + (comm, "get_attn_tp_context"), + lambda: SimpleNamespace(input_scattered=False), + ), + ((layernorm_sp, "layernorm_sp_enabled"), lambda: False), + ((comm, "use_symmetric_memory"), lambda *a, **k: nullcontext()), + ((comm, "is_allocation_symmetric"), lambda: False), + ( + (dsa_cp, "get_local_dp_buffer"), + lambda g: torch.empty(ROWS * CP_SIZE, HIDDEN).double(), + ), + ((dsa_cp, "attn_cp_all_gather_into_tensor"), collectives["gather"]), + ( + (dsa_cp, "attn_cp_reduce_scatter_tensor"), + collectives["reduce_scatter"], + ), + ]: + stack.enter_context(patch.object(*target, value)) + yield SimpleNamespace(parallel=parallel, flags=flags) + + def build(self, allow_reduce_scatter): + # A MoE layer after a dense one; dense layers run on every rank here. + return comm.LayerCommunicator( + layer_scatter_modes=SimpleNamespace( + is_first_layer=False, + is_last_layer=False, + is_layer_sparse=True, + is_previous_layer_sparse=False, + ), + input_layernorm=layernorm, + post_attention_layernorm=layernorm, + allow_reduce_scatter=allow_reduce_scatter, + ) + + def cp_extend(self): + return SimpleNamespace( + forward_mode=SimpleNamespace(is_context_parallel_extend=lambda: True) + ) + + def run_ranks(self, allow_reduce_scatter): + """Gather on each rank, run an FFN that leaves a partial sum or not, + finish the exit and return what each rank got back and published.""" + handed = {} + + def record_gather(cp): + def gather(output, local): + handed[cp] = local.clone() + + return gather + + def fill_gather(output, local): + output.copy_(torch.cat([handed[cp] for cp in range(CP_SIZE)])) + + def unused(*args): + raise AssertionError("no reduce-scatter here") + + # The all-gather: record each rank's shard, then give every rank all. + for cp in range(CP_SIZE): + with self.as_rank( + cp, dict(gather=record_gather(cp), reduce_scatter=unused) + ): + self.build(allow_reduce_scatter).prepare_mlp( + self.values[cp], self.residuals[cp], self.cp_extend() + ) + gathered, residuals = {}, {} + for cp in range(CP_SIZE): + with self.as_rank(cp, dict(gather=fill_gather, reduce_scatter=unused)): + gathered[cp], residuals[cp] = self.build( + allow_reduce_scatter + ).prepare_mlp(self.values[cp], self.residuals[cp], self.cp_extend()) + expected_rows = torch.cat( + [self.values[cp] + self.residuals[cp] for cp in range(CP_SIZE)] + ) + for cp in range(CP_SIZE): + torch.testing.assert_close(gathered[cp], expected_rows, rtol=0, atol=0) + + def ffn_output(cp, leaves): + return gathered[cp] * PARTIAL_WEIGHTS[cp] if leaves else gathered[cp] + + reduced = {} + + def record_reduce_scatter(cp): + def reduce_scatter(output, input_): + reduced[cp] = input_.clone() + + return reduce_scatter + + def fill_reduce_scatter(cp): + def reduce_scatter(output, input_): + total = sum(reduced[r] for r in range(CP_SIZE)) + output.copy_(total.tensor_split(CP_SIZE)[cp]) + + return reduce_scatter + + published, back = {}, {} + for phase in ("record", "fill"): + for cp in range(CP_SIZE): + reduce_scatter = ( + record_reduce_scatter(cp) + if phase == "record" + else fill_reduce_scatter(cp) + ) + with self.as_rank( + cp, dict(gather=fill_gather, reduce_scatter=reduce_scatter) + ) as rank: + communicator = self.build(allow_reduce_scatter) + with communicator.ffn_exit(self.cp_extend()) as exit_: + published[cp] = rank.flags.mlp_reduce_scatter + output = ffn_output(cp, leaves=published[cp]) + back[cp], _ = exit_.finish(output, residuals[cp]) + return published, back, reduced + + def test_the_reduce_scatter_completes_the_sum_the_moe_leaves(self): + published, back, reduced = self.run_ranks(allow_reduce_scatter=True) + self.assertEqual(published, {0: True, 1: True}, "the MoE leaves its sum") + self.assertEqual(sorted(reduced), [0, 1], "every rank joins the reduce-scatter") + for cp in range(CP_SIZE): + torch.testing.assert_close( + back[cp], self.values[cp] + self.residuals[cp], rtol=0, atol=0 + ) + + def test_a_complete_output_is_only_taken_back(self): + published, back, reduced = self.run_ranks(allow_reduce_scatter=False) + self.assertEqual(published, {0: False, 1: False}, "the MoE sums itself") + self.assertEqual(reduced, {}, "nothing is summed again") + for cp in range(CP_SIZE): + torch.testing.assert_close( + back[cp], self.values[cp] + self.residuals[cp], rtol=0, atol=0 + ) + + def test_a_sum_over_other_ranks_is_not_left_to_it(self): + # A MoE group narrower than attention CP, as when MoE DP splits CP. + def unused(*args): + raise AssertionError("not run") + + with self.as_rank( + 0, dict(gather=unused, reduce_scatter=unused), moe_group=group("moe", [0]) + ): + with self.assertRaises(NotImplementedError): + self.build(allow_reduce_scatter=True) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_declared_decoder_boundary.py b/test/registered/unit/layers/test_declared_decoder_boundary.py index 017cfe9645b9..d0b3cdaa9cd3 100644 --- a/test/registered/unit/layers/test_declared_decoder_boundary.py +++ b/test/registered/unit/layers/test_declared_decoder_boundary.py @@ -61,8 +61,11 @@ def parallel_of(*, attn_dp, attn_tp, attn_cp=1, **overrides): dwdp_size=1, enable_dp_attention=attn_dp > 1, enable_attn_tp_input_scattered=False, - tp_group=SimpleNamespace(name="tp"), - attn_tp_group=SimpleNamespace(name="attn_tp"), + tp_group=SimpleNamespace( + name="tp", ranks=list(range(attn_dp * attn_cp * attn_tp)) + ), + attn_tp_group=SimpleNamespace(name="attn_tp", ranks=list(range(attn_tp))), + attn_cp_group=SimpleNamespace(name="attn_cp", ranks=list(range(attn_cp))), ) fields.update(overrides) return SimpleNamespace(**fields) @@ -103,6 +106,10 @@ def moe_cp_gathers(): ), patch.object(comm, "is_enable_moe_cp_allgather", moe_cp_gathers), patch.object(comm, "get_lora", lambda: SimpleNamespace(enable_lora=False)), + # A MoE whose EP and TP sums merge: its output's group is the TP group. + patch.object( + comm, "post_experts_reduction_group", lambda: get_parallel().tp_group + ), # Planning asks whether a dense layer gathers for two-batch overlap. patch.object( comm,