Fix DeepSeek-V4/DeepSeek-V4-Pro DP-attention gather semantics (#31700)
This commit is contained in:
@@ -1942,7 +1942,9 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
if _do_shared_local and local_hidden_states.shape[0] > 0:
|
if _do_shared_local and local_hidden_states.shape[0] > 0:
|
||||||
_shared_local = self.mlp._forward_shared_experts(local_hidden_states)
|
_shared_local = self.mlp._forward_shared_experts(local_hidden_states)
|
||||||
dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
|
# self_attn has already reduced across attention TP, so these hidden
|
||||||
|
# states are replicated and must not be summed by a partial gather.
|
||||||
|
dp_gather_replicate(hidden_states, local_hidden_states, forward_batch)
|
||||||
_a2a_scatter_chunks: Optional[List[torch.Tensor]] = None
|
_a2a_scatter_chunks: Optional[List[torch.Tensor]] = None
|
||||||
if _use_tp_attn_a2a_scatter:
|
if _use_tp_attn_a2a_scatter:
|
||||||
s, r = get_parallel().attn_tp_size, get_parallel().attn_tp_rank
|
s, r = get_parallel().attn_tp_size, get_parallel().attn_tp_rank
|
||||||
@@ -2465,7 +2467,10 @@ class DeepseekV4Model(nn.Module):
|
|||||||
)
|
)
|
||||||
# Token ids are replicated within an attention-TP group. Use replicate
|
# Token ids are replicated within an attention-TP group. Use replicate
|
||||||
# gather here to avoid summing duplicated ids when attention_tp_size > 1.
|
# gather here to avoid summing duplicated ids when attention_tp_size > 1.
|
||||||
dp_gather_replicate(input_ids_global, input_ids[:, None], forward_batch)
|
# Clone because the MAX_LEN gather may zero its local input in place.
|
||||||
|
dp_gather_replicate(
|
||||||
|
input_ids_global, input_ids[:, None].clone(), forward_batch
|
||||||
|
)
|
||||||
input_ids_global = input_ids_global.squeeze(-1)
|
input_ids_global = input_ids_global.squeeze(-1)
|
||||||
else:
|
else:
|
||||||
input_ids_global = input_ids
|
input_ids_global = input_ids
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.cp.utils import (
|
|||||||
is_cp_v2_active,
|
is_cp_v2_active,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
dp_gather_partial,
|
dp_gather_replicate,
|
||||||
get_global_dp_buffer_len,
|
get_global_dp_buffer_len,
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
@@ -171,7 +171,12 @@ class DeepseekV4ModelNextN(nn.Module):
|
|||||||
dtype=input_ids.dtype,
|
dtype=input_ids.dtype,
|
||||||
device=input_ids.device,
|
device=input_ids.device,
|
||||||
)
|
)
|
||||||
dp_gather_partial(input_ids_global, input_ids[:, None], forward_batch)
|
# Token IDs are replicated within an attention-TP group. Use replicate
|
||||||
|
# gather to avoid summing duplicated IDs when attention_tp_size > 1.
|
||||||
|
# Clone because the MAX_LEN gather may zero its local input in place.
|
||||||
|
dp_gather_replicate(
|
||||||
|
input_ids_global, input_ids[:, None].clone(), forward_batch
|
||||||
|
)
|
||||||
input_ids_global = input_ids_global.squeeze(-1)
|
input_ids_global = input_ids_global.squeeze(-1)
|
||||||
else:
|
else:
|
||||||
input_ids_global = input_ids
|
input_ids_global = input_ids
|
||||||
|
|||||||
Reference in New Issue
Block a user