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>
1407 lines
57 KiB
Python
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)
|