Register cp-atten-allgather buffers with symm memory (#17756)

Signed-off-by: wangfakang <fakangwang@gmail.com>
This commit is contained in:
sky
2026-02-11 15:26:37 +08:00
committed by GitHub
parent a8eef53dc4
commit 72c1526657

View File

@@ -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,