Register cp-atten-allgather buffers with symm memory (#17756)
Signed-off-by: wangfakang <fakangwang@gmail.com>
This commit is contained in:
@@ -8,6 +8,9 @@ import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
attn_tp_all_gather_into_tensor,
|
||||
@@ -15,6 +18,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
get_attention_tp_group,
|
||||
get_attention_tp_rank,
|
||||
get_attention_tp_size,
|
||||
is_allocation_symmetric,
|
||||
)
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils.common import ceil_align, ceil_div
|
||||
@@ -294,12 +298,15 @@ def cp_attn_tp_all_gather_reorganazied_into_tensor(
|
||||
pad_size = max_len - input_.shape[0]
|
||||
if pad_size > 0:
|
||||
input_ = F.pad(input_, (0, 0, 0, pad_size), mode="constant", value=0)
|
||||
input_tensor_all = torch.empty(
|
||||
max_len * attn_tp_size,
|
||||
input_.shape[1],
|
||||
device=input_.device,
|
||||
dtype=input_.dtype,
|
||||
)
|
||||
with use_symmetric_memory(
|
||||
get_attention_tp_group(), disabled=not is_allocation_symmetric()
|
||||
):
|
||||
input_tensor_all = torch.empty(
|
||||
max_len * attn_tp_size,
|
||||
input_.shape[1],
|
||||
device=input_.device,
|
||||
dtype=input_.dtype,
|
||||
)
|
||||
# step2
|
||||
get_attention_tp_group().cp_all_gather_into_tensor_async(
|
||||
input_tensor_all, input_, stream_op
|
||||
@@ -348,9 +355,12 @@ def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
|
||||
| +-------------------------+
|
||||
"""
|
||||
if is_nsa_prefill_cp_round_robin_split():
|
||||
output_tensor = input_tensor.new_empty(
|
||||
(input_tensor.shape[0] * cp_size, *input_tensor.shape[1:]),
|
||||
)
|
||||
with use_symmetric_memory(
|
||||
get_attention_tp_group(), disabled=not is_allocation_symmetric()
|
||||
):
|
||||
output_tensor = input_tensor.new_empty(
|
||||
(input_tensor.shape[0] * cp_size, *input_tensor.shape[1:]),
|
||||
)
|
||||
attn_tp_all_gather_into_tensor(
|
||||
output_tensor,
|
||||
input_tensor,
|
||||
|
||||
Reference in New Issue
Block a user