From dbe9e3b7062df3120b3252c83fb6e3c48ac07c2c Mon Sep 17 00:00:00 2001 From: cctry Date: Thu, 25 Jun 2026 13:31:53 -0700 Subject: [PATCH] [PD] Early-send cached-prefix KV overlapping uncached prefill forward (#29316) --- python/sglang/srt/disaggregation/prefill.py | 19 ++++++++++++++++--- python/sglang/srt/environ.py | 1 + python/sglang/srt/managers/scheduler.py | 7 +++++++ 3 files changed, 24 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index b0715afea..ee8c82aaa 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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, diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 6756be8c0..a1768ce05 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7055e25d9..ec785ec65 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: