[PD] Optimistic prefill (#26780)

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
cctry
2026-06-02 01:16:14 -07:00
committed by GitHub
co-authored by gemini-code-assist[bot]
parent 5ae8d286d2
commit b55570d38e
11 changed files with 571 additions and 92 deletions
+230 -58
View File
@@ -19,7 +19,9 @@ Life cycle of a request in the prefill server
from __future__ import annotations from __future__ import annotations
import hashlib
import logging import logging
from array import array
from collections import deque from collections import deque
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional from typing import TYPE_CHECKING, List, Optional
@@ -37,6 +39,7 @@ from sglang.srt.disaggregation.utils import (
ReqToMetadataIdxAllocator, ReqToMetadataIdxAllocator,
TransferBackend, TransferBackend,
get_kv_class, get_kv_class,
is_aborted,
is_mla_backend, is_mla_backend,
poll_and_all_reduce_attn_cp_tp_group, poll_and_all_reduce_attn_cp_tp_group,
prepare_abort, prepare_abort,
@@ -67,7 +70,17 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def release_req_to_metadata_buffer( def should_force_retry(req: Req) -> bool:
"""Test hook to force a request into optimistic prefill retry."""
retry_prob = envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.get()
if retry_prob <= 0 or req.time_stats.prefill_retry_count > 0 or req.is_retracted:
return False
digest = hashlib.sha256(str(req.rid).encode()).digest()
return int.from_bytes(digest[:8], "big") < retry_prob * 2**64
def maybe_release_metadata_buffer(
req: Req, allocator: ReqToMetadataIdxAllocator req: Req, allocator: ReqToMetadataIdxAllocator
) -> None: ) -> None:
""" """
@@ -79,11 +92,7 @@ def release_req_to_metadata_buffer(
req: The request object that may have a metadata_buffer_index allocated req: The request object that may have a metadata_buffer_index allocated
allocator: The ReqToMetadataIdxAllocator instance to free the index allocator: The ReqToMetadataIdxAllocator instance to free the index
""" """
if ( if req.metadata_buffer_index >= 0:
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) allocator.free(req.metadata_buffer_index)
req.metadata_buffer_index = -1 req.metadata_buffer_index = -1
@@ -214,9 +223,11 @@ class PrefillBootstrapQueue:
) )
return kv_manager return kv_manager
def add(self, req: Req, num_kv_heads: int) -> None: def create_sender(self, req: Req, num_kv_heads: int) -> bool:
"""Create a KV sender for the request without enqueuing it.
Returns False if the request exceeds KV capacity."""
if self._check_if_req_exceed_kv_capacity(req): if self._check_if_req_exceed_kv_capacity(req):
return return False
backend = ( backend = (
TransferBackend.FAKE TransferBackend.FAKE
@@ -235,6 +246,42 @@ class PrefillBootstrapQueue:
pp_rank=self.pp_rank, pp_rank=self.pp_rank,
) )
self._process_req(req) self._process_req(req)
req.pending_bootstrap = True
return True
def ensure_metadata_buffer(self, req: Req) -> bool:
if req.metadata_buffer_index >= 0:
return True
if self.req_to_metadata_buffer_idx_allocator.available_size() == 0:
return False
req.metadata_buffer_index = self.req_to_metadata_buffer_idx_allocator.alloc()
assert req.metadata_buffer_index is not None
return True
def finalize_bootstrap(self, req: Req) -> bool:
"""Initialize the sender after bootstrap completes.
Returns False if no metadata buffer is available (non-terminal)."""
assert req.pending_bootstrap, f"finalize_bootstrap is not idempotent"
if not self.ensure_metadata_buffer(req):
return False
req.time_stats.set_bootstrap_done_time()
num_kv_indices = len(req.origin_input_ids)
decode_prefix_len = req.disagg_kv_sender.pop_decode_prefix_len()
req.start_send_idx = decode_prefix_len
num_kv_indices_to_send = num_kv_indices - decode_prefix_len
num_pages = kv_to_page_num(
num_kv_indices_to_send, self.token_to_kv_pool.page_size
)
req.disagg_kv_sender.init(num_pages, req.metadata_buffer_index)
req.pending_bootstrap = False
return True
def add(self, req: Req, num_kv_heads: int) -> None:
if not self.create_sender(req, num_kv_heads):
return
self.queue.append(req) self.queue.append(req)
def extend(self, reqs: List[Req], num_kv_heads: int) -> None: def extend(self, reqs: List[Req], num_kv_heads: int) -> None:
@@ -291,54 +338,31 @@ class PrefillBootstrapQueue:
if req.rid not in rids_to_check: if req.rid not in rids_to_check:
continue continue
if poll == KVPoll.Bootstrapping: if poll == KVPoll.Failed:
continue self.scheduler.handle_bootstrap_failure(req)
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.output_streamer.stream_output([req], req.return_logprob)
indices_to_remove.add(i) indices_to_remove.add(i)
failed_reqs.append(req) failed_reqs.append(req)
if self.scheduler.metrics_reporter.enable_metrics: elif poll == KVPoll.Bootstrapping:
self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() if (
if self.scheduler.enable_hicache_storage: req.time_stats.prefill_retry_count
# to release prefetch events associated with the request < self.scheduler.server_args.optimistic_prefill_retries
self.scheduler.tree_cache.release_aborted_request(req.rid) and not req.is_retracted # engine paused
continue ):
if not self.ensure_metadata_buffer(req):
# KV.WaitingForInput - decode is ready to receive. initialize the kv sender continue # no more metadata buffer
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:
break
req.metadata_buffer_index = (
self.req_to_metadata_buffer_idx_allocator.alloc()
)
assert req.metadata_buffer_index is not None
# Cal number of pages to send
# if decode has a cached prefix, we need to send the delta indices
# otherwise, send the entire request
decode_prefix_len = req.disagg_kv_sender.pop_decode_prefix_len()
req.start_send_idx = decode_prefix_len
num_kv_indices_to_send = num_kv_indices - decode_prefix_len
num_pages = kv_to_page_num(
num_kv_indices_to_send, self.token_to_kv_pool.page_size
)
req.disagg_kv_sender.init(num_pages, req.metadata_buffer_index)
bootstrapped_reqs.append(req) bootstrapped_reqs.append(req)
indices_to_remove.add(i) indices_to_remove.add(i)
req.time_stats.set_wait_queue_entry_time() req.time_stats.set_wait_queue_entry_time()
elif poll == KVPoll.WaitingForInput:
if not self.finalize_bootstrap(req):
continue
bootstrapped_reqs.append(req)
indices_to_remove.add(i)
req.time_stats.set_wait_queue_entry_time()
else:
raise RuntimeError(
f"Unexpected poll state {poll} for req {req.rid} in pop_bootstrapped"
)
self.queue = [ self.queue = [
entry for i, entry in enumerate(self.queue) if i not in indices_to_remove entry for i, entry in enumerate(self.queue) if i not in indices_to_remove
@@ -521,13 +545,50 @@ class SchedulerDisaggregationPrefillMixin:
v.tolist() for v in logits_output.next_token_token_ids_logprobs_val v.tolist() for v in logits_output.next_token_token_ids_logprobs_val
] ]
def advance_logprob_pt(i: int, req: Req) -> None:
nonlocal logprob_pt
if not req.return_logprob or extend_input_len_per_req is None:
return
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:
logprob_pt += extend_input_len - extend_logprob_start_len
# Poll optimistic prefill requests in this batch.
# Note: In overlap scheduling, a chunked request that was still pending
# during process_prefill_chunk is not checked again here.
# If it becomes ready in the gap, we still retry the request to keep
# chunked-prefill state management simple.
optimistic_polls = {}
optimistic_reqs = [
(i, req)
for i, req in enumerate(batch.reqs)
if req.pending_bootstrap and req.inflight_middle_chunks <= 0
]
if optimistic_reqs:
polls = poll_and_all_reduce_attn_cp_tp_group(
[req.disagg_kv_sender for _, req in optimistic_reqs],
self.attn_cp_cpu_group,
self.attn_tp_cpu_group,
)
optimistic_polls = {
idx: poll for (idx, _), poll in zip(optimistic_reqs, polls)
}
for i, (req, next_token_id) in enumerate( for i, (req, next_token_id) in enumerate(
zip(batch.reqs, next_token_ids, strict=True) zip(batch.reqs, next_token_ids, strict=True)
): ):
if req.inflight_middle_chunks <= 0: if req.inflight_middle_chunks <= 0:
req.time_stats.set_prefill_finished_time() req.time_stats.set_prefill_finished_time()
# There is no output_ids for prefill # For optimistic requests, check bootstrap before side effects
if i in optimistic_polls:
if not self.handle_pending_bootstrap(
req, optimistic_polls[i], defer_release=False
):
advance_logprob_pt(i, req)
continue
req.output_ids.append(next_token_id) req.output_ids.append(next_token_id)
maybe_cache_unfinished_req(req, self.tree_cache) maybe_cache_unfinished_req(req, self.tree_cache)
self.disagg_prefill_inflight_queue.append(req) self.disagg_prefill_inflight_queue.append(req)
@@ -558,12 +619,9 @@ class SchedulerDisaggregationPrefillMixin:
req.time_stats.set_prefill_transfer_queue_entry_time() req.time_stats.set_prefill_transfer_queue_entry_time()
if req.grammar is not None: if req.grammar is not None:
# FIXME: this try-except block is for handling unexpected xgrammar issue.
try: try:
req.grammar.accept_token(next_token_id) req.grammar.accept_token(next_token_id)
except ValueError as e: 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}" error_message = f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}"
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
prepare_abort( prepare_abort(
@@ -576,11 +634,24 @@ class SchedulerDisaggregationPrefillMixin:
# being chunked reqs' prefill is not finished # being chunked reqs' prefill is not finished
req.inflight_middle_chunks -= 1 req.inflight_middle_chunks -= 1
# Overlap deferred release for optimistic requests stopped in process_prefill_chunk
if req.pending_bootstrap:
advance_logprob_pt(i, req)
self.optimistic_release_and_requeue(req)
req.time_stats.set_last_chunked_prefill_finish_time()
continue
# Optimistic bootstrap can fail while this overlapped chunk is
# already running. Drop aborted chunks instead of sending KV.
if is_aborted(req):
advance_logprob_pt(i, req)
req.time_stats.set_last_chunked_prefill_finish_time()
continue
if req.return_logprob: if req.return_logprob:
extend_logprob_start_len = extend_logprob_start_len_per_req[i] extend_logprob_start_len = extend_logprob_start_len_per_req[i]
extend_input_len = extend_input_len_per_req[i] extend_input_len = extend_input_len_per_req[i]
if extend_logprob_start_len < extend_input_len: if extend_logprob_start_len < extend_input_len:
# Update input logprobs.
num_input_logprobs = extend_input_len - extend_logprob_start_len num_input_logprobs = extend_input_len - extend_logprob_start_len
self.batch_result_processor.logprob_result_processor.add_input_logprob_return_values( self.batch_result_processor.logprob_result_processor.add_input_logprob_return_values(
i, i,
@@ -593,6 +664,9 @@ class SchedulerDisaggregationPrefillMixin:
logprob_pt += num_input_logprobs logprob_pt += num_input_logprobs
if self.enable_overlap: if self.enable_overlap:
assert (
req.metadata_buffer_index >= 0
), f"Req {req.rid} does not have metadata buffer allocated"
self.send_kv_chunk(req, last_chunk=False, end_idx=req.tmp_end_idx) self.send_kv_chunk(req, last_chunk=False, end_idx=req.tmp_end_idx)
req.time_stats.set_last_chunked_prefill_finish_time() req.time_stats.set_last_chunked_prefill_finish_time()
@@ -708,7 +782,7 @@ class SchedulerDisaggregationPrefillMixin:
for req in done_reqs: for req in done_reqs:
req: Req req: Req
release_req_to_metadata_buffer( maybe_release_metadata_buffer(
req, self.req_to_metadata_buffer_idx_allocator req, self.req_to_metadata_buffer_idx_allocator
) )
@@ -734,12 +808,77 @@ class SchedulerDisaggregationPrefillMixin:
return transferred_rids return transferred_rids
def handle_bootstrap_failure(self: Scheduler, req: Req) -> None:
error_message = (
f"Prefill bootstrap failed for request rank={self.ps.tp_rank} "
f"{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})
if req.req_pool_idx is not None or self.tree_cache.supports_mamba():
release_kv_cache(req, self.tree_cache)
maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator)
req.pending_bootstrap = False
prepare_abort(req, error_message, status_code=HTTPStatus.INTERNAL_SERVER_ERROR)
self.output_streamer.stream_output([req], req.return_logprob)
if self.metrics_reporter.enable_metrics:
self.metrics_collector.increment_bootstrap_failed_reqs()
if self.enable_hicache_storage:
self.tree_cache.release_aborted_request(req.rid)
def handle_pending_bootstrap(
self: Scheduler, req: Req, poll: KVPoll, defer_release: bool
) -> bool:
"""Return True when bootstrap is finalized and KV transfer can proceed."""
if poll == KVPoll.Failed:
self.handle_bootstrap_failure(req)
return False
elif poll == KVPoll.Bootstrapping:
if not defer_release:
self.optimistic_release_and_requeue(req)
return False
elif poll == KVPoll.WaitingForInput:
force_retry = should_force_retry(req) # test hook
if force_retry:
if not defer_release:
self.optimistic_release_and_requeue(req)
return False
# Metadata buffer was allocated in pop_bootstrapped before
# the request entered the waiting queue, so finalize should not fail.
assert self.disagg_prefill_bootstrap_queue.finalize_bootstrap(req)
return True
else:
raise RuntimeError(
f"Unexpected poll state {poll} for req {req.rid} in handle_pending_bootstrap"
)
def check_bootstrap(self: Scheduler, req: Req) -> bool:
"""Check bootstrap status for an optimistic prefilled request.
Returns True if bootstrap is finished."""
if not req.pending_bootstrap:
return True
polls = poll_and_all_reduce_attn_cp_tp_group(
[req.disagg_kv_sender],
self.attn_cp_cpu_group,
self.attn_tp_cpu_group,
)
return self.handle_pending_bootstrap(
req, polls[0], defer_release=self.enable_overlap
)
def process_prefill_chunk(self: Scheduler) -> None: def process_prefill_chunk(self: Scheduler) -> None:
chunked_req_to_exclude = set() chunked_req_to_exclude = set()
if self.chunked_req: if self.chunked_req:
chunked_req_to_exclude.add(self.chunked_req) chunked_req_to_exclude.add(self.chunked_req)
maybe_cache_unfinished_req(self.chunked_req, self.tree_cache, chunked=True) maybe_cache_unfinished_req(self.chunked_req, self.tree_cache, chunked=True)
if self.enable_overlap:
if not self.check_bootstrap(self.chunked_req):
self.chunked_req = None # stop the current chunked prefill
elif self.enable_overlap:
# Delay KV transfer to process_batch_result_disagg_prefill when overlap is enabled to ensure results are resolved # 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( self.chunked_req.tmp_end_idx = min(
len(self.chunked_req.fill_ids), len(self.chunked_req.fill_ids),
@@ -747,6 +886,8 @@ class SchedulerDisaggregationPrefillMixin:
) )
else: else:
self.send_kv_chunk(self.chunked_req) self.send_kv_chunk(self.chunked_req)
if self.chunked_req is not None:
self.running_batch.batch_is_full = False self.running_batch.batch_is_full = False
if self.last_batch and self.last_batch.forward_mode.is_extend(): if self.last_batch and self.last_batch.forward_mode.is_extend():
@@ -853,3 +994,34 @@ class SchedulerDisaggregationPrefillMixin:
return return
req.disagg_kv_sender.send(page_indices, state_indices) req.disagg_kv_sender.send(page_indices, state_indices)
req.start_send_idx = end_idx req.start_send_idx = end_idx
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
"""Release KV cache and requeue an optimistic prefill request."""
max_retries = self.server_args.optimistic_prefill_retries
maybe_cache_unfinished_req(req, self.tree_cache)
release_kv_cache(req, self.tree_cache)
req.reset_for_retract()
req.output_ids = array("q")
req.start_send_idx = 0
req.tmp_end_idx = -1
req.hidden_states_tensor = None
req.pending_bootstrap = True
req.time_stats.reset_prefill_retry_time()
if req.time_stats.prefill_retry_count >= max_retries:
logger.info(
f"Req {req.rid} exhausted optimistic prefill retries "
"falling back to bootstrap queue"
)
# Reset it so the next real bootstrap done can be recorded.
req.time_stats.bootstrap_done_time = 0.0
self.disagg_prefill_bootstrap_queue.queue.append(req)
else:
req.time_stats.prefill_retry_count += 1
logger.info(
f"Req {req.rid} optimistic prefill retry "
f"{req.time_stats.prefill_retry_count}/{max_retries}"
)
if self.metrics_reporter.enable_metrics:
self.metrics_collector.increment_prefill_retries(1)
req.time_stats.set_wait_queue_entry_time()
self.waiting_queue.insert(0, req)
+28 -10
View File
@@ -50,8 +50,22 @@ class DisaggregationMode(Enum):
# Synchronization # Synchronization
######################### #########################
# env var for testing failure, convert to float explicitly
FAILURE_PROB = float(os.getenv("DISAGGREGATION_TEST_FAILURE_PROB", 0)) def _get_failure_prob() -> float:
try:
return float(envs.SGLANG_TEST_DISAGG_FAILURE_PROB.get())
except Exception:
# fallback to legacy env var
return float(os.getenv("DISAGGREGATION_TEST_FAILURE_PROB", "0"))
def _poll_with_failure_injection(pollers) -> List[int]:
if (failure_prob := _get_failure_prob()) > 0:
return [
int(KVPoll.Failed) if random.random() < failure_prob else int(poller.poll())
for poller in pollers
]
return [int(poller.poll()) for poller in pollers]
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool: def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
@@ -87,13 +101,7 @@ def poll_and_all_reduce(
server_args: Optional[ServerArgs] = None, server_args: Optional[ServerArgs] = None,
): ):
# at a certain prob, the poll is failed to simulate failure # at a certain prob, the poll is failed to simulate failure
if FAILURE_PROB > 0: polls = _poll_with_failure_injection(pollers)
polls = [
int(KVPoll.Failed) if random.random() < FAILURE_PROB else int(poller.poll())
for poller in pollers
]
else:
polls = [int(poller.poll()) for poller in pollers]
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed. # Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if ( if (
@@ -141,7 +149,9 @@ def poll_and_all_reduce_with_staging(
): ):
staging_handler.advance_scatter(decode_req) staging_handler.advance_scatter(decode_req)
raw_polls = [int(dr.kv_receiver.poll()) for dr in decode_reqs] # allow test injection of failure probability at runtime
receivers = [dr.kv_receiver for dr in decode_reqs]
raw_polls = _poll_with_failure_injection(receivers)
for i, decode_req in enumerate(decode_reqs): for i, decode_req in enumerate(decode_reqs):
if raw_polls[i] == int(KVPoll.Success): if raw_polls[i] == int(KVPoll.Success):
if decode_req.kv_receiver.require_staging and not staging_handler.is_done( if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
@@ -689,3 +699,11 @@ def prepare_abort(req: Req, error_message: str, status_code=None):
req.logprob.input_top_logprobs_idx = [] req.logprob.input_top_logprobs_idx = []
req.logprob.input_token_ids_logprobs_val = [] req.logprob.input_token_ids_logprobs_val = []
req.logprob.input_token_ids_logprobs_idx = [] req.logprob.input_token_ids_logprobs_idx = []
def is_aborted(req: Req) -> bool:
from sglang.srt.managers.schedule_batch import FINISH_ABORT
return isinstance(req.to_finish, FINISH_ABORT) or isinstance(
req.finished_reason, FINISH_ABORT
)
+2
View File
@@ -250,6 +250,7 @@ class Envs:
SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64) SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64)
SGLANG_NATIVE_MOVE_KV_CACHE = EnvBool(False) SGLANG_NATIVE_MOVE_KV_CACHE = EnvBool(False)
SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(True) SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(True)
SGLANG_TEST_DISAGG_FAILURE_PROB = EnvFloat(0.0)
# Scheduler: memory leak test # Scheduler: memory leak test
SGLANG_TEST_RETRACT = EnvBool(False) SGLANG_TEST_RETRACT = EnvBool(False)
@@ -323,6 +324,7 @@ class Envs:
# Test: pd-disaggregation # Test: pd-disaggregation
SGLANG_TEST_PD_DISAGG_BACKEND = EnvStr("mooncake") SGLANG_TEST_PD_DISAGG_BACKEND = EnvStr("mooncake")
SGLANG_TEST_PD_DISAGG_DEVICES = EnvStr(None) SGLANG_TEST_PD_DISAGG_DEVICES = EnvStr(None)
SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB = EnvFloat(0.0)
# Model Parallel # Model Parallel
SGLANG_USE_MESSAGE_QUEUE_BROADCASTER = EnvBool(True) SGLANG_USE_MESSAGE_QUEUE_BROADCASTER = EnvBool(True)
@@ -955,6 +955,9 @@ class Req(ReqDllmMixin):
# We use `tmp_end_idx` to store the end index of the kv cache to send. # We use `tmp_end_idx` to store the end index of the kv cache to send.
self.tmp_end_idx: int = -1 self.tmp_end_idx: int = -1
self.metadata_buffer_index: int = -1 self.metadata_buffer_index: int = -1
# Used in overlap sequence to signal that an optimistic request should
# abort chunking. Set in create_sender, consumed in process_batch_result.
self.pending_bootstrap = False
# For Matryoshka embeddings # For Matryoshka embeddings
self.dimensions = dimensions self.dimensions = dimensions
+23 -5
View File
@@ -443,8 +443,10 @@ class PrefillAdder:
self.preempt_list = [] self.preempt_list = []
self.new_chunked_req = None self.new_chunked_req = None
self.log_hit_tokens = 0 self.log_hit_tokens = 0
self.reprocessed_log_hit_tokens = 0
# TODO(lsyin): report the real input tokens excluding page alignment # TODO(lsyin): report the real input tokens excluding page alignment
self.log_input_tokens = 0 self.log_input_tokens = 0
self.reprocessed_log_input_tokens = 0
if running_batch is not None: if running_batch is not None:
# Estimate the offset in the remaining token space # Estimate the offset in the remaining token space
@@ -580,7 +582,11 @@ class PrefillAdder:
return AddReqResult.CONTINUE return AddReqResult.CONTINUE
def _update_prefill_budget( def _update_prefill_budget(
self, prefix_len: int, extend_input_len: int, max_new_tokens: int self,
prefix_len: int,
extend_input_len: int,
max_new_tokens: int,
retracted_stain: bool,
): ):
# TODO(lsyin): check this workaround logic, which only ensures the prefill will not out of memory, and may be too conservative # TODO(lsyin): check this workaround logic, which only ensures the prefill will not out of memory, and may be too conservative
extend_input_len = self.ceil_paged_tokens(extend_input_len) extend_input_len = self.ceil_paged_tokens(extend_input_len)
@@ -599,8 +605,13 @@ class PrefillAdder:
elif self.rem_chunk_tokens is not None: elif self.rem_chunk_tokens is not None:
self.rem_chunk_tokens -= extend_input_len self.rem_chunk_tokens -= extend_input_len
# reprocessed_log_* is a subset of log_*; metrics_reporter subtracts it
# when computing the first-attempt prefix cache hit rate.
self.log_hit_tokens += prefix_len self.log_hit_tokens += prefix_len
self.log_input_tokens += extend_input_len self.log_input_tokens += extend_input_len
if retracted_stain:
self.reprocessed_log_hit_tokens += prefix_len
self.reprocessed_log_input_tokens += extend_input_len
def _get_dllm_remain_tokens(self) -> int: def _get_dllm_remain_tokens(self) -> int:
_rem_tokens = min( _rem_tokens = min(
@@ -628,7 +639,7 @@ class PrefillAdder:
self.can_run_list.append(req) self.can_run_list.append(req)
self._update_prefill_budget(prefix_len, trunc_len, 0) self._update_prefill_budget(prefix_len, trunc_len, 0, req.retracted_stain)
def _req_inc_lock_ref(self, req: Req): def _req_inc_lock_ref(self, req: Req):
result = self.tree_cache.inc_lock_ref(req.last_node) result = self.tree_cache.inc_lock_ref(req.last_node)
@@ -654,7 +665,9 @@ class PrefillAdder:
if not truncated if not truncated
else 0 else 0
) )
self._update_prefill_budget(0, req.extend_input_len, max_new_tokens) self._update_prefill_budget(
0, req.extend_input_len, max_new_tokens, req.retracted_stain
)
# Return based on remaining token availability # Return based on remaining token availability
return ( return (
@@ -693,6 +706,7 @@ class PrefillAdder:
if not truncated if not truncated
else 0 else 0
), ),
req.retracted_stain,
) )
# Return if chunked prefill not finished # Return if chunked prefill not finished
@@ -794,6 +808,7 @@ class PrefillAdder:
0, 0,
req.extend_input_len, req.extend_input_len,
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS), min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
req.retracted_stain,
) )
else: else:
if self.rem_chunk_tokens <= 0: if self.rem_chunk_tokens <= 0:
@@ -806,7 +821,7 @@ class PrefillAdder:
req.fill_ids = req.fill_ids[:trunc_len] req.fill_ids = req.fill_ids[:trunc_len]
self.can_run_list.append(req) self.can_run_list.append(req)
self.new_chunked_req = req self.new_chunked_req = req
self._update_prefill_budget(0, trunc_len, 0) self._update_prefill_budget(0, trunc_len, 0, req.retracted_stain)
return self.budget_state() return self.budget_state()
@@ -926,6 +941,7 @@ class PrefillAdder:
req.sampling_params.max_new_tokens, req.sampling_params.max_new_tokens,
CLIP_MAX_NEW_TOKENS, CLIP_MAX_NEW_TOKENS,
), ),
req.retracted_stain,
) )
else: else:
# Make sure at least one page is available # Make sure at least one page is available
@@ -960,7 +976,9 @@ class PrefillAdder:
self.new_chunked_req = req self.new_chunked_req = req
self._req_inc_lock_ref(req) self._req_inc_lock_ref(req)
self._update_prefill_budget(prefix_len, trunc_len, 0) self._update_prefill_budget(
prefix_len, trunc_len, 0, req.retracted_stain
)
return self.budget_state() return self.budget_state()
+10 -2
View File
@@ -54,7 +54,7 @@ from sglang.srt.disaggregation.encode_receiver import create_mm_receiver
from sglang.srt.disaggregation.prefill import ( from sglang.srt.disaggregation.prefill import (
PrefillBootstrapQueue, PrefillBootstrapQueue,
SchedulerDisaggregationPrefillMixin, SchedulerDisaggregationPrefillMixin,
release_req_to_metadata_buffer, maybe_release_metadata_buffer,
) )
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
DisaggregationMode, DisaggregationMode,
@@ -3586,9 +3586,17 @@ class Scheduler(
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
# For disaggregation prefill mode, free the metadata buffer index # For disaggregation prefill mode, free the metadata buffer index
if self.disaggregation_mode == DisaggregationMode.PREFILL: if self.disaggregation_mode == DisaggregationMode.PREFILL:
release_req_to_metadata_buffer( bootstrap_pending = req.pending_bootstrap
maybe_release_metadata_buffer(
req, self.req_to_metadata_buffer_idx_allocator req, self.req_to_metadata_buffer_idx_allocator
) )
if (
bootstrap_pending
and hasattr(req, "disagg_kv_sender")
and req.disagg_kv_sender is not None
):
if hasattr(req.disagg_kv_sender, "abort"):
req.disagg_kv_sender.abort()
# For mamba radix cache # For mamba radix cache
if ( if (
@@ -60,6 +60,8 @@ class PrefillStats:
new_token_ratio: float new_token_ratio: float
num_running_reqs: QueueCount num_running_reqs: QueueCount
num_new_seqs: int # len(can_run_list) num_new_seqs: int # len(can_run_list)
reprocessed_log_input_tokens: int = 0
reprocessed_log_hit_tokens: int = 0
num_pending_tokens: int = 0 num_pending_tokens: int = 0
@classmethod @classmethod
@@ -73,6 +75,8 @@ class PrefillStats:
return cls( return cls(
log_input_tokens=adder.log_input_tokens, log_input_tokens=adder.log_input_tokens,
log_hit_tokens=adder.log_hit_tokens, log_hit_tokens=adder.log_hit_tokens,
reprocessed_log_input_tokens=adder.reprocessed_log_input_tokens,
reprocessed_log_hit_tokens=adder.reprocessed_log_hit_tokens,
new_token_ratio=adder.new_token_ratio, new_token_ratio=adder.new_token_ratio,
num_running_reqs=QueueCount.from_reqs( num_running_reqs=QueueCount.from_reqs(
running_reqs, enable_priority_scheduling running_reqs, enable_priority_scheduling
@@ -578,9 +582,16 @@ class SchedulerMetricsReporter:
) )
priority_enabled = self.scheduler.enable_priority_scheduling priority_enabled = self.scheduler.enable_priority_scheduling
total_tokens = prefill_stats.log_input_tokens + prefill_stats.log_hit_tokens effective_input_tokens = (
prefill_stats.log_input_tokens
- prefill_stats.reprocessed_log_input_tokens
)
effective_hit_tokens = (
prefill_stats.log_hit_tokens - prefill_stats.reprocessed_log_hit_tokens
)
total_tokens = effective_input_tokens + effective_hit_tokens
cache_hit_rate = ( cache_hit_rate = (
prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0 effective_hit_tokens / total_tokens if total_tokens > 0 else 0.0
) )
# Basics # Basics
@@ -21,6 +21,8 @@ import uuid
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing_extensions import Self
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import ( from sglang.srt.observability.metrics_collector import (
@@ -230,7 +232,7 @@ class ReqTimeStatsBase:
diff_realtime_monotonic: float = 0.0 diff_realtime_monotonic: float = 0.0
@classmethod @classmethod
def new_from_obj(cls, obj: ReqTimeStatsBase, *args, **kwargs) -> "ReqTimeStatsBase": def new_from_obj(cls, obj: Optional[ReqTimeStatsBase], *args, **kwargs) -> Self:
calibrate_time_diff() calibrate_time_diff()
new_obj = cls(*args, **kwargs) new_obj = cls(*args, **kwargs)
if obj is None: if obj is None:
@@ -668,6 +670,17 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
if self.trace_ctx.tracing_enable: if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_event("retract", 1, convert_time_to_realtime_ns(ts)) self.trace_ctx.trace_event("retract", 1, convert_time_to_realtime_ns(ts))
def reset_prefill_retry_time(self):
self.wait_queue_entry_time = 0.0
self.forward_entry_time = 0.0
self.prefill_finished_time = 0.0
self.completion_time = 0.0
self.prefill_transfer_queue_entry_time = 0.0
self.prefill_kv_transfer_finish_time = 0.0
self.last_forward_entry_time = 0.0
self.last_prefill_finished_time = 0.0
self.last_chunked_prefill_finish_time = 0.0
def set_wait_queue_entry_time(self, ts=None): def set_wait_queue_entry_time(self, ts=None):
ts = ts or time.perf_counter() ts = ts or time.perf_counter()
if self.wait_queue_entry_time == 0.0: if self.wait_queue_entry_time == 0.0:
@@ -1008,22 +1021,22 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
and forward_duration >= 0 and forward_duration >= 0
), f"bootstrap_queue_duration={bootstrap_queue_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0" ), f"bootstrap_queue_duration={bootstrap_queue_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0"
# Break down bootstrap_queue_duration into sub-phases if (
if self.bootstrap_done_time > 0: self.bootstrap_done_time > 0
and self.prefill_bootstrap_queue_entry_time > 0
):
bootstrap_duration = self.duration_between( bootstrap_duration = self.duration_between(
self.prefill_bootstrap_queue_entry_time, self.bootstrap_done_time self.prefill_bootstrap_queue_entry_time, self.bootstrap_done_time
) )
alloc_wait_duration = self.duration_between(
self.bootstrap_done_time, self.wait_queue_entry_time
)
if SGLANG_TEST_REQUEST_TIME_STATS: if SGLANG_TEST_REQUEST_TIME_STATS:
assert ( assert (
bootstrap_duration >= 0 and alloc_wait_duration >= 0 bootstrap_duration >= 0
), f"bootstrap_duration={bootstrap_duration} < 0 or alloc_wait_duration={alloc_wait_duration} < 0" ), f"bootstrap_duration={bootstrap_duration} < 0"
bootstrap_fields = ( bootstrap_fields = (
f"bootstrap_duration={self.format_duration(bootstrap_duration)}, " f"bootstrap_duration={self.format_duration(bootstrap_duration)}, "
f"alloc_wait_duration={self.format_duration(alloc_wait_duration)}, "
) )
elif self.bootstrap_done_time > 0:
bootstrap_fields = f"bootstrap_done_time={self.format_wallclock(self.bootstrap_done_time)}, "
else: else:
bootstrap_fields = f"bootstrap_queue_duration={self.format_duration(bootstrap_queue_duration)}, " bootstrap_fields = f"bootstrap_queue_duration={self.format_duration(bootstrap_queue_duration)}, "
+33 -2
View File
@@ -36,7 +36,9 @@ from sglang.srt.arg_groups.argparse_actions import (
DeprecatedStoreTrueAction, DeprecatedStoreTrueAction,
LoRAPathAction, LoRAPathAction,
) )
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch from sglang.srt.configs.linear_attn_model_registry import (
get_linear_attn_spec_by_arch,
)
from sglang.srt.connector import ConnectorType from sglang.srt.connector import ConnectorType
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
parse_ib_device_config, parse_ib_device_config,
@@ -818,6 +820,7 @@ class ServerArgs:
num_reserved_decode_tokens: int = 512 # used for decode kv cache offload in PD num_reserved_decode_tokens: int = 512 # used for decode kv cache offload in PD
# FIXME: hack to reduce ITL when decode bs is small # FIXME: hack to reduce ITL when decode bs is small
disaggregation_decode_polling_interval: int = 1 disaggregation_decode_polling_interval: int = 1
optimistic_prefill_retries: int = 0
# Encode prefill disaggregation # Encode prefill disaggregation
encoder_only: bool = False encoder_only: bool = False
@@ -1776,6 +1779,7 @@ class ServerArgs:
is_deepseek_dsa, is_deepseek_dsa,
) )
self.uses_mamba_radix_cache = False
if parse_connector_type(self.model_path) == ConnectorType.INSTANCE: if parse_connector_type(self.model_path) == ConnectorType.INSTANCE:
return return
@@ -1783,7 +1787,7 @@ class ServerArgs:
model_arch = hf_config.architectures[0] model_arch = hf_config.architectures[0]
_hybrid_spec = get_linear_attn_spec_by_arch(model_arch) _hybrid_spec = get_linear_attn_spec_by_arch(model_arch)
if _hybrid_spec is not None: if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
self._handle_mamba_radix_cache( self._handle_mamba_radix_cache(
model_arch=model_arch, model_arch=model_arch,
support_mamba_cache=_hybrid_spec.support_mamba_cache, support_mamba_cache=_hybrid_spec.support_mamba_cache,
@@ -2620,6 +2624,8 @@ class ServerArgs:
sm100_default_attention_backend: str = None, sm100_default_attention_backend: str = None,
fallback_attention_backend: str = "triton", fallback_attention_backend: str = "triton",
): ):
self.uses_mamba_radix_cache = True
if ( if (
is_sm100_supported() is_sm100_supported()
and self.attention_backend is None and self.attention_backend is None
@@ -4308,6 +4314,24 @@ class ServerArgs:
) )
def _handle_other_validations(self): def _handle_other_validations(self):
# Handle optimistic prefill validation
if (
self.optimistic_prefill_retries > 0
and self.disaggregation_mode == "prefill"
):
if self.pp_size > 1:
logger.warning("Optimistic prefill does not support pp_size > 1")
self.optimistic_prefill_retries = 0
elif self.enable_hierarchical_cache:
logger.warning("Optimistic prefill does not support hierarchical cache")
self.optimistic_prefill_retries = 0
elif getattr(self, "uses_mamba_radix_cache", False):
logger.warning(
"Optimistic prefill does not support models that use "
"mamba radix cache."
)
self.optimistic_prefill_retries = 0
# Handle model inference tensor dump. # Handle model inference tensor dump.
if self.debug_tensor_dump_output_folder is not None: if self.debug_tensor_dump_output_folder is not None:
logger.warning( logger.warning(
@@ -6892,6 +6916,13 @@ class ServerArgs:
help="The interval to poll requests in decode server. Can be set to >1 to reduce the overhead of this.", help="The interval to poll requests in decode server. Can be set to >1 to reduce the overhead of this.",
) )
parser.add_argument(
"--optimistic-prefill-retries",
type=int,
default=ServerArgs.optimistic_prefill_retries,
help="Number of optimistic prefill retries that will skip the bootstrap wait. ",
)
# Encode prefill disaggregation # Encode prefill disaggregation
parser.add_argument( parser.add_argument(
"--encoder-only", "--encoder-only",
@@ -0,0 +1,202 @@
import time
import unittest
import uuid
from concurrent.futures import ThreadPoolExecutor, as_completed
from types import SimpleNamespace
import requests
from prometheus_client.parser import text_string_to_metric_families
from sglang.srt.disaggregation.prefill import should_force_retry
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST
register_cuda_ci(est_time=120, stage="base-b", runner_config="2-gpu-large")
FORCE_RETRY_PROB = 0.1
def rid_that_forces_retry(prefix: str) -> str:
"""Return a rid that the test retry sampler will select."""
for _ in range(1000):
rid = f"{prefix}{uuid.uuid4().hex}"
req = SimpleNamespace(
rid=rid,
is_retracted=False,
time_stats=SimpleNamespace(prefill_retry_count=0),
)
if should_force_retry(req):
return rid
raise RuntimeError("Failed to sample an optimistic prefill retry rid")
class OptimisticPrefillRetryCounterMixin:
def _get_retry_counter(self) -> float:
response = requests.get(f"{self.prefill_url}/metrics")
response.raise_for_status()
total = 0.0
for family in text_string_to_metric_families(response.text):
if family.name != "sglang:num_prefill_retries":
continue
for sample in family.samples:
if sample.name == "sglang:num_prefill_retries_total":
total += sample.value
return total
def assert_retry_counter_increases(self, fn):
before_retries = self._get_retry_counter()
result = fn()
after_retries = self._get_retry_counter()
self.assertGreater(after_retries, before_retries)
return result
class TestOptimisticPrefill(
OptimisticPrefillRetryCounterMixin, PDDisaggregationServerBase
):
@classmethod
def setUpClass(cls):
super().setUpClass()
cls._force_retry_prob_was_set = (
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.is_set()
)
cls._force_retry_prob_value = (
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.get()
)
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.set(FORCE_RETRY_PROB)
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
cls.extra_prefill_args = [
"--optimistic-prefill-retries",
"3",
"--chunked-prefill-size",
"128",
"--enable-metrics",
"--enable-request-time-stats-logging",
]
cls.launch_all()
@classmethod
def tearDownClass(cls):
try:
super().tearDownClass()
finally:
if getattr(cls, "_force_retry_prob_was_set", False):
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.set(
cls._force_retry_prob_value
)
else:
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.clear()
def test_gsm8k(self):
args = SimpleNamespace(
base_url=f"http://{self.base_host}:{self.lb_port}",
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=200,
num_threads=128,
)
metrics = self.assert_retry_counter_increases(lambda: run_eval(args))
print(f"Evaluation metrics: {metrics}")
self.assertGreater(metrics["score"], 0.62)
time.sleep(1) # trigger memory check
def test_logprob(self):
request_id = rid_that_forces_retry("logprob-retry-")
prompt = f"{request_id}: " + "The capital of France is Paris. " * 900
j = self.assert_retry_counter_increases(
lambda: requests.post(
self.lb_url + "/generate",
json={
"rid": request_id,
"text": prompt,
"sampling_params": {"temperature": 0, "max_new_tokens": 8},
"return_logprob": True,
"return_input_logprob": True,
"logprob_start_len": 0,
},
).json()
)
completion_tokens = j["meta_info"]["completion_tokens"]
input_logprobs = j["meta_info"]["input_token_logprobs"]
output_logprobs = j["meta_info"]["output_token_logprobs"]
self.assertGreater(j["meta_info"]["prompt_tokens"], 512)
assert len(output_logprobs) == completion_tokens
assert len(input_logprobs) > 0
class TestOptimisticPrefillFailure(PDDisaggregationServerBase):
@classmethod
def setUpClass(cls):
super().setUpClass()
# enable optimistic prefill retry sampling and disagg failure prob
cls._force_retry_ctx = (
envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.override(
FORCE_RETRY_PROB
)
)
cls._force_retry_ctx.__enter__()
cls._disagg_failure_ctx = envs.SGLANG_TEST_DISAGG_FAILURE_PROB.override(
FORCE_RETRY_PROB
)
cls._disagg_failure_ctx.__enter__()
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
cls.extra_prefill_args = [
"--optimistic-prefill-retries",
"3",
"--chunked-prefill-size",
"128",
"--enable-metrics",
"--enable-request-time-stats-logging",
"--load-format",
"dummy",
]
cls.launch_all()
@classmethod
def tearDownClass(cls):
try:
super().tearDownClass()
finally:
if getattr(cls, "_force_retry_ctx", None):
cls._force_retry_ctx.__exit__(None, None, None)
if getattr(cls, "_disagg_failure_ctx", None):
cls._disagg_failure_ctx.__exit__(None, None, None)
def test_survive_requests(self):
# send many small requests to ensure the engine survives injected failures
n = 100
with ThreadPoolExecutor(max_workers=32) as executor:
futures = []
for i in range(n):
rid = f"survive-{i}-{uuid.uuid4().hex}"
futures.append(
executor.submit(
requests.post,
self.lb_url + "/generate",
json={
"rid": rid,
"text": "Hello world",
"sampling_params": {"temperature": 0, "max_new_tokens": 4},
},
timeout=30,
)
)
for future in as_completed(futures):
try:
_ = future.result()
except Exception:
pass
time.sleep(1) # trigger memory check
if __name__ == "__main__":
unittest.main()
@@ -81,6 +81,7 @@ class TestPrefillAdder(CustomTestCase):
req.output_ids = [0] * output_len req.output_ids = [0] * output_len
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens) req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time) req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time)
req.retracted_stain = False
req.finished.return_value = False req.finished.return_value = False
return req return req