diff --git a/python/sglang/srt/disaggregation/common/staging_buffer.py b/python/sglang/srt/disaggregation/common/staging_buffer.py index e41071cba..cba5dd7a6 100644 --- a/python/sglang/srt/disaggregation/common/staging_buffer.py +++ b/python/sglang/srt/disaggregation/common/staging_buffer.py @@ -771,3 +771,29 @@ def resolve_total_kv_heads( "nor kv_head_num. " "Ensure DecodePreallocQueue._init_kv_manager sets kv_args.kv_head_num." ) + + +def staging_grid_tokens(chunked_prefill_size: Optional[int], page_size: int) -> int: + """Token width of one staging grid slot; shared by prefetch and the + sender's grid alignment.""" + cps = chunked_prefill_size or 8192 + return max(1, cps // page_size) * page_size + + +def compute_grid_segments( + start_idx: int, end_idx: int, base: int, grid_tokens: int +) -> List[Tuple[int, int]]: + """Split [start_idx, end_idx) at grid boundaries base + k * grid_tokens + so each segment maps to exactly one staging slot. An empty range yields + one empty segment (a metadata-only last chunk still needs a send). + """ + segments: List[Tuple[int, int]] = [] + seg_start = start_idx + while seg_start < end_idx: + next_boundary = base + ((seg_start - base) // grid_tokens + 1) * grid_tokens + seg_end = min(next_boundary, end_idx) + segments.append((seg_start, seg_end)) + seg_start = seg_end + if not segments: + segments = [(start_idx, end_idx)] + return segments diff --git a/python/sglang/srt/disaggregation/common/staging_handler.py b/python/sglang/srt/disaggregation/common/staging_handler.py index b2ef46ad1..823bd9460 100644 --- a/python/sglang/srt/disaggregation/common/staging_handler.py +++ b/python/sglang/srt/disaggregation/common/staging_handler.py @@ -12,6 +12,7 @@ import dataclasses import logging import struct import threading +import time from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -82,11 +83,21 @@ class DecodeStagingHandler: self.total_kv_heads = total_kv_heads self.tp_rank = tp_rank self.scheduler = scheduler + # Same stall->Failed semantics (and knob) as _check_waiting_timeout, + # which is unreachable once the receiver has concluded Success. + from sglang.srt.environ import envs + + self.completion_timeout = float( + envs.SGLANG_DISAGGREGATION_WAITING_TIMEOUT.get() + ) self._room_to_decode_req: dict = {} # Stashed at registration: removal paths null decode_req.kv_receiver # before unregister runs, but release_room still needs it. self._room_to_receiver: dict = {} self._wm_subscribers: dict = {} + # room -> chunk_idx -> [(page_start, num_pages, writer_id)] fan-in + # arrivals; handler-owned so room teardown can purge them. + self._writer_counts: dict = {} def register_wm_subscriber(self, receiver, session_id: str) -> None: """Register a prefill's bootstrap connection for watermark broadcasts.""" @@ -96,9 +107,9 @@ class DecodeStagingHandler: if key not in self._wm_subscribers: self._wm_subscribers[key] = (receiver, session_id) - def num_writers_for(self, decode_req) -> int: + def num_writers_for(self, receiver) -> int: """Compute num_writers for a specific request based on its prefill TP.""" - prefill_tp = decode_req.kv_receiver.prefill_info.attn_tp_size + prefill_tp = receiver.prefill_info.attn_tp_size if prefill_tp > self.decode_tp: return prefill_tp // max(1, self.decode_tp) return 1 @@ -141,15 +152,33 @@ class DecodeStagingHandler: def register_decode_req(self, room: int, decode_req: DecodeRequest) -> None: # Called once per room from pop_preallocated, before send_metadata. + decode_req._staging_all_success = False + decode_req._staging_success_ts = 0.0 + decode_req._staging_failed = False decode_req._staging_scatter_done = False decode_req._chunk_events = [] self._room_to_decode_req[room] = decode_req self._room_to_receiver[room] = decode_req.kv_receiver + # Scatter offsets shift suffix-relative page_start by the decode prefix, + # exact only when the prefix is page-aligned. Fail just this request on a + # mismatch instead of raising, which would kill the prefill scheduler. + page_size = self.kv_buffer_info["page_size"] + if decode_req.req.cache_protected_len % page_size != 0: + logger.error( + "[STAGING] decode prefix length %s is not page-aligned " + "(page_size=%s); failing room=%s (staging scatter offsets " + "would be wrong).", + decode_req.req.cache_protected_len, + page_size, + room, + ) + decode_req._staging_failed = True def unregister_decode_req(self, room: int) -> None: # Pop before release_room so no new arrival can start consuming the slots. decode_req = self._room_to_decode_req.pop(room, None) receiver = self._room_to_receiver.pop(room, None) + self._writer_counts.pop(room, None) if decode_req is not None: self.release_room(room, decode_req, receiver) self.kv_manager._staging_ctx.room_receivers.pop(room, None) @@ -214,7 +243,9 @@ class DecodeStagingHandler: if staging_offset < 0 or alloc_id < 0: return False - ok = self._scatter_region(staging_offset, page_start, num_pages, decode_req) + ok = self._scatter_region( + staging_offset, page_start, num_pages, decode_req, receiver + ) if ok: event = torch.cuda.Event() event.record(self.staging_allocator._scatter_stream) @@ -242,38 +273,36 @@ class DecodeStagingHandler: page_start: int, num_pages: int, writer_id: str, - chunk_writer_counts: dict, ) -> bool: """Process a staging chunk arrival from any transport (NIXL RDMA notif or ZMQ CHUNK_READY). - Accumulates writer arrivals in *chunk_writer_counts* and submits scatter - once all writers for this chunk have reported in. Returns True if scatter - was submitted. + Accumulates writer arrivals and submits scatter once all writers for + this chunk have reported in. Returns True if scatter was submitted. """ - chunk_writer_counts[room][chunk_idx].append((page_start, num_pages, writer_id)) - decode_req = self._room_to_decode_req.get(room) - if decode_req is None: + # Read from the stash, not decode_req.kv_receiver: a concurrent teardown + # nulls the latter before unregister removes the room. + receiver = self._room_to_receiver.get(room) + if receiver is None: logger.warning( - "Staging chunk arrived for unregistered room=%s chunk=%d, skipping", + "Staging chunk arrived for unregistered room=%s chunk=%d, " "skipping", room, chunk_idx, ) return False - writers_arrived = len(chunk_writer_counts[room][chunk_idx]) - num_writers = self.num_writers_for(decode_req) - if writers_arrived >= num_writers: + room_counts = self._writer_counts.setdefault(room, {}) + arrivals = room_counts.setdefault(chunk_idx, []) + arrivals.append((page_start, num_pages, writer_id)) + num_writers = self.num_writers_for(receiver) + if len(arrivals) >= num_writers: self.submit_chunk_scatter(room, chunk_idx, page_start, num_pages) - del chunk_writer_counts[room][chunk_idx] + del room_counts[chunk_idx] return True return False def submit_last_scatter_async(self, room: int) -> bool: - """Submit scatter for the last chunk when all ranks report Success. - - Called from decode_thread. Sets ``_scatter_event`` **before** - ``_staging_last_scatter_submitted`` so the main thread sees the - event when it checks the flag (CPython GIL guarantees ordering). - """ + """Record all-ranks Success. Scatter is fully arrival-driven (every + chunk, including the last); advance_scatter completes the room once + no allocation is still waiting for its arrival.""" decode_req = self._room_to_decode_req.get(room) if decode_req is None: logger.warning( @@ -283,15 +312,11 @@ class DecodeStagingHandler: room, ) return False - alloc_id = self._submit_last_scatter(decode_req) - if alloc_id >= 0: - event = torch.cuda.Event() - event.record(self.staging_allocator._scatter_stream) - decode_req._scatter_event = event - decode_req._scatter_alloc_id = alloc_id - decode_req._staging_last_scatter_submitted = True - else: - decode_req._staging_scatter_done = True + if not decode_req._staging_all_success: + # Set the timestamp before the flag so the deadline check never + # reads a zero ts. + decode_req._staging_success_ts = time.monotonic() + decode_req._staging_all_success = True return True # ------------------------------------------------------------------ @@ -302,15 +327,19 @@ class DecodeStagingHandler: """Return True if staging scatter is complete for this request.""" return decode_req._staging_scatter_done and not decode_req._chunk_events - def advance_scatter(self, decode_req: DecodeRequest) -> None: - """Check CUDA events and free completed staging allocations. + def is_failed(self, decode_req: DecodeRequest) -> bool: + """Return True if staging completion timed out for this request.""" + return decode_req._staging_failed - Scatter kernels have already been submitted by the decode_thread - (via submit_chunk_scatter / submit_last_scatter_async). This - method only polls the recorded events and releases staging memory. + def advance_scatter(self, decode_req: DecodeRequest) -> None: + """Poll scatter events, free completed allocations, detect completion. + + The room is done once all ranks reported Success AND every allocation + was scattered AND every event fired; gating on outstanding allocations + keeps it open while a CHUNK_READY is still in flight after Success. + Rooms incomplete past the disaggregation waiting timeout are failed. """ - room = decode_req.req.bootstrap_room - chunk_events = getattr(decode_req, "_chunk_events", None) + chunk_events = decode_req._chunk_events if chunk_events: for i in range(len(chunk_events) - 1, -1, -1): event, alloc_id = chunk_events[i] @@ -318,15 +347,24 @@ class DecodeStagingHandler: chunk_events.pop(i) self._free_and_send_watermark(alloc_id, decode_req) - if not getattr(decode_req, "_staging_last_scatter_submitted", False): + if not decode_req._staging_all_success: return - - event = getattr(decode_req, "_scatter_event", None) - if event is not None and event.query(): - self._free_and_send_watermark(decode_req._scatter_alloc_id, decode_req) - decode_req._scatter_event = None - decode_req._scatter_alloc_id = -1 + room = decode_req.req.bootstrap_room + receiver = self._room_to_receiver.get(room) + chunk_infos = receiver.chunk_staging_infos if receiver is not None else [] + incomplete = bool(chunk_events) or any(info[0] >= 0 for info in chunk_infos) + if not incomplete: decode_req._staging_scatter_done = True + return + elapsed = time.monotonic() - decode_req._staging_success_ts + if elapsed > self.completion_timeout: + logger.error( + "[STAGING] room=%s not complete %.0fs after all-ranks Success " + "(a scatter never arrived); failing the request.", + room, + elapsed, + ) + decode_req._staging_failed = True # ------------------------------------------------------------------ # Internal methods @@ -338,6 +376,7 @@ class DecodeStagingHandler: page_start: int, num_pages: int, decode_req: DecodeRequest, + receiver, ) -> bool: """Submit scatter kernels for a staging region to scatter_stream. @@ -365,9 +404,12 @@ class DecodeStagingHandler: staging_view = self.staging_allocator.buffer.buffer[staging_offset:] req_pool_idx = decode_req.req.req_pool_idx - token_start = page_start * page_size + # page_start is suffix-relative (pages after the decode-side cached + # prefix); req_to_token rows are absolute. + prefix_tokens = decode_req.req.cache_protected_len + token_start = prefix_tokens + page_start * page_size token_end = token_start + num_pages * page_size - prefill_tp = decode_req.kv_receiver.prefill_info.attn_tp_size + prefill_tp = receiver.prefill_info.attn_tp_size with torch.cuda.stream(scatter_stream): kv_indices = self.scheduler.req_to_token_pool.req_to_token[ @@ -392,28 +434,6 @@ class DecodeStagingHandler: return True - def _submit_last_scatter(self, decode_req: DecodeRequest) -> int: - """Submit scatter for the last chunk. Returns alloc_id >= 0, or -1.""" - receiver = decode_req.kv_receiver - chunk_infos = receiver.chunk_staging_infos if receiver is not None else [] - if not chunk_infos: - return -1 - - last_info = chunk_infos[-1] - alloc_id, staging_offset, _, _, last_num_pages = last_info - if staging_offset < 0 or alloc_id < 0: - return -1 - - seq_len = len(decode_req.req.origin_input_ids) - ps = self.scheduler.token_to_kv_pool_allocator.page_size - total_pages = (seq_len + ps - 1) // ps - page_start = total_pages - last_num_pages - - ok = self._scatter_region( - staging_offset, page_start, last_num_pages, decode_req - ) - return alloc_id if ok else -1 - def _free_and_send_watermark( self, alloc_id: int, decode_req: DecodeRequest ) -> None: @@ -542,11 +562,17 @@ class PrefillStagingStrategy: """ def __init__(self, kv_manager, staging_buffer): + from sglang.srt.disaggregation.common.staging_buffer import ( + staging_grid_tokens, + ) + self.kv_manager = kv_manager self.staging_buffer = staging_buffer page_size = kv_manager.kv_buffer_tensors["page_size"] - cps = kv_manager.server_args.chunked_prefill_size or 8192 - self.full_chunk_pages = max(1, cps // page_size) + self.full_chunk_pages = ( + staging_grid_tokens(kv_manager.server_args.chunked_prefill_size, page_size) + // page_size + ) def check_ready( self, @@ -831,11 +857,11 @@ def prefetch_staging_reqs( """ import zmq + from sglang.srt.disaggregation.common.staging_buffer import staging_grid_tokens from sglang.srt.utils.network import NetworkAddress page_size = kv_buffer_tensors["page_size"] - cps = chunked_prefill_size or 8192 - full_chunk_pages = max(1, cps // page_size) + full_chunk_pages = staging_grid_tokens(chunked_prefill_size, page_size) // page_size for session_id, tinfo in transfer_infos[room].items(): # mooncake exposes is_dummy as a dataclass bool field, NIXL exposes it diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index 650f226e0..8b574904f 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -29,6 +29,9 @@ class TransferKVChunk: trace_ctx: Union[TraceReqContext, TraceNullContext] = dataclasses.field( default_factory=TraceNullContext ) + # Set when the staging worker first counts this chunk toward the per-room + # outstanding count; stays set across re-enqueue on a watermark defer. + staging_counted: bool = False def pack_list_of_buffers(buffers: List[bytes]) -> bytes: diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 1907e0ee6..353e50bbe 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -207,6 +207,9 @@ class MooncakeKVManager(CommonKVManager): self.start_prefill_thread() self.session_failures = defaultdict(int) self.failed_sessions = set() + # Per-room count of chunks not yet transferred; teardown waits for + # zero so a deferred chunk is not dropped by an early conclude. + self._staging_outstanding = defaultdict(int) self.session_lock = threading.Lock() # Determine the number of threads to use for kv sender cpu_count = os.cpu_count() @@ -270,7 +273,6 @@ class MooncakeKVManager(CommonKVManager): if self.enable_staging: self._init_staging_allocator() self._staging_handler = None - self._chunk_writer_counts: dict = defaultdict(lambda: defaultdict(list)) self.start_decode_thread() def init_engine(self): @@ -402,7 +404,8 @@ class MooncakeKVManager(CommonKVManager): return PrefillStagingStrategy(self, staging_buffer) def _send_chunk_ready(self, req, chunk_idx, kv_chunk, prefill_unique_rank): - """Notify decode that a non-last staging chunk RDMA is complete.""" + """Notify decode that a staging chunk RDMA is complete (every chunk; + scatter is arrival-driven).""" na = NetworkAddress(req.endpoint, req.dst_port) self._send_multipart_locked( na.to_tcp(), @@ -477,7 +480,7 @@ class MooncakeKVManager(CommonKVManager): "reduce chunked_prefill_size." ) return (-1, False) - if ret == 0 and not kv_chunk.is_last_chunk: + if ret == 0: self._send_chunk_ready(req, chunk_idx, kv_chunk, prefill_unique_rank) return (ret, False) @@ -1549,8 +1552,14 @@ class MooncakeKVManager(CommonKVManager): MooncakeRequestStage.MOONCAKE_WORKER_SEND.level, thread_finish_flag=True, ) + self._staging_outstanding.pop(kv_chunk.room, None) continue + # Count each chunk once; the flag survives re-enqueue on defer. + if not kv_chunk.staging_counted: + self._staging_outstanding[kv_chunk.room] += 1 + kv_chunk.staging_counted = True + if ( self.enable_staging and staging_strategy is None @@ -1812,10 +1821,21 @@ class MooncakeKVManager(CommonKVManager): if staging_deferred: continue - if ( + self._staging_outstanding[kv_chunk.room] -= 1 + # Tear down only when no chunk is still outstanding and the room + # has concluded: already cleared, Success, or a Failed *last* + # chunk. A non-last Failed chunk keeps the room (more chunks may + # follow), not on the last chunk alone since an earlier deferred + # chunk may still need to transfer. + if self._staging_outstanding.get(kv_chunk.room, 0) <= 0 and ( kv_chunk.room not in self.request_status or self.check_status(kv_chunk.room) == KVPoll.Success + or ( + kv_chunk.is_last_chunk + and self.check_status(kv_chunk.room) == KVPoll.Failed + ) ): + self._staging_outstanding.pop(kv_chunk.room, None) if kv_chunk.room in self.transfer_infos: self.transfer_infos.pop(kv_chunk.room) self.req_to_decode_prefix_len.pop(kv_chunk.room, None) @@ -1969,7 +1989,6 @@ class MooncakeKVManager(CommonKVManager): page_start, num_pages, session_id, - self._chunk_writer_counts, ) continue @@ -2004,7 +2023,6 @@ class MooncakeKVManager(CommonKVManager): handler = self._staging_handler if handler.is_staging_room(bootstrap_room): handler.submit_last_scatter_async(bootstrap_room) - self._chunk_writer_counts.pop(bootstrap_room, None) self.update_status(bootstrap_room, KVPoll.Success) elif status == KVPoll.Failed: self.record_failure( @@ -2177,6 +2195,13 @@ class MooncakeKVSender(CommonKVSender): def poll(self) -> KVPoll: if self.conclude_state is None: status = self.kv_mgr.check_status(self.bootstrap_room) + # Hold Success until all staging chunks transferred: a deferred + # chunk can still be pending, and concluding now would drop it. + if ( + status == KVPoll.Success + and self.kv_mgr._staging_outstanding.get(self.bootstrap_room, 0) > 0 + ): + return KVPoll.Transferring if status in (KVPoll.Success, KVPoll.Failed): self.conclude_state = status self.trace_ctx.trace_req_finish() diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 86a5b2904..447f88935 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -481,6 +481,9 @@ class NixlKVManager(CommonKVManager): FastQueue() for _ in range(transfer_queue_size) ] self.exceptions: Dict[int, Exception] = {} + # Per-room count of chunks not yet transferred; teardown waits for + # zero so a deferred chunk is not dropped by an early conclude. + self._staging_outstanding = defaultdict(int) # Mirror mooncake: one staging buffer per worker queue, all # built before workers spawn so each worker owns a private # buffer (no cross-worker contention on the staging ring). @@ -506,7 +509,6 @@ class NixlKVManager(CommonKVManager): if self.enable_staging: self._init_staging_decode_ctx() self._staging_handler = None - self._chunk_writer_counts: dict = defaultdict(lambda: defaultdict(list)) self._start_decode_staging_thread() self._start_heartbeat_checker_thread() else: @@ -1115,10 +1117,16 @@ class NixlKVManager(CommonKVManager): handles: List[Any] = [] try: if self.check_status(room) == KVPoll.Failed: + self._staging_outstanding.pop(room, None) continue assert room in self.transfer_infos + # Count each chunk once; the flag survives re-enqueue on defer. + if not kv_chunk.staging_counted: + self._staging_outstanding[room] += 1 + kv_chunk.staging_counted = True + # Lazily build a per-worker staging strategy bound to this # worker's private staging buffer (matches mooncake). if ( @@ -1329,10 +1337,26 @@ class NixlKVManager(CommonKVManager): break time.sleep(0) + self._staging_outstanding[room] -= 1 if kv_chunk.is_last_chunk: self.update_status(room, KVPoll.Success) - # Drop per-room state on Success (parity with mooncake - # transfer_worker; staging prefetch sets are NIXL-only). + elif self.check_status(room) != KVPoll.Success: + # A deferred earlier chunk can complete after the last chunk + # already concluded Success; don't regress the status. + self.update_status(room, KVPoll.Transferring) + + # Drop per-room state only when no chunk is still outstanding and + # the room has concluded: Success, or a Failed *last* chunk. A + # non-last Failed chunk keeps the room (more chunks may follow); a + # late chunk for an already-Failed room is skipped at loop top. + if self._staging_outstanding.get(room, 0) <= 0 and ( + self.check_status(room) == KVPoll.Success + or ( + kv_chunk.is_last_chunk + and self.check_status(room) == KVPoll.Failed + ) + ): + self._staging_outstanding.pop(room, None) self.transfer_infos.pop(room, None) self.req_to_decode_prefix_len.pop(room, None) if self.enable_staging and self._staging_ctx is not None: @@ -1341,8 +1365,6 @@ class NixlKVManager(CommonKVManager): for k in list(self._staging_ctx.prefetch_requested): if k[0] == room: self._staging_ctx.prefetch_requested.discard(k) - else: - self.update_status(room, KVPoll.Transferring) except Exception as e: # Catch all exceptions to prevent silently killing this # worker thread, but still propagate via failure_exception(). @@ -2461,10 +2483,12 @@ class NixlKVManager(CommonKVManager): page_start = int(components[6]) num_pages = int(components[7]) agent_name = components[8] if len(components) > 8 else "" - self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank) + # Count this notif's own arrival BEFORE _track_kv_arrival, which can + # conclude the transfer and record all-ranks Success. self._handle_staging_chunk_arrived( room, chunk_idx, page_start, num_pages, agent_name ) + self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank) def _handle_aux_notification(self, room: int, components: List[str]): """Handle an aux notification and trigger last scatter if staging is complete. @@ -2567,7 +2591,6 @@ class NixlKVManager(CommonKVManager): page_start, num_pages, agent_name, - self._chunk_writer_counts, ) def _maybe_submit_last_scatter(self, room: int): @@ -2587,7 +2610,6 @@ class NixlKVManager(CommonKVManager): handler = self._staging_handler if handler is not None and handler.is_staging_room(room): handler.submit_last_scatter_async(room) - self._chunk_writer_counts.pop(room, None) def check_transfer_done(self, room: int): if room not in self.transfer_statuses: @@ -2762,6 +2784,13 @@ class NixlKVSender(CommonKVSender): if self._send_failed: return KVPoll.Failed # type: ignore status = self.kv_mgr.check_status(self.bootstrap_room) + # Hold Success until all staging chunks transferred: a deferred chunk + # can still be pending, and concluding now would drop it. + if ( + status == KVPoll.Success + and self.kv_mgr._staging_outstanding.get(self.bootstrap_room, 0) > 0 + ): + return KVPoll.Transferring # type: ignore if ( status == KVPoll.Success and self._transfer_start_time is not None diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 5642c64ef..02db16862 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -32,6 +32,10 @@ import torch from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base.conn import StateType from sglang.srt.disaggregation.common.conn import CommonKVManager +from sglang.srt.disaggregation.common.staging_buffer import ( + compute_grid_segments, + staging_grid_tokens, +) from sglang.srt.disaggregation.utils import ( FAKE_BOOTSTRAP_HOST, DisaggregationMode, @@ -334,6 +338,8 @@ class PrefillBootstrapQueue: decode_prefix_len = req.disagg_kv_sender.pop_decode_prefix_len() num_kv_indices = len(req.origin_input_ids) req.start_send_idx = decode_prefix_len + # Base of the staging chunk grid (suffix-relative send coordinates). + req.disagg_decode_prefix_len = 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, @@ -1075,16 +1081,26 @@ class SchedulerDisaggregationPrefillMixin: running_batch.batch_is_full = False def maybe_send_cached_prefix_chunk(self: Scheduler, req: Req) -> None: - # Only bootstrap-finalized requests; staging excluded. - if ( - not envs.SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX.get() - or self.enable_staging - or req.pending_bootstrap - ): + if not envs.SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX.get(): + return + + # Staging sends into positional grid slots, so the early-send boundary + # must stay stable across the request's batches: snapshot the at-rest + # prefix on the first batch. Non-staging reads the live prefix. + if self.enable_staging and req.early_send_prefix_end is None: + req.early_send_prefix_end = max( + 0, len(req.prefix_indices) - req.host_hit_length + ) + + if req.pending_bootstrap: return # Device-resident prefix only; page-aligned so start_send_idx stays exact. - cached_end = len(req.prefix_indices) - req.host_hit_length + cached_end = ( + req.early_send_prefix_end + if self.enable_staging + else len(req.prefix_indices) - req.host_hit_length + ) if cached_end <= req.start_send_idx: return if cached_end % self.token_to_kv_pool_allocator.page_size != 0: @@ -1124,6 +1140,15 @@ class SchedulerDisaggregationPrefillMixin: 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 + if self.enable_staging: + # Staging identifies chunks positionally against a uniform + # prefetched grid, so non-last sends must end on a grid + # boundary; the remainder rides with the next send. + grid_tokens = staging_grid_tokens( + self.server_args.chunked_prefill_size, page_size + ) + base = req.disagg_decode_prefix_len + end_idx = base + ((end_idx - base) // grid_tokens) * grid_tokens if end_idx < start_idx: logger.debug( @@ -1238,17 +1263,35 @@ class SchedulerDisaggregationPrefillMixin: payloads[st]() if st in payloads else None for st in state_types ] - kv_indices = self.req_to_token_pool.req_to_token[ - req.req_pool_idx, start_idx:end_idx - ] - page_indices = kv_to_page_indices(kv_indices, page_size) - if not req.disagg_kv_sender.should_send_kv_chunk(len(page_indices), last_chunk): - return - req.disagg_kv_sender.send( - page_indices, - state_indices, - num_kv_tokens=end_idx - start_idx, - ) + if self.enable_staging: + # One sender.send per grid slot; the sender's cumulative page + # counter marks only the final sub-send of the final chunk as + # is_last, routing aux/state correctly. + segments = compute_grid_segments( + start_idx, + end_idx, + req.disagg_decode_prefix_len, + staging_grid_tokens(self.server_args.chunked_prefill_size, page_size), + ) + else: + segments = [(start_idx, end_idx)] + + for seg_start, seg_end in segments: + is_final_segment = seg_end == end_idx + kv_indices = self.req_to_token_pool.req_to_token[ + req.req_pool_idx, seg_start:seg_end + ] + page_indices = kv_to_page_indices(kv_indices, page_size) + segment_is_last = last_chunk and is_final_segment + if not req.disagg_kv_sender.should_send_kv_chunk( + len(page_indices), segment_is_last + ): + continue + req.disagg_kv_sender.send( + page_indices, + state_indices if segment_is_last else None, + num_kv_tokens=seg_end - seg_start, + ) req.start_send_idx = end_idx def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None: @@ -1260,6 +1303,8 @@ class SchedulerDisaggregationPrefillMixin: req.output_ids = array("q") req.start_send_idx = 0 req.tmp_end_idx = -1 + req.disagg_decode_prefix_len = 0 + req.early_send_prefix_end = None req.hidden_states_tensor = None req.output_dsa_topk_indices = None req.pending_bootstrap = True diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 237f293ec..e6f35fac7 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -218,6 +218,13 @@ def poll_and_all_reduce_with_staging( receivers = [dr.kv_receiver for dr in decode_reqs] raw_polls = _poll_with_failure_injection(receivers) for i, decode_req in enumerate(decode_reqs): + if decode_req.kv_receiver.require_staging and staging_handler.is_failed( + decode_req + ): + # Staging completion timed out; KVPoll.Failed == 0 propagates + # through the MIN all_reduce. + raw_polls[i] = int(KVPoll.Failed) + continue if raw_polls[i] == int(KVPoll.Success): if decode_req.kv_receiver.require_staging and not staging_handler.is_done( decode_req diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 7ceeabf77..15e9f8805 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1149,6 +1149,12 @@ class Req(ReqDllmMixin): # This is because kv is not ready in `process_prefill_chunk`. # We use `tmp_end_idx` to store the end index of the kv cache to send. self.tmp_end_idx: int = -1 + # Decode-side cached-prefix length; base of the staging chunk grid + # (start_send_idx starts here but advances with every send). + self.disagg_decode_prefix_len: int = 0 + # At-rest device-resident prefix end, snapshotted on the request's + # first prefill batch; the cached-prefix early-send never goes past it. + self.early_send_prefix_end: Optional[int] = None 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. diff --git a/test/registered/disaggregation/test_disaggregation_different_tp.py b/test/registered/disaggregation/test_disaggregation_different_tp.py index a1972895d..2849c3c50 100644 --- a/test/registered/disaggregation/test_disaggregation_different_tp.py +++ b/test/registered/disaggregation/test_disaggregation_different_tp.py @@ -595,5 +595,104 @@ class TestDisaggregationGDNHybridHeteroTP(PDDisaggregationServerBase): self.assertGreater(metrics["score"], 0.60) +class TestDisaggregationStagingRadixPrefillLargerTP(PDDisaggregationServerBase): + """Prefill TP=4 -> Decode TP=2, staging + radix cache on both sides. + + The gsm8k few-shot preamble is a prefix shared by every request. With a + small chunked-prefill-size it spans several staging grid slots, so a + prefill radix hit bundles a multi-slot prefix into the first send and the + decode side reuses its own cached prefix -- the exact grid-split and + decode-prefix scatter-offset paths that a default (single-slot) prefix + never reaches. Unpatched, this configuration corrupts KV or wedges; the + gsm8k score guards against both. + """ + + # Small enough that the shared gsm8k few-shot prefix (~900 tokens) spans + # several staging grid slots, so a radix hit exercises the grid-split path. + CHUNKED_PREFILL_SIZE = "256" + + @classmethod + def setUpClass(cls): + super().setUpClass() + envs.SGLANG_ENABLE_JIT_DEEPGEMM.set(False) + + cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST) + + cls.start_prefill() + cls.start_decode() + + cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill) + cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode) + + cls.launch_lb() + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "4", + "--chunked-prefill-size", + cls.CHUNKED_PREFILL_SIZE, + "--enable-metrics", + "--enable-request-time-stats-logging", + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + env = {**os.environ, **STAGING_ENV} + cls.process_prefill = popen_launch_pd_server( + cls.model, + cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + env=env, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "2", + "--base-gpu-id", + "4", + "--chunked-prefill-size", + cls.CHUNKED_PREFILL_SIZE, + "--disaggregation-decode-enable-radix-cache", + "--enable-metrics", + "--enable-request-time-stats-logging", + ] + decode_args += cls.transfer_backend + cls.rdma_devices + env = {**os.environ, **STAGING_ENV} + cls.process_decode = popen_launch_pd_server( + cls.model, + cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + env=env, + ) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + metrics = run_eval(args) + print(f"[Staging Radix PrefillLargerTP] Evaluation metrics: {metrics}") + self.assertGreater(metrics["score"], 0.60) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index 61c8d2f5c..763842ea2 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -466,6 +466,7 @@ class TestNixlTransferWorker(CustomTestCase): mgr.req_to_decode_prefix_len = {room: 4} mgr.enable_staging = False mgr._staging_ctx = None + mgr._staging_outstanding = defaultdict(int) mgr.is_mla_backend = False mgr.is_hybrid_mla_backend = False mgr.attn_tp_size = 1