Bugfix: fix symm not enabled due to incorrect registration of comm (#19329)
Signed-off-by: wangfakang <fakangwang@gmail.com>
This commit is contained in:
@@ -46,6 +46,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
get_attention_cp_rank,
|
get_attention_cp_rank,
|
||||||
get_attention_cp_size,
|
get_attention_cp_size,
|
||||||
get_attention_dp_size,
|
get_attention_dp_size,
|
||||||
|
get_attention_tp_group,
|
||||||
get_attention_tp_rank,
|
get_attention_tp_rank,
|
||||||
get_attention_tp_size,
|
get_attention_tp_size,
|
||||||
get_dp_global_num_tokens,
|
get_dp_global_num_tokens,
|
||||||
@@ -855,7 +856,7 @@ class CommunicateSimpleFn:
|
|||||||
return tuple(gathered_hidden_states)
|
return tuple(gathered_hidden_states)
|
||||||
|
|
||||||
hidden_states, local_hidden_states = (
|
hidden_states, local_hidden_states = (
|
||||||
get_local_dp_buffer(),
|
get_local_dp_buffer(get_attention_tp_group()),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
attn_tp_all_gather_into_tensor(
|
attn_tp_all_gather_into_tensor(
|
||||||
@@ -965,7 +966,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
|||||||
|
|
||||||
if residual_input_mode == ScatterMode.SCATTERED and context.attn_tp_size > 1:
|
if residual_input_mode == ScatterMode.SCATTERED and context.attn_tp_size > 1:
|
||||||
residual, local_residual = (
|
residual, local_residual = (
|
||||||
get_local_dp_buffer(),
|
get_local_dp_buffer(get_attention_tp_group()),
|
||||||
residual,
|
residual,
|
||||||
)
|
)
|
||||||
attn_tp_all_gather_into_tensor(residual, local_residual)
|
attn_tp_all_gather_into_tensor(residual, local_residual)
|
||||||
@@ -982,7 +983,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
|||||||
hidden_states += residual
|
hidden_states += residual
|
||||||
|
|
||||||
hidden_states, local_hidden_states = (
|
hidden_states, local_hidden_states = (
|
||||||
get_global_dp_buffer(),
|
get_global_dp_buffer(get_tp_group()),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
|
dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
|
||||||
@@ -1210,8 +1211,12 @@ class CommunicateSummableTensorPairFn:
|
|||||||
context: CommunicateContext,
|
context: CommunicateContext,
|
||||||
allow_reduce_scatter: bool = False,
|
allow_reduce_scatter: bool = False,
|
||||||
):
|
):
|
||||||
|
if get_tensor_model_parallel_world_size() == get_attention_dp_size():
|
||||||
|
group = get_tp_group()
|
||||||
|
else:
|
||||||
|
group = get_attention_tp_group()
|
||||||
hidden_states, global_hidden_states = (
|
hidden_states, global_hidden_states = (
|
||||||
get_local_dp_buffer(),
|
get_local_dp_buffer(group),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
if should_use_dp_reduce_scatterv():
|
if should_use_dp_reduce_scatterv():
|
||||||
@@ -1237,7 +1242,7 @@ class CommunicateSummableTensorPairFn:
|
|||||||
hidden_states += residual
|
hidden_states += residual
|
||||||
residual = None
|
residual = None
|
||||||
hidden_states, local_hidden_states = (
|
hidden_states, local_hidden_states = (
|
||||||
get_local_dp_buffer(),
|
get_local_dp_buffer(get_attention_tp_group()),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
attn_tp_all_gather_into_tensor(
|
attn_tp_all_gather_into_tensor(
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ from sglang.srt.layers.communicator import (
|
|||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
attn_cp_all_gather_into_tensor,
|
attn_cp_all_gather_into_tensor,
|
||||||
attn_cp_reduce_scatter_tensor,
|
attn_cp_reduce_scatter_tensor,
|
||||||
|
get_attention_cp_group,
|
||||||
get_local_dp_buffer,
|
get_local_dp_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
@@ -154,7 +155,7 @@ class NSACPCommunicateWithAllReduceAndLayerNormFn(
|
|||||||
if nsa_use_prefill_cp(forward_batch):
|
if nsa_use_prefill_cp(forward_batch):
|
||||||
assert context.attn_dp_size == 1
|
assert context.attn_dp_size == 1
|
||||||
hidden_states, local_hidden_states = (
|
hidden_states, local_hidden_states = (
|
||||||
get_local_dp_buffer(),
|
get_local_dp_buffer(get_attention_cp_group()),
|
||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
attn_cp_all_gather_into_tensor(
|
attn_cp_all_gather_into_tensor(
|
||||||
|
|||||||
@@ -126,8 +126,8 @@ class _DpGatheredBufferWrapper:
|
|||||||
cls._global_num_tokens = global_num_tokens
|
cls._global_num_tokens = global_num_tokens
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_global_dp_buffer(cls) -> torch.Tensor:
|
def get_global_dp_buffer(cls, group: GroupCoordinator) -> torch.Tensor:
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not cls._dp_max_padding):
|
with use_symmetric_memory(group, disabled=not cls._dp_max_padding):
|
||||||
buffer = torch.empty(
|
buffer = torch.empty(
|
||||||
(cls._global_dp_buffer_len, cls._hidden_size),
|
(cls._global_dp_buffer_len, cls._hidden_size),
|
||||||
dtype=cls._dtype,
|
dtype=cls._dtype,
|
||||||
@@ -136,8 +136,8 @@ class _DpGatheredBufferWrapper:
|
|||||||
return buffer
|
return buffer
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_local_dp_buffer(cls) -> torch.Tensor:
|
def get_local_dp_buffer(cls, group: GroupCoordinator) -> torch.Tensor:
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not cls._dp_max_padding):
|
with use_symmetric_memory(group, disabled=not cls._dp_max_padding):
|
||||||
buffer = torch.empty(
|
buffer = torch.empty(
|
||||||
(cls._local_dp_buffer_len, cls._hidden_size),
|
(cls._local_dp_buffer_len, cls._hidden_size),
|
||||||
dtype=cls._dtype,
|
dtype=cls._dtype,
|
||||||
@@ -193,12 +193,12 @@ def set_dp_buffer_len(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_global_dp_buffer() -> torch.Tensor:
|
def get_global_dp_buffer(group: GroupCoordinator) -> torch.Tensor:
|
||||||
return _DpGatheredBufferWrapper.get_global_dp_buffer()
|
return _DpGatheredBufferWrapper.get_global_dp_buffer(group=group)
|
||||||
|
|
||||||
|
|
||||||
def get_local_dp_buffer() -> torch.Tensor:
|
def get_local_dp_buffer(group: GroupCoordinator) -> torch.Tensor:
|
||||||
return _DpGatheredBufferWrapper.get_local_dp_buffer()
|
return _DpGatheredBufferWrapper.get_local_dp_buffer(group=group)
|
||||||
|
|
||||||
|
|
||||||
def get_global_dp_buffer_len() -> int:
|
def get_global_dp_buffer_len() -> int:
|
||||||
|
|||||||
@@ -220,7 +220,10 @@ class StandardDispatcher(BaseDispatcher):
|
|||||||
def combine(self, combine_input: StandardCombineInput) -> torch.Tensor:
|
def combine(self, combine_input: StandardCombineInput) -> torch.Tensor:
|
||||||
(hidden_states,) = combine_input
|
(hidden_states,) = combine_input
|
||||||
if should_use_flashinfer_cutlass_moe_fp4_allgather():
|
if should_use_flashinfer_cutlass_moe_fp4_allgather():
|
||||||
hidden_states, global_hidden_states = get_local_dp_buffer(), hidden_states
|
hidden_states, global_hidden_states = (
|
||||||
|
get_local_dp_buffer(get_tp_group()),
|
||||||
|
hidden_states,
|
||||||
|
)
|
||||||
get_tp_group().reduce_scatterv(
|
get_tp_group().reduce_scatterv(
|
||||||
global_hidden_states,
|
global_hidden_states,
|
||||||
output=hidden_states,
|
output=hidden_states,
|
||||||
|
|||||||
Reference in New Issue
Block a user