[PD] Early-send cached-prefix KV overlapping uncached prefill forward (#29316)

This commit is contained in:
cctry
2026-06-25 13:31:53 -07:00
committed by GitHub
parent e6efe10072
commit dbe9e3b706
3 changed files with 24 additions and 3 deletions
+16 -3
View File
@@ -469,8 +469,6 @@ class SchedulerDisaggregationPrefillMixin:
@torch.no_grad()
def event_loop_normal_disagg_prefill(self: Scheduler) -> None:
"""A normal scheduler loop for prefill worker in disaggregation mode."""
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
while True:
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
@@ -502,7 +500,6 @@ class SchedulerDisaggregationPrefillMixin:
@torch.no_grad()
def event_loop_overlap_disagg_prefill(self: Scheduler) -> None:
self.result_queue = deque()
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
while True:
# Receive requests
@@ -956,6 +953,22 @@ class SchedulerDisaggregationPrefillMixin:
if self.last_batch.batch_size() < last_bs:
self.running_batch.batch_is_full = False
def maybe_send_cached_prefix_chunk(self: Scheduler, req: Req) -> None:
# Only bootstrap-finalized requests; staging excluded.
if (
not envs.SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX.get()
or self.enable_staging
or req.pending_bootstrap
):
return
# Device-resident prefix only; page-aligned so start_send_idx stays exact.
cached_end = len(req.prefix_indices) - req.host_hit_length
if cached_end <= req.start_send_idx:
return
assert cached_end % self.token_to_kv_pool_allocator.page_size == 0
self.send_kv_chunk(req, last_chunk=False, end_idx=cached_end)
def send_kv_chunk(
self: Scheduler,
req: Req,
+1
View File
@@ -309,6 +309,7 @@ class Envs:
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}")
SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX = EnvBool(True)
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False)
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
+7
View File
@@ -1212,6 +1212,8 @@ class Scheduler(
# The prefill requests that are in the middle of kv sending
self.disagg_prefill_inflight_queue: List[Req] = []
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
# Init mm receiver for EPD disaggregation mode
if (
self.server_args.language_only
@@ -3204,6 +3206,11 @@ class Scheduler(
if batch.forward_mode.is_prebuilt():
return self._run_batch_prebuilt(batch)
# PD prefill: early-send cached prefix KV, overlapping the suffix forward.
if self.disaggregation_mode == DisaggregationMode.PREFILL:
for req in batch.reqs:
self.maybe_send_cached_prefix_chunk(req)
# Run forward
if self.is_generation:
if self.enable_overlap: