[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
from sglang.srt.environ import envs
from sglang.srt.utils import get_bool_env_var
if TYPE_CHECKING:
@@ -45,6 +46,7 @@ class PrefillDelayer:
token_usage_low_watermark: Optional[float],
metrics_collector: Optional["SchedulerMetricsCollector"] = None,
device: Optional["torch.device"] = "cpu",
device_group=None,
):
self._max_delay_passes = max_delay_passes
self._token_usage_low_watermark = token_usage_low_watermark
@@ -68,15 +70,32 @@ class PrefillDelayer:
self.dp_size = dp_size
self.enable_dp_attention = server_args.enable_dp_attention
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,
# token_watermark_force_allow, running_batch, max_prefill_bs,
# waiting_queue_len.
self._global_info_buffer = torch.empty(
(dp_size_dim, attn_tp_size, 5),
dtype=torch.int64,
device=device,
device=self._gather_device,
)
self._cpu_group = cpu_group
self._metrics_collector = metrics_collector
@@ -277,13 +296,13 @@ class PrefillDelayer:
max_prefill_bs,
waiting_queue_len,
],
device="cpu",
device=self._gather_device,
dtype=torch.int64,
)
torch.distributed.all_gather_into_tensor(
self._global_info_buffer.flatten(),
local_info,
group=self._cpu_group,
group=self._gather_group,
)
tp0_info = self._global_info_buffer[:, 0, :]
return tp0_info
+2 -5
View File
@@ -1119,17 +1119,14 @@ class Scheduler(
dp_size=self.dp_size,
attn_tp_size=self.attn_tp_size,
cpu_group=self.tp_cpu_group,
device_group=self.tp_group.device_group,
server_args=self.server_args,
metrics_collector=(
self.metrics_collector if self.enable_metrics else None
),
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
device=(
self.tp_group.device
if self.server_args.disable_overlap_schedule
else "cpu"
),
device=self.tp_group.device,
)
# NOTE: preemption is enabled by default for priority scheduling.