Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 43 additions & 1 deletion python/sglang/srt/managers/cache_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,6 +222,8 @@ class PrefetchAck:
# Number of hits in extra pools.
pool_hits: Optional[dict[str, int]] = None
completed_req: Optional[bool] = None
# Only emitted for a single-worker backend exception after I/O returns.
failed: bool = False


class StorageOperation:
Expand Down Expand Up @@ -249,6 +251,9 @@ def __init__(
self.stats_requested_tokens = 0
# Absolute token offset at which this storage-prefetched span starts.
self.storage_start = 0
# Terminal consumption is scheduler-owned, including cancelled/stale ops.
self.terminal_ack_consumed = False
self.terminal_outcome: Optional[str] = None

self.id = StorageOperation.counter
StorageOperation.counter += 1
Expand Down Expand Up @@ -1125,7 +1130,12 @@ def _page_transfer(self, operation: PrefetchOperation) -> int:
kv_derived_transfers,
)
except Exception:
if not get_memory().enable_unified_memory:
# The resident single-worker drain distinguishes failure
# from a miss and reclaims acknowledged progress plus tail.
if (
self._supports_local_prefetch_failure()
or not get_memory().enable_unified_memory
):
raise
logger.exception(
"HiCache prefetch transfer failed for request %s",
Expand Down Expand Up @@ -1196,17 +1206,36 @@ def prefetch_io_aux_func(self):
continue
if operation is None:
continue
failed = False
try:
self._page_transfer(operation)
except Exception:
# Local recovery does not change multi-rank collective ordering.
if not self._supports_local_prefetch_failure():
raise
logger.exception(
"HiCache prefetch read failed: %s", operation.request_id
)
failed = True
finally:
self.prefetch_sync_queue.put(
PrefetchAck(
rid=operation.request_id,
completed_req=True,
operation=operation,
failed=failed,
)
)

def _supports_local_prefetch_failure(self) -> bool:
# _create_sync_groups excludes single-rank groups. PP tickets also
# carry their own completion protocol, even before local allocation.
return getattr(self, "_supports_prefetch_failure_ack", False) and not (
self.prefetch_hits_sync_groups
or self.prefetch_completion_sync_groups
or getattr(self, "pp_prefetch_command_group", None) is not None
)

def prefetch_rate_limited(self) -> bool:
"""
Rate limit the prefetching operations to avoid overwhelming the storage backend.
Expand Down Expand Up @@ -1287,6 +1316,19 @@ def prefetch_thread_func(self):
operation
)
except Exception:
if self._supports_local_prefetch_failure():
logger.exception(
"HiCache prefetch query failed: %s", operation.request_id
)
self.prefetch_sync_queue.put(
PrefetchAck(
rid=operation.request_id,
operation=operation,
completed_req=True,
failed=True,
)
)
continue
if not get_memory().enable_unified_memory:
raise
logger.exception(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,13 @@ def __init__(
host_memory_mode: str = "cache",
):
startup_storage_backend = storage_backend
# Only the resident Unified FULL drain currently consumes failure ACKs.
# Keep legacy and multi-pool consumers on their existing protocol.
self._supports_prefetch_failure_ack = (
host_memory_mode == "cache"
and len(mem_pool_host.entries) == 1
and mem_pool_host.anchor_entry.name == PoolName.KV
)
self.extra_host_mem_release_queues: dict[PoolName, Queue[torch.Tensor]] = {}
self.pp_prefetch_command_group = None
self.pp_prefetch_command_thread = None
Expand Down
26 changes: 24 additions & 2 deletions python/sglang/srt/mem_cache/unified_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -2578,6 +2578,8 @@ def release_aborted_request(self, request: CacheRequestHandle) -> None:
return

info = self.ongoing_prefetch[request]
if info.operation.terminal_outcome is None:
info.operation.terminal_outcome = "CANCELLED"
if info.operation.host_indices is None:
self.cache_controller.terminate_prefetch(info.operation)
self.revoke_pending_prefetch(request)
Expand Down Expand Up @@ -2754,6 +2756,8 @@ def revoke_pending_prefetch(self, request: CacheRequestHandle) -> None:
info = self.ongoing_prefetch.get(request)
if info is None:
return
if info.operation.terminal_outcome is None:
info.operation.terminal_outcome = "MISS"
self._invalidate_absent_from_hit_query(info.operation)
# Every revoke path runs before the bounce alloc, so buffer mode
# holds no occupancy here; post-alloc aborts go through
Expand Down Expand Up @@ -2985,6 +2989,8 @@ def _drain_and_alloc_storage_hit():
def _drain_ack_prefetch():
for ack in _drain_queue(cc.ack_prefetch_queue, n_ack_prefetch):
operation = ack.operation
if operation.terminal_ack_consumed:
continue
info = self.ongoing_prefetch.get(operation.handle)
is_current = info is not None and info.operation is operation
if ack.completed_tokens is not None:
Expand All @@ -2998,11 +3004,27 @@ def _drain_ack_prefetch():
)
operation.pool_transfers_done = True
if ack.completed_req:
if is_current:
operation.terminal_ack_consumed = True
if operation.terminal_outcome is None:
operation.terminal_outcome = (
"FAILURE"
if ack.failed
else "SUCCESS"
if operation.completed_tokens > 0
else "MISS"
)
if is_current and ack.failed:
# The backend has returned/raised: abort frees only
# acknowledged pages; the terminal drain owns the tail.
self.release_aborted_request(operation.handle)
elif is_current:
# check_prefetch_progress() is not called for this rid yet.
# Let us insert the prefetch result into the radix tree.
self._handle_prefetch_result(operation)
if operation.ack_releases_incomplete_host_indices:
if (
operation.host_indices is not None
and operation.ack_releases_incomplete_host_indices
):
cc.append_host_mem_release(
operation.host_indices[operation.completed_tokens :],
(
Expand Down
Loading
Loading