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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/layers/boundary_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
241 changes: 201 additions & 40 deletions python/sglang/srt/layers/communicator.py

Large diffs are not rendered by default.

153 changes: 0 additions & 153 deletions python/sglang/srt/layers/communicator_dsa_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,23 +13,13 @@
# ==============================================================================


from functools import partial
from typing import Optional

import torch

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,
Expand Down Expand Up @@ -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
2 changes: 1 addition & 1 deletion python/sglang/srt/model_executor/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
9 changes: 5 additions & 4 deletions python/sglang/srt/models/deepseek_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading