[PD] Early-send cached-prefix KV overlapping uncached prefill forward (#29316)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user