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
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
|
use_symmetric_memory,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
attn_tp_all_gather_into_tensor,
|
attn_tp_all_gather_into_tensor,
|
||||||
@@ -15,6 +18,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
get_attention_tp_group,
|
get_attention_tp_group,
|
||||||
get_attention_tp_rank,
|
get_attention_tp_rank,
|
||||||
get_attention_tp_size,
|
get_attention_tp_size,
|
||||||
|
is_allocation_symmetric,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils.common import ceil_align, ceil_div
|
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]
|
pad_size = max_len - input_.shape[0]
|
||||||
if pad_size > 0:
|
if pad_size > 0:
|
||||||
input_ = F.pad(input_, (0, 0, 0, pad_size), mode="constant", value=0)
|
input_ = F.pad(input_, (0, 0, 0, pad_size), mode="constant", value=0)
|
||||||
input_tensor_all = torch.empty(
|
with use_symmetric_memory(
|
||||||
max_len * attn_tp_size,
|
get_attention_tp_group(), disabled=not is_allocation_symmetric()
|
||||||
input_.shape[1],
|
):
|
||||||
device=input_.device,
|
input_tensor_all = torch.empty(
|
||||||
dtype=input_.dtype,
|
max_len * attn_tp_size,
|
||||||
)
|
input_.shape[1],
|
||||||
|
device=input_.device,
|
||||||
|
dtype=input_.dtype,
|
||||||
|
)
|
||||||
# step2
|
# step2
|
||||||
get_attention_tp_group().cp_all_gather_into_tensor_async(
|
get_attention_tp_group().cp_all_gather_into_tensor_async(
|
||||||
input_tensor_all, input_, stream_op
|
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():
|
if is_nsa_prefill_cp_round_robin_split():
|
||||||
output_tensor = input_tensor.new_empty(
|
with use_symmetric_memory(
|
||||||
(input_tensor.shape[0] * cp_size, *input_tensor.shape[1:]),
|
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(
|
attn_tp_all_gather_into_tensor(
|
||||||
output_tensor,
|
output_tensor,
|
||||||
input_tensor,
|
input_tensor,
|
||||||
|
|||||||
Reference in New Issue
Block a user