[PrefillDelayer] support NCCL all-gather for cross-DP info sync (#24768)

Co-authored-by: Byron Hsu <byron@periodiclabs.ai>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Byron Hsu
2026-05-09 21:20:03 -07:00
committed by GitHub
co-authored by Byron Hsu Cursor
parent 7edb4c3cea
commit 47483001b6
2 changed files with 25 additions and 9 deletions
+23 -4
View File
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
import torch import torch
from sglang.srt.environ import envs
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -45,6 +46,7 @@ class PrefillDelayer:
token_usage_low_watermark: Optional[float], token_usage_low_watermark: Optional[float],
metrics_collector: Optional["SchedulerMetricsCollector"] = None, metrics_collector: Optional["SchedulerMetricsCollector"] = None,
device: Optional["torch.device"] = "cpu", device: Optional["torch.device"] = "cpu",
device_group=None,
): ):
self._max_delay_passes = max_delay_passes self._max_delay_passes = max_delay_passes
self._token_usage_low_watermark = token_usage_low_watermark self._token_usage_low_watermark = token_usage_low_watermark
@@ -68,15 +70,32 @@ class PrefillDelayer:
self.dp_size = dp_size self.dp_size = dp_size
self.enable_dp_attention = server_args.enable_dp_attention self.enable_dp_attention = server_args.enable_dp_attention
dp_size_dim = dp_size if self.enable_dp_attention else 1 dp_size_dim = dp_size if self.enable_dp_attention else 1
# Mirror scheduler_dp_attn_mixin's NCCL all-gather path: when the
# env flag is on (or overlap scheduling is disabled), ride the NCCL
# device group on `device` instead of gloo on CPU.
use_nccl = (
server_args.disable_overlap_schedule
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
)
if use_nccl:
assert (
device_group is not None
), "device_group is required when using NCCL for PrefillDelayer all-gather"
self._gather_group = device_group
self._gather_device = device
else:
self._gather_group = cpu_group
self._gather_device = "cpu"
# Fields packed per rank into the all-gather tensor: prefillable, # Fields packed per rank into the all-gather tensor: prefillable,
# token_watermark_force_allow, running_batch, max_prefill_bs, # token_watermark_force_allow, running_batch, max_prefill_bs,
# waiting_queue_len. # waiting_queue_len.
self._global_info_buffer = torch.empty( self._global_info_buffer = torch.empty(
(dp_size_dim, attn_tp_size, 5), (dp_size_dim, attn_tp_size, 5),
dtype=torch.int64, dtype=torch.int64,
device=device, device=self._gather_device,
) )
self._cpu_group = cpu_group
self._metrics_collector = metrics_collector self._metrics_collector = metrics_collector
@@ -277,13 +296,13 @@ class PrefillDelayer:
max_prefill_bs, max_prefill_bs,
waiting_queue_len, waiting_queue_len,
], ],
device="cpu", device=self._gather_device,
dtype=torch.int64, dtype=torch.int64,
) )
torch.distributed.all_gather_into_tensor( torch.distributed.all_gather_into_tensor(
self._global_info_buffer.flatten(), self._global_info_buffer.flatten(),
local_info, local_info,
group=self._cpu_group, group=self._gather_group,
) )
tp0_info = self._global_info_buffer[:, 0, :] tp0_info = self._global_info_buffer[:, 0, :]
return tp0_info return tp0_info
+2 -5
View File
@@ -1119,17 +1119,14 @@ class Scheduler(
dp_size=self.dp_size, dp_size=self.dp_size,
attn_tp_size=self.attn_tp_size, attn_tp_size=self.attn_tp_size,
cpu_group=self.tp_cpu_group, cpu_group=self.tp_cpu_group,
device_group=self.tp_group.device_group,
server_args=self.server_args, server_args=self.server_args,
metrics_collector=( metrics_collector=(
self.metrics_collector if self.enable_metrics else None self.metrics_collector if self.enable_metrics else None
), ),
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes, max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark, token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
device=( device=self.tp_group.device,
self.tp_group.device
if self.server_args.disable_overlap_schedule
else "cpu"
),
) )
# NOTE: preemption is enabled by default for priority scheduling. # NOTE: preemption is enabled by default for priority scheduling.