Files
sglang/python/sglang/srt/disaggregation/prefill.py
T
leaveletandClaude Fable 5 58f2350738 Give prefill inflight transfers a liveness bound
E2e caught one request wedged FOREVER in disagg_prefill_inflight_queue
(probes 21s apart with zero traffic both showed #inflight-req: 1, no
reap/timeout warnings ever logged).  Mechanism, established by reading
the full state machine: prefill Success is set locally by the transfer
worker on the LAST chunk; if the decode peer is torn down between the
handshake and the prefill's final send(), add_transfer_request silently
drops the chunk (no transfer destinations) — Success becomes
unreachable.  The only external rescue, the decode ABORT notification,
is best-effort (silently swallowed on send error, no-op if it races the
room registration), there is no prefill-side heartbeat of decode
sessions, and the sender's only timeout covers Bootstrapping — the
inflight queue itself has no liveness bound.  The orphan pins the
request's KV pages and rides every poll collective.

Two fixes, both reaped through the existing Failed branch via the
CP/TP MIN-reduce poll consensus (Failed=0 wins, so one rank concluding
flips every rank together — rank-uniform by construction):

- add_transfer_request: a room with no transfer destinations that is
  NOT already Success (the dummy-rank handshake marking) now concludes
  Failed loudly instead of dropping the chunk silently.
- Inflight residency timeout: entries stuck in a non-terminal poll
  state past SGLANG_DISAGGREGATION_INFLIGHT_TIMEOUT (default 300s,
  matching the sibling BOOTSTRAP/WAITING timeouts) get sender.abort()
  and reap on the next poll.  Covers what the hardening cannot: lost
  ABORT datagrams, decode crashes.

Known sibling gaps left for follow-up: the decode transfer queue has
no Transferring liveness bound, and an abort that matches no queue is
still a silent no-op (much narrower race than first thought — work
requests are ordered before control requests within a tick).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-12 04:55:35 +00:00

1407 lines
57 KiB
Python

