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