[AMD] Feat/dp moe reduce scatter (#28216)

Co-authored-by: wunhuang <wunhuang@amd.com>
Co-authored-by: Wang, FangYuan <39615225+At1a8@users.noreply.github.com>
This commit is contained in:
kk
2026-06-16 18:12:39 -07:00
committed by GitHub
co-authored by wunhuang Wang, FangYuan
parent d86a7e7018
commit 37ef295c78
3 changed files with 157 additions and 2 deletions
+96
View File
@@ -524,12 +524,100 @@ def _dp_gather_via_all_gather(
get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens)
# Variable-length DP-MoE gather (reference https://github.com/ROCm/ATOM/pull/930): instead of padding every
# rank to max_len (all_gather) or all-reducing a sum_len zero-buffer (all_reduce),
# gather exactly sum(per-rank tokens) via all_gatherv. Env-gated; only the simple
# tp_size==dp_size (attn_tp_size==1) case is supported for now (e.g. tp8dp8).
_USE_DP_GATHERV = get_bool_env_var("SGLANG_DP_USE_GATHERV")
def is_dp_gatherv_active() -> bool:
"""Variable-length DP-MoE gather/scatter (all_gatherv + reduce_scatterv) is
enabled and the current parallel layout (attn_tp_size==1, tp_size==dp_size)
is supported. Env-gated by SGLANG_DP_USE_GATHERV; default off."""
return (
_USE_DP_GATHERV
and get_attention_tp_size() == 1
and get_tensor_model_parallel_world_size() == get_attention_dp_size()
)
def _dp_gatherv_sizes(forward_batch) -> Optional[List[int]]:
"""Per-rank CPU token counts for the buffer being gathered. The MoE gather
passes a ForwardBatch (global_num_tokens_cpu); the logits gather passes a
LogitsMetadata (global_num_tokens_for_logprob_cpu). Return the sizes that
match the LOCAL tensor for this context, or None to fall back."""
sizes = getattr(forward_batch, "global_num_tokens_for_logprob_cpu", None)
if sizes is None:
sizes = getattr(forward_batch, "global_num_tokens_cpu", None)
if sizes is None:
return None
try:
return [int(x) for x in sizes]
except (TypeError, ValueError):
return None
def _dp_gather_via_all_gatherv(
global_tokens: torch.Tensor,
local_tokens: torch.Tensor,
forward_batch: ForwardBatch,
is_partial: bool,
sizes: List[int],
):
# attn_tp_size == 1: each DP rank contributes exactly `sizes[rank]` rows.
# CRITICAL: the MoE downstream runs on the WHOLE `global_tokens` buffer
# (M = global_tokens.shape[0]), so the gather MUST fill every row. We pad
# each rank's local tensor up to sizes[rank] with zeros (matching the
# buffer's reserved per-rank slot) so sum(sizes) == buffer rows and there
# is no uninitialized tail for the MoE to read.
rank = get_attention_dp_rank()
local_rows = sizes[rank]
if local_tokens.shape[0] == local_rows:
local_real = local_tokens
elif local_tokens.shape[0] > local_rows:
local_real = local_tokens[:local_rows]
else:
local_real = local_tokens.new_zeros((local_rows, *local_tokens.shape[1:]))
local_real[: local_tokens.shape[0]].copy_(local_tokens)
gathered = get_tp_group().all_gatherv(local_real, sizes=sizes)
if isinstance(gathered, list):
# all_gatherv may return a list of per-rank tensors; concatenate them
# along the token dim (taking [0] would drop all but rank 0's tokens).
gathered = torch.cat(gathered, dim=0)
# gathered rows == sum(sizes); must equal the buffer length.
global_tokens[: gathered.shape[0]].copy_(gathered)
def _dp_gather(
global_tokens: torch.Tensor,
local_tokens: torch.Tensor,
forward_batch: ForwardBatch,
is_partial: bool,
):
if (
is_dp_gatherv_active()
and forward_batch.dp_padding_mode is not None
and not forward_batch.dp_padding_mode.is_max_len()
):
# The gatherv per-rank sizes MUST sum to the pre-allocated global buffer
# (the MoE runs on the whole buffer, so any unfilled tail = garbage).
# The buffer was sized from the ceil_align'd global_num_tokens stored via
# set_dp_buffer_len (forward_batch_info), so the authoritative sizes are
# get_dp_global_num_tokens() — the SAME source the reduce_scatterv combine
# uses (symmetric). _dp_gatherv_sizes() reads the raw (un-aligned, and for
# the MoE-gather context the logprob-token) counts, which do NOT match the
# buffer for prefill steps -> would force an all_reduce fallback.
# Prefer the buffer-aligned sizes; fall back to the per-batch sizes only
# if they happen to match (e.g. the logits gather path).
_gatherv_sizes = get_dp_global_num_tokens()
if _gatherv_sizes is None or sum(_gatherv_sizes) != global_tokens.shape[0]:
_gatherv_sizes = _dp_gatherv_sizes(forward_batch)
if _gatherv_sizes is not None and sum(_gatherv_sizes) == global_tokens.shape[0]:
_dp_gather_via_all_gatherv(
global_tokens, local_tokens, forward_batch, is_partial, _gatherv_sizes
)
return
if forward_batch.dp_padding_mode.is_max_len():
_dp_gather_via_all_gather(
global_tokens, local_tokens, forward_batch, is_partial
@@ -579,6 +667,14 @@ def dp_scatter(
def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
if is_dp_gatherv_active():
# Variable-length combine matching all_gatherv dispatch: scatter the
# global (sum_len) tensor back to per-rank token counts. Fall through to
# the default reduce-scatter path if per-rank sizes are unavailable.
sizes = get_dp_global_num_tokens()
if sizes is not None:
get_tp_group().reduce_scatterv(input, output=output, sizes=sizes)
return
if get_tensor_model_parallel_world_size() == get_attention_dp_size():
get_tp_group().reduce_scatter_tensor(output, input)
else:
+22 -2
View File
@@ -62,6 +62,7 @@ from sglang.srt.layers.dp_attention import (
get_global_dp_buffer,
get_local_dp_buffer,
is_dp_attention_enabled,
is_dp_gatherv_active,
)
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
@@ -1502,6 +1503,19 @@ class DeepseekV4DecoderLayer(nn.Module):
and get_attention_tp_size() > 1
and not get_moe_a2a_backend().is_none()
)
# symmetric gather+scatter for the no-EP TP-MoE dp-attn path:
# all_gatherv gather (in self.mlp's dp_gather) + reduce_scatterv combine.
# The experts ARE TP-sharded by intermediate (moe_tp_size==tp_size), so
# the post-experts reduce is a SUM. reduce_scatterv does that sum+scatter
# in ONE op, REPLACING the MoE-internal post-experts all_reduce — so we
# MUST tell the MoE to skip it (use_reduce_scatter=True) or it
# double-reduces. Env-gated via SGLANG_DP_USE_GATHERV, default OFF.
_use_gatherv_pair = (
_use_tp_moe_gather
and is_dp_gatherv_active()
and forward_batch.dp_padding_mode is not None
and not forward_batch.dp_padding_mode.is_max_len()
)
if _use_cp:
if get_moe_a2a_backend().is_none():
hidden_states = dsa_cp_gather_hidden_states(hidden_states)
@@ -1528,7 +1542,9 @@ class DeepseekV4DecoderLayer(nn.Module):
forward_batch,
input_ids=input_ids,
input_ids_global=input_ids_global,
use_reduce_scatter=_use_cp,
# Skip the MoE-internal post-experts all_reduce when we will do the
# reduce via reduce_scatterv at the combine below (else double-reduce).
use_reduce_scatter=_use_cp or _use_gatherv_pair,
)
if _use_cp and get_moe_a2a_backend().is_none():
hidden_states = dsa_cp_reduce_scatter_hidden_states(hidden_states)
@@ -1537,7 +1553,11 @@ class DeepseekV4DecoderLayer(nn.Module):
get_local_dp_buffer(get_tp_group()),
hidden_states,
)
if should_use_dp_reduce_scatterv():
if should_use_dp_reduce_scatterv() or _use_gatherv_pair:
# SUM the TP-sharded per-rank partial expert outputs AND scatter
# each rank its own token slice, in one op. Correct because the
# MoE-internal all_reduce was skipped (use_reduce_scatter above).
# This is the symmetric inverse of the all_gatherv gather.
get_tp_group().reduce_scatterv(
global_hidden_states,
output=hidden_states,