"""
Life cycle of a request in the prefill server
1. Bootstrap Queue
a. Initialize a sender for each request
b. Use the queue to store requests whose bootstrap (handshake and preallocation) has not finished
c. Poll senders to check bootstrap state
d. Once bootstrap is complete, move request to Waiting Queue
2. Waiting Queue
a. Use PrefillAdder to pop requests
b. Run forward
c. Add the request to Inflight Queue
3. Inflight Queue
a. Poll (non-blocking) the sender of the request
b. Once the transfer has finished, return the request
"""
from __future__ import annotations
import logging
import time
from collections import deque
from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional
import torch
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST,
DisaggregationMode,
KVClassType,
MetadataBuffers,
ReqToMetadataIdxAllocator,
TransferBackend,
append_cp_draft_state_buffers,
get_kv_class,
is_mla_backend,
kv_to_page_num,
poll_and_all_reduce_attn_cp_tp_group,
prepare_abort,
)
from sglang.srt.managers.schedule_batch import (
FINISH_ABORT,
FINISH_LENGTH,
Req,
ScheduleBatch,
)
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.observability.req_time_stats import set_schedule_time_batch
if TYPE_CHECKING:
from torch.distributed import ProcessGroup
from sglang.srt.managers.scheduler import GenerationBatchResult, Scheduler
from sglang.srt.mem_cache.memory_pool import KVCache
logger = logging.getLogger(__name__)
def _cp_draft_shared_kv_debug(message: str, *args) -> None:
if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get():
logger.info("[CP_DRAFT_SHARED_KV] " + message, *args)
_CP_SHARED_KV_BS_GT1_PREFILL_DEBUG_COUNTS = {}
def _cp_shared_kv_bs_gt1_prefill_debug(
key: str,
message: str,
*args,
) -> None:
if not envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get():
return
limit = int(envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG_LIMIT.get())
count = _CP_SHARED_KV_BS_GT1_PREFILL_DEBUG_COUNTS.get(key, 0)
if limit > 0 and count >= limit:
return
_CP_SHARED_KV_BS_GT1_PREFILL_DEBUG_COUNTS[key] = count + 1
logger.info("[CP_SHARED_KV_BS_GT1_DEBUG] event=%s " + message, key, *args)
_CP_SHARED_KV_BS_GT1_PREFILL_TIMING_COUNTS = {}
def _cp_shared_kv_bs_gt1_prefill_timing_start() -> Optional[float]:
if not envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING.get():
return None
return time.perf_counter()
def _cp_shared_kv_bs_gt1_prefill_timing(
key: str,
start_time: Optional[float],
message: str,
*args,
) -> None:
if start_time is None or not envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING.get():
return
elapsed_ms = (time.perf_counter() - start_time) * 1000.0
slow_ms = float(envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING_SLOW_MS.get())
if slow_ms > 0 and elapsed_ms < slow_ms:
return
limit = int(envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING_LIMIT.get())
count = _CP_SHARED_KV_BS_GT1_PREFILL_TIMING_COUNTS.get(key, 0)
if limit > 0 and count >= limit:
return
_CP_SHARED_KV_BS_GT1_PREFILL_TIMING_COUNTS[key] = count + 1
logger.info(
"[CP_SHARED_KV_BS_GT1_TIMING] event=%s elapsed_ms=%.3f " + message,
key,
elapsed_ms,
*args,
)
def _cp_shared_kv_bs_gt1_prefill_marker(
key: str,
message: str,
*args,
) -> None:
if not envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING.get():
return
limit = int(envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING_LIMIT.get())
count = _CP_SHARED_KV_BS_GT1_PREFILL_TIMING_COUNTS.get(key, 0)
if limit > 0 and count >= limit:
return
_CP_SHARED_KV_BS_GT1_PREFILL_TIMING_COUNTS[key] = count + 1
logger.info(
"[CP_SHARED_KV_BS_GT1_TIMING] event=%s elapsed_ms=0.000 " + message,
key,
*args,
)
def _seq_summary(values) -> str:
if values is None:
return "None"
try:
size = len(values)
except TypeError:
return str(values)
if size == 0:
return "size=0"
try:
head = list(values[: min(8, size)])
except TypeError:
head = list(values)[: min(8, size)]
try:
min_val = min(values)
max_val = max(values)
return f"size={size} min={min_val} max={max_val} head={head}"
except (TypeError, ValueError):
return f"size={size} head={head}"
def _pool_summary(pool) -> str:
if pool is None:
return "None"
parts = [pool.__class__.__name__]
for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"):
if hasattr(pool, attr):
parts.append(f"{attr}={getattr(pool, attr)}")
return " ".join(parts)
def _pool_state_type(pool) -> str:
if pool is None or not hasattr(pool, "get_state_buf_infos"):
return "none"
if isinstance(pool, SWAKVPool):
return "swa"
if isinstance(pool, HybridLinearKVPool):
return "mamba"
if isinstance(pool, NSATokenToKVPool):
return "nsa"
return "unknown"
def _state_buf_infos(pool):
state_type = _pool_state_type(pool)
if state_type == "none":
return state_type, [], [], []
state_data_ptrs, state_data_lens, state_item_lens = pool.get_state_buf_infos()
return state_type, state_data_ptrs, state_data_lens, state_item_lens
def _state_layer_ids(pool):
get_state_layer_ids = getattr(pool, "get_state_layer_ids", None)
if get_state_layer_ids is None:
return []
return list(get_state_layer_ids())
def _kv_locs_to_page_indices_cpu(
kv_locs: torch.Tensor,
page_size: int,
):
"""Return int32 page indices without materializing token-level CPU indices."""
if page_size == 1:
page_locs = kv_locs
else:
page_locs = kv_locs[::page_size] // page_size
return page_locs.to(dtype=torch.int32).cpu().numpy()
def release_req_to_metadata_buffer(
req: Req, allocator: ReqToMetadataIdxAllocator
) -> None:
"""
Release the metadata buffer index allocated for a request in prefill disaggregation mode.
This function safely releases the metadata buffer index if it was allocated.
Args:
req: The request object that may have a metadata_buffer_index allocated
allocator: The ReqToMetadataIdxAllocator instance to free the index
"""
if (
hasattr(req, "metadata_buffer_index")
and req.metadata_buffer_index is not None
and req.metadata_buffer_index >= 0
):
allocator.free(req.metadata_buffer_index)
req.metadata_buffer_index = -1
class PrefillBootstrapQueue:
"""
Store the requests in bootstrapping
"""
def __init__(
self,
token_to_kv_pool: KVCache,
draft_token_to_kv_pool: Optional[KVCache],
req_to_metadata_buffer_idx_allocator: ReqToMetadataIdxAllocator,
metadata_buffers: MetadataBuffers,
tp_rank: int,
tp_size: int,
gpu_id: int,
bootstrap_port: int,
gloo_group: ProcessGroup,
max_total_num_tokens: int,
scheduler: Scheduler,
pp_rank: int,
pp_size: int,
transfer_backend: TransferBackend,
):
self.token_to_kv_pool = token_to_kv_pool
self.draft_token_to_kv_pool = draft_token_to_kv_pool
self.is_mla_backend = is_mla_backend(token_to_kv_pool)
self.metadata_buffers = metadata_buffers
self.req_to_metadata_buffer_idx_allocator = req_to_metadata_buffer_idx_allocator
self.tp_rank = tp_rank
self.tp_size = tp_size
self.pp_rank = pp_rank
self.pp_size = pp_size
self.gpu_id = gpu_id
self.bootstrap_port = bootstrap_port
self.queue: List[Req] = []
self.gloo_group = gloo_group
self.max_total_num_tokens = max_total_num_tokens
self.scheduler = scheduler
self.transfer_backend = transfer_backend
self.kv_manager = self._init_kv_manager()
self._maybe_init_per_layer_transfer_manager()
if self.scheduler.tp_worker.is_hybrid_swa:
# FIXME: current SWA allocation allocate full kv cache size in prefill
self.max_total_num_tokens = min(
self.max_total_num_tokens,
self.scheduler.tp_worker.model_runner.swa_max_total_num_tokens,
)
def _maybe_init_per_layer_transfer_manager(self) -> None:
# Lever A: create the per-layer overlapped-transfer manager and register its
# per-layer notifier on the KV pool. The forward fires
# notify_layer_end_for_backup -> layer_backup_notifiers(local_layer_id), which
# calls manager.on_layer_end. Additive/no-op until a request registers a
# transfer context (A3 wiring); gated by SGLANG_CP_SHARED_KV_PER_LAYER_TRANSFER.
self.kv_manager.per_layer_transfer_manager = None
if not envs.SGLANG_CP_SHARED_KV_PER_LAYER_TRANSFER.get():
return
import torch
from sglang.srt.disaggregation.cp_per_layer_transfer import (
PerLayerTransferManager,
)
device = torch.cuda.current_device()
manager = PerLayerTransferManager(
event_factory=torch.cuda.Event,
current_stream=torch.cuda.current_stream,
worker_init=lambda: torch.cuda.set_device(device),
group_size=envs.SGLANG_CP_SHARED_KV_PER_LAYER_GROUP.get(),
)
pool = self.token_to_kv_pool
if hasattr(pool, "register_layer_backup_notifier"):
pool.register_layer_backup_notifier(manager.on_layer_end)
self.kv_manager.per_layer_transfer_manager = manager
logger.info(
"[CP_PER_LAYER_TRANSFER] registered per-layer transfer manager notifier"
)
else:
logger.warning(
"[CP_PER_LAYER_TRANSFER] kv pool lacks register_layer_backup_notifier; "
"per-layer overlap disabled"
)
def _init_kv_manager(self) -> CommonKVManager:
kv_args_class = get_kv_class(self.transfer_backend, KVClassType.KVARGS)
kv_args = kv_args_class()
kv_args.engine_rank = self.tp_rank
kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.dp_rank
kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer
kv_data_ptrs, kv_data_lens, kv_item_lens = (
self.token_to_kv_pool.get_contiguous_buf_infos()
)
target_kv_buffer_count = len(kv_data_ptrs)
draft_kv_data_lens = []
draft_kv_item_lens = []
draft_kv_buffer_count = 0
if self.draft_token_to_kv_pool is not None:
# We should also transfer draft model kv cache. The indices are
# always shared with a target model.
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
self.draft_token_to_kv_pool.get_contiguous_buf_infos()
)
draft_kv_buffer_count = len(draft_kv_data_ptrs)
kv_data_ptrs += draft_kv_data_ptrs
kv_data_lens += draft_kv_data_lens
kv_item_lens += draft_kv_item_lens
kv_args.draft_kv_buffer_start = target_kv_buffer_count
kv_args.draft_kv_buffer_count = draft_kv_buffer_count
_cp_draft_shared_kv_debug(
"prefill_kv_manager cp_rank=%s target_pool=(%s) draft_pool=(%s) "
"target_bufs=%s draft_bufs=%s total_bufs=%s target_lens=%s "
"draft_lens=%s target_item_lens=%s draft_item_lens=%s",
self.tp_rank,
_pool_summary(self.token_to_kv_pool),
_pool_summary(self.draft_token_to_kv_pool),
target_kv_buffer_count,
draft_kv_buffer_count,
len(kv_data_ptrs),
_seq_summary(kv_data_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_data_lens),
_seq_summary(kv_item_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_item_lens),
)
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
if not self.is_mla_backend:
kv_args.kv_head_num = self.token_to_kv_pool.head_num
kv_args.total_kv_head_num = (
self.scheduler.model_config.get_total_num_kv_heads()
)
kv_args.page_size = self.token_to_kv_pool.page_size
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
self.metadata_buffers.get_buf_infos()
)
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
kv_args.gpu_id = self.scheduler.gpu_id
if hasattr(self.token_to_kv_pool, "get_state_buf_infos"):
state_data_ptrs, state_data_lens, state_item_lens = (
self.token_to_kv_pool.get_state_buf_infos()
)
kv_args.state_data_ptrs = state_data_ptrs
kv_args.state_data_lens = state_data_lens
kv_args.state_item_lens = state_item_lens
kv_args.state_layer_ids = _state_layer_ids(self.token_to_kv_pool)
if isinstance(self.token_to_kv_pool, SWAKVPool):
kv_args.state_type = "swa"
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
kv_args.state_type = "mamba"
# Get state dimension info for cross-TP slice transfer
if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"):
kv_args.state_dim_per_tensor = (
self.token_to_kv_pool.get_state_dim_per_tensor()
)
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
kv_args.state_type = "nsa"
else:
kv_args.state_type = "none"
else:
kv_args.state_data_ptrs = []
kv_args.state_data_lens = []
kv_args.state_item_lens = []
kv_args.state_layer_ids = []
kv_args.state_type = "none"
draft_state_type = "none"
draft_state_data_ptrs = []
draft_state_data_lens = []
draft_state_item_lens = []
if self.draft_token_to_kv_pool is not None:
(
draft_state_type,
draft_state_data_ptrs,
draft_state_data_lens,
draft_state_item_lens,
) = _state_buf_infos(self.draft_token_to_kv_pool)
draft_state_registered = append_cp_draft_state_buffers(
kv_args,
draft_state_type,
draft_state_data_ptrs,
draft_state_data_lens,
draft_state_item_lens,
role="prefill",
cp_rank=self.tp_rank,
)
if draft_state_data_ptrs and not draft_state_registered:
_cp_draft_shared_kv_debug(
"prefill_draft_state_skipped cp_rank=%s target_state_type=%s "
"draft_state_type=%s draft_state_bufs=%s "
"reason=unsupported_state_type",
self.tp_rank,
kv_args.state_type,
draft_state_type,
len(draft_state_data_ptrs),
)
if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get():
_cp_draft_shared_kv_debug(
"prefill_state_manager cp_rank=%s target_state_type=%s "
"draft_state_type=%s draft_state_bufs=%s draft_state_lens=%s "
"draft_state_item_lens=%s draft_state_start=%s registered_state_bufs=%s "
"registered_state_lens=%s registered_state_item_lens=%s",
self.tp_rank,
kv_args.state_type,
kv_args.draft_state_type,
kv_args.draft_state_buffer_count,
_seq_summary(draft_state_data_lens),
_seq_summary(draft_state_item_lens),
kv_args.draft_state_buffer_start,
len(kv_args.state_data_ptrs),
_seq_summary(kv_args.state_data_lens),
_seq_summary(kv_args.state_item_lens),
)
if kv_args.draft_state_buffer_count > 0:
_cp_draft_shared_kv_debug(
"prefill_draft_state_registered cp_rank=%s "
"draft_state_type=%s draft_state_bufs=%s draft_state_start=%s "
"registered_state_type=%s registered_state_bufs=%s",
self.tp_rank,
kv_args.draft_state_type,
kv_args.draft_state_buffer_count,
kv_args.draft_state_buffer_start,
kv_args.state_type,
len(kv_args.state_data_ptrs),
)
if envs.SGLANG_EAGLE_ACCEPT_DEBUG.get() and self.scheduler.spec_algorithm.is_eagle():
logger.info(
"[EAGLE_ACCEPT_DEBUG] prefill_kv_manager cp_rank=%s "
"target_kv_bufs=%s draft_kv_bufs=%s total_kv_bufs=%s "
"target_state_type=%s registered_state_bufs=%s "
"draft_state_type=%s draft_state_bufs=%s",
self.tp_rank,
target_kv_buffer_count,
draft_kv_buffer_count,
len(kv_args.kv_data_ptrs),
kv_args.state_type,
len(kv_args.state_data_ptrs),
kv_args.draft_state_type,
kv_args.draft_state_buffer_count,
)
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
DisaggregationMode.PREFILL,
self.scheduler.server_args,
self.is_mla_backend,
)
return kv_manager
def add(self, req: Req, num_kv_heads: int) -> None:
if self._check_if_req_exceed_kv_capacity(req):
return
backend = (
TransferBackend.FAKE
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST
else self.transfer_backend
)
kv_sender_class = get_kv_class(backend, KVClassType.SENDER)
dest_tp_ranks = [self.tp_rank]
req.disagg_kv_sender = kv_sender_class(
mgr=self.kv_manager,
bootstrap_addr=f"{req.bootstrap_host}:{self.bootstrap_port}",
bootstrap_room=req.bootstrap_room,
dest_tp_ranks=dest_tp_ranks,
pp_rank=self.pp_rank,
)
self._process_req(req)
self.queue.append(req)
def extend(self, reqs: List[Req], num_kv_heads: int) -> None:
for req in reqs:
self.add(req, num_kv_heads)
def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool:
if len(req.origin_input_ids) > self.max_total_num_tokens:
message = f"Request {req.rid} exceeds the maximum number of tokens: {len(req.origin_input_ids)} > {self.max_total_num_tokens}"
logger.error(message)
req.time_stats.trace_ctx.abort(abort_info={"reason": message})
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
self.scheduler.stream_output([req], req.return_logprob)
return True
return False
def _process_req(self, req: Req) -> None:
"""
Set max_new_tokens = 1, so PrefillAdder memory estimation is accurate
"""
req.sampling_params.max_new_tokens = 1
def pop_bootstrapped(
self,
return_failed_reqs: bool = False,
rids_to_check: Optional[List[str]] = None,
) -> List[Req]:
"""
pop the reqs which has finished bootstrapping
return_failed_reqs: For PP, on rank 0, also return the failed reqs to notify the next rank
rids_to_check: For PP, on rank > 0, check the rids from the previous rank has consensus with the current rank.
"""
bootstrapped_reqs = []
failed_reqs = []
remaining_queue = []
if len(self.queue) == 0:
if return_failed_reqs is False:
return []
else:
return [], []
poll_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"bootstrap_poll_start",
"queue=%s return_failed=%s rids_to_check=%s",
len(self.queue),
return_failed_reqs,
len(rids_to_check) if rids_to_check is not None else None,
)
polls = poll_and_all_reduce_attn_cp_tp_group(
[req.disagg_kv_sender for req in self.queue],
self.scheduler.attn_cp_cpu_group,
self.scheduler.attn_tp_cpu_group,
debug_label="bootstrap",
debug_ids=[req.rid for req in self.queue],
)
_cp_shared_kv_bs_gt1_prefill_timing(
"bootstrap_poll_done",
poll_start,
"queue=%s polls_head=%s",
len(self.queue),
polls[:8],
)
for i, (req, poll) in enumerate(zip(self.queue, polls)):
if rids_to_check is not None:
# if req not in reqs_info_to_check, skip
if req.rid not in rids_to_check:
remaining_queue.append(req)
continue
if poll == KVPoll.Bootstrapping:
remaining_queue.append(req)
continue
elif poll == KVPoll.Failed:
error_message = f"Prefill bootstrap failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}"
try:
req.disagg_kv_sender.failure_exception()
except Exception as e:
error_message += f" with exception {e}"
logger.error(error_message)
req.time_stats.trace_ctx.abort(abort_info={"reason": error_message})
prepare_abort(
req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR
)
self.scheduler.stream_output([req], req.return_logprob)
failed_reqs.append(req)
if self.scheduler.enable_metrics:
self.scheduler.metrics_collector.increment_bootstrap_failed_reqs()
if self.scheduler.enable_hicache_storage:
# to release prefetch events associated with the request
self.scheduler.tree_cache.release_aborted_request(req.rid)
continue
# KV.WaitingForInput - init here
req.time_stats.set_bootstrap_done_time()
num_kv_indices = len(req.origin_input_ids)
if self.req_to_metadata_buffer_idx_allocator.available_size() == 0:
remaining_queue.append(req)
remaining_queue.extend(self.queue[i + 1 :])
break
req.metadata_buffer_index = (
self.req_to_metadata_buffer_idx_allocator.alloc()
)
assert req.metadata_buffer_index is not None
num_pages = kv_to_page_num(num_kv_indices, self.token_to_kv_pool.page_size)
req.disagg_kv_sender.init(num_pages, req.metadata_buffer_index)
bootstrapped_reqs.append(req)
req.time_stats.set_wait_queue_entry_time()
self.queue = remaining_queue
if return_failed_reqs is False:
return bootstrapped_reqs
else:
return bootstrapped_reqs, failed_reqs
class SchedulerDisaggregationPrefillMixin:
"""
Mixin for Scheduler to handle disaggregation prefill
"""
def get_next_disagg_prefill_batch_to_run(
self: Scheduler,
) -> Optional[ScheduleBatch]:
# HACK (byronhsu): reset the batch_is_full flag because we never enter update_running_batch which resets it
# Otherwise, it hangs under high concurrency
self.running_batch.batch_is_full = False
self.process_prefill_chunk()
batch = self.get_new_batch_prefill()
batch = self.maybe_prepare_mlp_sync_batch(batch)
if batch:
set_schedule_time_batch(batch)
return batch
def _register_per_layer_transfers(self: Scheduler, batch) -> None:
"""Lever A: before run_batch, register a per-layer transfer context for each
eligible request so the per-layer notifier overlaps its main-KV transfer with
the forward. Scoped (first impl) to single-forward, no-cached-prefix requests:
the per-layer notifier transfers FORWARD-written pages, so chunked or cached-
prefix requests (whose prefix is HiCache-loaded, not forward-written) are left
on the monolithic post-forward path. No-op unless the flag is on."""
kv_manager = self.disagg_prefill_bootstrap_queue.kv_manager
mgr = getattr(kv_manager, "per_layer_transfer_manager", None)
if mgr is None:
return
page_size = self.token_to_kv_pool_allocator.page_size
for req in batch.reqs:
try:
if getattr(req, "disagg_kv_sender", None) is None:
continue
room = getattr(req, "bootstrap_room", None)
if room is None:
continue
start_idx = getattr(req, "start_send_idx", 0)
if mgr.has_chunk(room, start_idx):
continue # this chunk already registered (bs>1 batch-forming re-iterates)
# Register the EXACT range this forward's send_kv_chunk will transmit:
# req_to_token[start_send_idx : end_idx], page-floored for a NON-last chunk
# (mirrors send_kv_chunk). UNIFIED: non-chunked = one full range [0:end];
# chunked = one range per chunk. The per-layer backup hook fires after each
# layer's full processing (HiCache-loaded prefix KV + forward-written new KV
# both final), so transferring layer L's pages then is correct and overlaps
# the transfer with the forward. Multiple chunks of a room get separate
# contexts (keyed by start_send_idx), finished FIFO.
end_idx = min(len(req.fill_ids), len(req.origin_input_ids))
if getattr(req, "is_chunked", 0) > 0: # not the last chunk
end_idx -= end_idx % page_size
if end_idx <= start_idx:
continue
page_indices = _kv_locs_to_page_indices_cpu(
self.req_to_token_pool.req_to_token[
req.req_pool_idx, start_idx:end_idx
],
page_size,
)
kv_manager.register_per_layer_transfer(
room, page_indices, chunk_key=start_idx
)
except Exception as e:
# Never let lever-A setup crash the scheduler; fall back to the
# monolithic post-forward transfer for this request.
logger.warning(
"[CP_PER_LAYER_TRANSFER] register skipped for room %s: %r",
getattr(req, "bootstrap_room", None),
e,
)
@torch.no_grad()
def event_loop_normal_disagg_prefill(self: Scheduler) -> None:
"""A normal scheduler loop for prefill worker in disaggregation mode."""
while True:
# Receive requests
recv_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"event_loop_recv_start",
"waiting=%s inflight=%s bootstrap=%s",
len(self.waiting_queue),
len(self.disagg_prefill_inflight_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
recv_reqs = self.recv_requests()
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_recv_done",
recv_start,
"recv=%s waiting=%s inflight=%s bootstrap=%s",
len(recv_reqs),
len(self.waiting_queue),
len(self.disagg_prefill_inflight_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
self.process_input_requests(recv_reqs)
pop_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
bootstrapped_reqs = self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_pop_bootstrapped_done",
pop_start,
"bootstrapped=%s waiting_before=%s inflight=%s bootstrap_remaining=%s",
len(bootstrapped_reqs),
len(self.waiting_queue),
len(self.disagg_prefill_inflight_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
self.waiting_queue.extend(bootstrapped_reqs)
# Get the next batch to run
batch_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"event_loop_get_batch_start",
"waiting=%s inflight=%s bootstrap=%s",
len(self.waiting_queue),
len(self.disagg_prefill_inflight_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
batch = self.get_next_disagg_prefill_batch_to_run()
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_get_batch_done",
batch_start,
"has_batch=%s batch_size=%s waiting=%s inflight=%s bootstrap=%s",
batch is not None,
len(batch.reqs) if batch is not None else 0,
len(self.waiting_queue),
len(self.disagg_prefill_inflight_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
self.cur_batch = batch
# Launch the current batch
if batch:
run_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"event_loop_run_batch_start",
"batch_size=%s extend_lens=%s prefix_lens=%s inflight=%s",
len(batch.reqs),
_seq_summary(getattr(batch, "extend_lens", None)),
_seq_summary(getattr(batch, "prefix_lens", None)),
len(self.disagg_prefill_inflight_queue),
)
self._register_per_layer_transfers(batch)
result = self.run_batch(batch)
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_run_batch_done",
run_start,
"batch_size=%s inflight=%s",
len(batch.reqs),
len(self.disagg_prefill_inflight_queue),
)
result_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"event_loop_process_result_start",
"batch_size=%s inflight_before=%s",
len(batch.reqs),
len(self.disagg_prefill_inflight_queue),
)
self.process_batch_result(batch, result)
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_process_result_done",
result_start,
"batch_size=%s inflight_after=%s",
len(batch.reqs),
len(self.disagg_prefill_inflight_queue),
)
else:
self.self_check_during_idle()
inflight_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"event_loop_process_inflight_start",
"inflight=%s waiting=%s bootstrap=%s",
len(self.disagg_prefill_inflight_queue),
len(self.waiting_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
self.process_disagg_prefill_inflight_queue()
_cp_shared_kv_bs_gt1_prefill_timing(
"event_loop_process_inflight_done",
inflight_start,
"inflight=%s waiting=%s bootstrap=%s",
len(self.disagg_prefill_inflight_queue),
len(self.waiting_queue),
len(self.disagg_prefill_bootstrap_queue.queue),
)
# Update last_batch
self.last_batch = batch
@torch.no_grad()
def event_loop_overlap_disagg_prefill(self: Scheduler) -> None:
self.result_queue = deque()
while True:
# Receive requests
recv_reqs = self.recv_requests()
self.process_input_requests(recv_reqs)
self.waiting_queue.extend(
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
)
# Get the next batch to run
batch = self.get_next_disagg_prefill_batch_to_run()
self.cur_batch = batch
# Launch the current batch
if batch:
self._register_per_layer_transfers(batch)
batch_result = self.run_batch(batch)
self.result_queue.append((batch.copy(), batch_result))
else:
batch_result = None
# Process the last batch
if self.last_batch:
tmp_batch, tmp_result = self.result_queue.popleft()
self.process_batch_result(tmp_batch, tmp_result)
elif batch is None:
# When the server is idle, do self-check and re-init some states
self.self_check_during_idle()
self.process_disagg_prefill_inflight_queue()
# Run sample of the current batch
# It depends on the result of the last batch (e.g., grammar), so we run it after the last batch is processed.
self.launch_batch_sample_if_needed(batch_result)
# Update last_batch
self.last_batch = batch
def process_batch_result_disagg_prefill(
self: Scheduler,
batch: ScheduleBatch,
result: GenerationBatchResult,
) -> None:
"""
Transfer kv for prefill completed requests and add it into disagg_prefill_inflight_queue
Adapted from process_batch_result_prefill
"""
(
logits_output,
next_token_ids,
extend_input_len_per_req,
extend_logprob_start_len_per_req,
copy_done,
) = (
result.logits_output,
result.next_token_ids,
result.extend_input_len_per_req,
result.extend_logprob_start_len_per_req,
result.copy_done,
)
if copy_done is not None:
copy_done.synchronize()
logprob_pt = 0
# Transfer kv for prefill completed requests and add it into disagg_prefill_inflight_queue
next_token_ids = result.next_token_ids.tolist()
if envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get():
spec_info = getattr(batch, "spec_info", None)
spec_hidden = getattr(spec_info, "hidden_states", None)
spec_topk = getattr(spec_info, "topk_index", None)
_cp_shared_kv_bs_gt1_prefill_debug(
"prefill_result_handoff",
"bs=%s rids=%s extend_lens=%s prefix_lens=%s next_token_ids=%s "
"has_spec=%s hidden_shape=%s topk_shape=%s out_cache_tokens=%s "
"inflight_before=%s",
len(batch.reqs),
[req.rid for req in batch.reqs[:8]],
list(getattr(batch, "extend_lens", []) or []),
list(getattr(batch, "prefix_lens", []) or []),
next_token_ids[:8],
spec_info is not None,
tuple(spec_hidden.shape) if spec_hidden is not None else None,
tuple(spec_topk.shape) if spec_topk is not None else None,
int(batch.out_cache_loc.numel())
if getattr(batch, "out_cache_loc", None) is not None
else None,
len(self.disagg_prefill_inflight_queue),
)
if batch.return_logprob:
if logits_output.next_token_logprobs is not None:
logits_output.next_token_logprobs = (
logits_output.next_token_logprobs.tolist()
)
if logits_output.input_token_logprobs is not None:
logits_output.input_token_logprobs = tuple(
logits_output.input_token_logprobs.tolist()
)
# Batch the D2H copies of the EAGLE spec outputs: one gathered copy per
# tensor for all finishing requests instead of one synchronous device
# sync per request here (hidden_states) and per request inside
# MetadataBuffers.set_buf (topk_p/topk_index, which were stored as GPU
# row views and copied into the CPU metadata buffers one by one).
spec_cpu_row = None
if self.spec_algorithm.is_eagle() and batch.spec_info is not None:
finished_idx = [i for i, r in enumerate(batch.reqs) if r.is_chunked <= 0]
if finished_idx:
spec_cpu_row = {i: k for k, i in enumerate(finished_idx)}
if len(finished_idx) == len(batch.reqs):
gather = lambda t: t
else:
gather_idx = torch.tensor(
finished_idx,
dtype=torch.int64,
device=batch.spec_info.hidden_states.device,
)
gather = lambda t: t.index_select(0, gather_idx)
# .to(copy=True) detaches from spec_info even when the batch is
# already on CPU (set_buf reads these synchronously below).
spec_topk_p_cpu = gather(batch.spec_info.topk_p).to("cpu", copy=True)
spec_topk_index_cpu = gather(batch.spec_info.topk_index).to(
"cpu", copy=True
)
spec_hidden_cpu = gather(batch.spec_info.hidden_states).to(
"cpu", copy=True
)
for i, (req, next_token_id) in enumerate(
zip(batch.reqs, next_token_ids, strict=True)
):
if req.is_chunked <= 0:
req.time_stats.set_prefill_finished_time()
# There is no output_ids for prefill
req.output_ids.append(next_token_id)
self.tree_cache.cache_unfinished_req(req) # update the tree and lock
self.disagg_prefill_inflight_queue.append(req)
# Residency clock for the inflight liveness bound (every rank
# stamps at the same logical point in the same tick, so the
# timeout decision below stays effectively rank-uniform; the
# MIN-reduced poll consensus absorbs any microsecond skew).
req.disagg_inflight_enter_time = time.perf_counter()
if spec_cpu_row is not None:
k = spec_cpu_row[i]
req.output_topk_p = spec_topk_p_cpu[k]
req.output_topk_index = spec_topk_index_cpu[k]
req.hidden_states_tensor = spec_hidden_cpu[k]
else:
req.hidden_states_tensor = None
if req.return_logprob:
assert extend_logprob_start_len_per_req is not None
assert extend_input_len_per_req is not None
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
extend_input_len = extend_input_len_per_req[i]
num_input_logprobs = extend_input_len - extend_logprob_start_len
self.add_logprob_return_values(
i,
req,
logprob_pt,
next_token_ids,
num_input_logprobs,
logits_output,
)
logprob_pt += num_input_logprobs
self.send_kv_chunk(req, last_chunk=True)
req.time_stats.set_prefill_transfer_queue_entry_time()
if req.grammar is not None:
# FIXME: this try-except block is for handling unexpected xgrammar issue.
try:
req.grammar.accept_token(next_token_id)
except ValueError as e:
# Grammar accept_token can raise ValueError if the token is not in the grammar.
# This can happen if the grammar is not set correctly or the token is invalid.
error_message = f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}"
release_kv_cache(req, self.tree_cache)
prepare_abort(
req,
error_message,
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
)
req.grammar.finished = req.finished()
else:
# being chunked reqs' prefill is not finished
req.is_chunked -= 1
if req.return_logprob:
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
extend_input_len = extend_input_len_per_req[i]
if extend_logprob_start_len < extend_input_len:
# Update input logprobs.
num_input_logprobs = extend_input_len - extend_logprob_start_len
self.add_input_logprob_return_values(
i,
req,
logits_output,
logprob_pt,
num_input_logprobs,
last_prefill_chunk=False,
)
logprob_pt += num_input_logprobs
if self.enable_overlap:
self.send_kv_chunk(req, last_chunk=False, end_idx=req.tmp_end_idx)
req.time_stats.set_last_chunked_prefill_finish_time()
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
self.report_prefill_stats(
prefill_stats=batch.prefill_stats,
can_run_cuda_graph=can_run_cuda_graph,
dp_cooperation_info=batch.dp_cooperation_info,
)
def process_disagg_prefill_inflight_queue(
self: Scheduler, rids_to_check: Optional[List[str]] = None
) -> List[Req]:
"""
Poll the requests in the middle of transfer. If done, return the request.
rids_to_check: For PP, on rank > 0, check the rids from the previous rank has consensus with the current rank.
"""
if len(self.disagg_prefill_inflight_queue) == 0:
return []
done_reqs = []
poll_start = _cp_shared_kv_bs_gt1_prefill_timing_start()
_cp_shared_kv_bs_gt1_prefill_marker(
"inflight_poll_start",
"inflight=%s rids_head=%s rids_to_check=%s",
len(self.disagg_prefill_inflight_queue),
[req.rid for req in self.disagg_prefill_inflight_queue[:8]],
len(rids_to_check) if rids_to_check is not None else None,
)
polls = poll_and_all_reduce_attn_cp_tp_group(
[req.disagg_kv_sender for req in self.disagg_prefill_inflight_queue],
self.attn_cp_cpu_group,
self.attn_tp_cpu_group,
debug_label="inflight",
debug_ids=[req.rid for req in self.disagg_prefill_inflight_queue],
)
_cp_shared_kv_bs_gt1_prefill_timing(
"inflight_poll_done",
poll_start,
"inflight=%s polls_head=%s",
len(self.disagg_prefill_inflight_queue),
polls[:8],
)
undone_reqs: List[Req] = []
# Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue
for req, poll in zip(self.disagg_prefill_inflight_queue, polls):
if rids_to_check is not None:
if req.rid not in rids_to_check:
undone_reqs.append(req)
continue
# In PP mode, the previous rank may have reached a terminal
# state (Success/Failed) while this rank's local poll is still
# in a transient state due to clock skew or propagation delay.
# Treat non-terminal states as undone instead of crashing.
if poll not in (
KVPoll.Success,
KVPoll.Failed,
):
logger.warning(
f"PP rank {self.pp_rank}: unexpected poll state {poll} for rid {req.rid} "
f"from consensus; treating as undone"
)
undone_reqs.append(req)
continue
if poll in [KVPoll.WaitingForInput, KVPoll.Transferring]:
# Liveness bound: an inflight entry whose transfer never
# concludes (decode peer torn down with its abort
# notification lost, last chunk silently dropped, decode
# crash) would otherwise be re-queued FOREVER — it pins the
# request's KV pages and rides every poll collective. On
# expiry, conclude the sender locally and keep the entry
# undone: Failed(0) wins the CP/TP MIN-reduce, so every rank
# reaps it together through the Failed branch on the NEXT
# poll — one removal path, rank-uniform by construction.
enter_time = getattr(req, "disagg_inflight_enter_time", None)
if (
enter_time is not None
and time.perf_counter() - enter_time
>= envs.SGLANG_DISAGGREGATION_INFLIGHT_TIMEOUT.get()
and hasattr(req.disagg_kv_sender, "abort")
):
logger.warning(
"Prefill inflight transfer timed out; failing the "
"request. rid=%s bootstrap_room=%s poll=%s "
"elapsed=%.1fs (raise "
"SGLANG_DISAGGREGATION_INFLIGHT_TIMEOUT if transfers "
"can legitimately take longer)",
req.rid,
req.bootstrap_room,
poll,
time.perf_counter() - enter_time,
)
req.disagg_kv_sender.abort()
undone_reqs.append(req)
elif poll == KVPoll.Success: # transfer done
release_kv_cache(req, self.tree_cache) # unlock the tree
req.finished_reason = FINISH_LENGTH(length=0)
# FIXME: clean up req's data in transfer engine
if hasattr(req.disagg_kv_sender, "clear"):
req.disagg_kv_sender.clear()
done_reqs.append(req)
req.time_stats.set_prefill_kv_transfer_finish_time()
elif poll == KVPoll.Failed:
error_message = f"Prefill transfer failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}"
try:
req.disagg_kv_sender.failure_exception()
except Exception as e:
error_message += f" with exception {e}"
logger.warning(error_message)
req.time_stats.trace_ctx.abort(abort_info={"reason": error_message})
release_kv_cache(req, self.tree_cache) # unlock the tree
prepare_abort(
req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR
)
done_reqs.append(req)
if self.enable_metrics:
self.metrics_collector.increment_transfer_failed_reqs()
else:
logger.warning(
f"Unexpected polling state {poll} for rid {req.rid} in inflight queue; "
f"treating as undone"
)
undone_reqs.append(req)
for req in done_reqs:
req.time_stats.set_completion_time()
for req in done_reqs:
if isinstance(req.finished_reason, FINISH_ABORT):
continue
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST:
continue
kv_mgr = getattr(req.disagg_kv_sender, "kv_mgr", None)
if kv_mgr and getattr(kv_mgr, "is_dummy_cp_rank", False):
# Dummy CP ranks transfer nothing; skip so they don't pollute the metric.
continue
metrics = req.time_stats.compute_and_observe_kv_transfer_metrics(
req.disagg_kv_sender.get_transfer_metric()
)
if metrics:
# Update last-value for REST API
if "latency_ms" in metrics:
self.kv_transfer_latency_ms = metrics["latency_ms"]
if "speed_gb_s" in metrics:
self.kv_transfer_speed_gb_s = metrics["speed_gb_s"]
# Stream requests which have finished transfer
self.stream_output(
done_reqs,
any(req.return_logprob for req in done_reqs),
None,
)
for req in done_reqs:
req: Req
release_req_to_metadata_buffer(
req, self.req_to_metadata_buffer_idx_allocator
)
self.disagg_prefill_inflight_queue = undone_reqs
return done_reqs
def get_transferred_rids(self: Scheduler) -> List[str]:
"""
Used by PP, get the transferred rids but **do not pop**
"""
polls = poll_and_all_reduce_attn_cp_tp_group(
[req.disagg_kv_sender for req in self.disagg_prefill_inflight_queue],
self.attn_cp_cpu_group,
self.attn_tp_cpu_group,
debug_label="get_transferred_rids",
debug_ids=[req.rid for req in self.disagg_prefill_inflight_queue],
)
transferred_rids: List[str] = []
for req, poll in zip(self.disagg_prefill_inflight_queue, polls):
if poll == KVPoll.Success or poll == KVPoll.Failed:
transferred_rids.append(req.rid)
return transferred_rids
def process_prefill_chunk(self: Scheduler) -> None:
chunked_req_to_exclude = set()
if self.chunked_req:
chunked_req_to_exclude.add(self.chunked_req)
self.tree_cache.cache_unfinished_req(self.chunked_req, chunked=True)
if self.enable_overlap:
# Delay KV transfer to process_batch_result_disagg_prefill when overlap is enabled to ensure results are resolved
self.chunked_req.tmp_end_idx = min(
len(self.chunked_req.fill_ids),
len(self.chunked_req.origin_input_ids),
)
else:
self.send_kv_chunk(self.chunked_req)
self.running_batch.batch_is_full = False
if self.last_batch and self.last_batch.forward_mode.is_extend():
if self.last_batch.chunked_req:
# In the context pipeline parallelism, after the last chunk, the current microbatch still track outdated chunked_req.
# We need to discard it.
chunked_req_to_exclude.add(self.last_batch.chunked_req)
last_bs = self.last_batch.batch_size()
self.last_batch.filter_batch(
chunked_req_to_exclude=list(chunked_req_to_exclude)
)
if self.last_batch.batch_size() < last_bs:
self.running_batch.batch_is_full = False
def send_kv_chunk(
self: Scheduler,
req: Req,
last_chunk: bool = False,
end_idx: Optional[int] = None,
) -> None:
"""
Send a prefilled chunk to the decode server
"""
page_size = self.token_to_kv_pool_allocator.page_size
start_idx = req.start_send_idx
end_idx = (
end_idx
if end_idx is not None
else min(len(req.fill_ids), len(req.origin_input_ids))
)
if not last_chunk:
# if not the last chunk and the last page is partial, delay the last partial page to the next send
end_idx = end_idx - end_idx % page_size
page_indices = _kv_locs_to_page_indices_cpu(
self.req_to_token_pool.req_to_token[req.req_pool_idx, start_idx:end_idx],
page_size,
)
req.start_send_idx = end_idx
state_indices = None
if last_chunk:
self.disagg_metadata_buffers.set_buf(req)
# Prepare extra pool indices for hybrid models
if isinstance(
self.token_to_kv_pool_allocator.get_kvcache(), HybridLinearKVPool
):
# Mamba hybrid model: send single mamba state index
state_indices = [
self.req_to_token_pool.req_index_to_mamba_index_mapping[
req.req_pool_idx
]
.cpu()
.numpy()
]
elif isinstance(self.token_to_kv_pool_allocator.get_kvcache(), SWAKVPool):
# SWA hybrid model: send last window KV indices
seq_len = len(req.fill_ids)
window_size = self.sliding_window_size
window_start = max(0, seq_len - window_size)
window_start = (window_start // page_size) * page_size
window_kv_indices_full = self.req_to_token_pool.req_to_token[
req.req_pool_idx, window_start:seq_len
]
# Translate to SWA pool indices
window_kv_indices_swa = (
self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
window_kv_indices_full
)
)
state_indices = _kv_locs_to_page_indices_cpu(
window_kv_indices_swa,
page_size,
)
elif isinstance(
self.token_to_kv_pool_allocator.get_kvcache(), NSATokenToKVPool
):
seq_len = len(req.fill_ids)
state_indices = _kv_locs_to_page_indices_cpu(
self.req_to_token_pool.req_to_token[req.req_pool_idx, :seq_len],
page_size,
)
if len(page_indices) == 0 and not last_chunk:
logger.info(
f"Skip sending non-final kv chunk for request {req.rid=} {req.bootstrap_room=} because page_indices is empty"
)
return
if len(page_indices) == 0:
logger.warning(
"[CP_SHARED_KV_TRANSFER][final_empty_chunk] "
"sending final aux metadata without new KV pages: "
f"{req.rid=} {req.bootstrap_room=} {start_idx=} {end_idx=} "
f"origin_len={len(req.origin_input_ids)} fill_len={len(req.fill_ids)}"
)
prefill_queue = getattr(self, "disagg_prefill_bootstrap_queue", None)
has_draft_pool = (
getattr(prefill_queue, "draft_token_to_kv_pool", None) is not None
)
prefix_len = len(getattr(req, "prefix_indices", ()))
host_hit_length = int(getattr(req, "host_hit_length", 0) or 0)
draft_prefix_overlap = max(0, min(end_idx, prefix_len) - start_idx)
_cp_draft_shared_kv_debug(
"prefill_send_kv_chunk rid=%s room=%s start_idx=%s end_idx=%s "
"last_chunk=%s page_size=%s pages=%s state_pages=%s "
"has_draft_pool=%s prefix_len=%s host_hit_length=%s "
"cache_protected_len=%s extend_input_len=%s fill_len=%s "
"origin_input_len=%s already_computed=%s draft_prefix_overlap=%s",
req.rid,
req.bootstrap_room,
start_idx,
end_idx,
last_chunk,
page_size,
_seq_summary(page_indices),
_seq_summary(state_indices),
has_draft_pool,
prefix_len,
host_hit_length,
getattr(req, "cache_protected_len", None),
getattr(req, "extend_input_len", None),
len(req.fill_ids),
len(req.origin_input_ids),
getattr(req, "already_computed", None),
draft_prefix_overlap,
)
if envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get():
_cp_shared_kv_bs_gt1_prefill_debug(
"send_kv_chunk",
"rid=%s room=%s start_idx=%s end_idx=%s last_chunk=%s "
"page_size=%s pages=%s state_pages=%s prefix_len=%s "
"host_hit_length=%s extend_input_len=%s fill_len=%s "
"origin_input_len=%s has_draft_pool=%s",
req.rid,
req.bootstrap_room,
start_idx,
end_idx,
last_chunk,
page_size,
_seq_summary(page_indices),
_seq_summary(state_indices),
prefix_len,
host_hit_length,
getattr(req, "extend_input_len", None),
len(req.fill_ids),
len(req.origin_input_ids),
has_draft_pool,
)
if has_draft_pool and draft_prefix_overlap > 0:
_cp_draft_shared_kv_debug(
"prefill_send_cachehit_draft_prefix rid=%s room=%s "
"draft_prefix_overlap=%s prefix_len=%s host_hit_length=%s "
"start_idx=%s end_idx=%s note=transfer_reads_draft_pool_for_cached_prefix",
req.rid,
req.bootstrap_room,
draft_prefix_overlap,
prefix_len,
host_hit_length,
start_idx,
end_idx,
)
req.disagg_kv_sender.send(page_indices, state_indices)