[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:
co-authored by
Byron Hsu
Cursor
parent
7edb4c3cea
commit
47483001b6
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user