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/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..22020e90230d 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, @@ -774,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 @@ -814,9 +828,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 +838,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,36 +854,56 @@ 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()) - # 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 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 not enable_moe_dense_fully_dp() 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 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: + # 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), - 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, 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=( @@ -1081,7 +1116,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 @@ -1292,13 +1327,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.ffn_output_move_completes_sum: 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 @@ -1949,15 +1989,66 @@ 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 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 @@ -1972,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 @@ -1988,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 @@ -2014,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, @@ -2063,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 @@ -2096,8 +2197,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 CP moves gather them. + if cp_moves is None: + raise NotImplementedError(f"{produced=} {need=}") on_chunk, fused = _select_ffn_input( produced, residual=residual, @@ -2107,7 +2209,7 @@ def _select_ffn_input( fusions=fusions, residual_joins_sum=residual_joins_sum, ) - 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 @@ -2165,32 +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}: - # 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): @@ -2345,6 +2462,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, @@ -2561,6 +2696,32 @@ 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, + 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_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 8e181d953b3c..d0b3cdaa9cd3 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 @@ -62,18 +61,22 @@ 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) @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 +84,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 +106,18 @@ 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, + "get_exec", + lambda: SimpleNamespace( + overlap=SimpleNamespace(enable_two_batch_overlap=False) + ), + ), ): yield @@ -168,8 +183,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 +203,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, @@ -401,6 +432,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 +502,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 @@ -456,11 +534,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) @@ -948,6 +1022,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 @@ -982,16 +1149,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( @@ -1003,13 +1184,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): @@ -1031,13 +1222,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) ), @@ -1047,7 +1251,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) # ---------------------------------------------------------------------------