Skip to content
Merged
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
24 changes: 4 additions & 20 deletions python/sglang/srt/mem_cache/unified_cache/storage_attachment.py
Original file line number Diff line number Diff line change
Expand Up @@ -421,31 +421,15 @@ def _release_pending_storage_ops(self) -> None:
cache.revoke_pending_prefetch(handle)
continue
completed_tokens, _ = controller.terminate_prefetch(info.operation)
del cache.ongoing_prefetch[handle]
if info.anchor_lock_params is not None:
cache.dec_host_lock_ref(
info.anchor_node_id, info.anchor_lock_params
)
if cache.buffer_pipeline is not None:
cache.buffer_pipeline.pop_prefix_ctx(handle)
cache.buffer_pipeline.release_anchor_lock(handle)
controller.append_host_mem_release(
host_indices=info.host_indices[:completed_tokens],
extra_pools=(
cache._retire_ongoing_prefetch(
handle,
info.host_indices[:completed_tokens],
(
[x for xfers in info.comp_xfers.values() for x in xfers]
if info.operation.pool_transfers_done
else None
),
)
controller.prefetch_tokens_occupied = max(
0,
controller.prefetch_tokens_occupied
- cache._prefetch_occupied_span(
info.prefetch_key,
info.host_indices,
operation=info.operation,
),
)
except Exception:
logger.exception("Failed to release pending prefetch %s", handle.rid)
cache.ongoing_prefetch.pop(handle, None)
Expand Down
121 changes: 44 additions & 77 deletions python/sglang/srt/mem_cache/unified_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -2252,14 +2252,7 @@ def _handle_prefetch_result(self, operation: PrefetchOperation) -> None:
# All PP/TP ranks will get the same `min_completed_tokens`, because `completed_tokens`
# and `pool_hits` in their operations are same. No need to sync cross-rank here.
if not self._check_hybrid_prefetch_result(
request,
operation,
completed_tokens,
hash_value,
host_indices,
last_host_node_id,
anchor_lock_params,
prefetch_key,
request, operation, hash_value, host_indices
):
# Hybrid all-or-nothing check failed; result already discarded.
return
Expand Down Expand Up @@ -2353,12 +2346,8 @@ def _check_hybrid_prefetch_result(
self,
request: CacheRequestHandle,
operation: PrefetchOperation,
completed_tokens: int,
hash_value: list[str],
host_indices: torch.Tensor,
last_host_node_id: NodeId,
anchor_lock_params: DecLockRefParams,
prefetch_key: RadixKey,
) -> bool:
"""Decide the length of usable prefix.

Expand Down Expand Up @@ -2421,25 +2410,15 @@ def _check_hybrid_prefetch_result(
self.storage_existence_cache.invalidate_beyond(
PoolName.KV, invalidation_hashes, keep_pages=keep_pages
)
# The controller's prefetch IO thread already releases the untransferred
# tail (host_indices[completed_tokens:])
self.cache_controller.append_host_mem_release(
host_indices=host_indices[:completed_tokens],
extra_pools=pool_transfers if operation.pool_transfers_done else None,
)
self._finish_storage_prefetch(
request, fulfilled_tokens=0, reason="storage_transfer"
)
if anchor_lock_params is not None:
self.dec_host_lock_ref(last_host_node_id, anchor_lock_params)
if self.buffer_pipeline is not None:
self.buffer_pipeline.pop_prefix_ctx(request)
self.buffer_pipeline.release_anchor_lock(request)
del self.ongoing_prefetch[request]
self.cache_controller.prefetch_tokens_occupied -= (
self._prefetch_occupied_span(
prefetch_key, host_indices, operation=operation
)
# The controller's prefetch IO thread already releases the untransferred
# tail (host_indices[completed_tokens:])
self._retire_ongoing_prefetch(
request,
host_indices[:completed_tokens],
pool_transfers if operation.pool_transfers_done else None,
)
self.prefetch_loaded_tokens_by_reqid[request] = 0
self.prefetch_loaded_storage_start_by_reqid.pop(request, None)
Expand Down Expand Up @@ -2598,35 +2577,44 @@ def release_aborted_request(self, request: CacheRequestHandle) -> None:
if request not in self.ongoing_prefetch:
return

(
last_host_node_id,
prefetch_key,
host_indices,
operation,
anchor_lock_params,
comp_xfers,
) = self.ongoing_prefetch[request]
if operation.host_indices is None:
self.cache_controller.terminate_prefetch(operation)
info = self.ongoing_prefetch[request]
if info.operation.host_indices is None:
self.cache_controller.terminate_prefetch(info.operation)
self.revoke_pending_prefetch(request)
return

completed_tokens, _ = self.cache_controller.terminate_prefetch(operation)
if anchor_lock_params is not None:
self.dec_host_lock_ref(last_host_node_id, anchor_lock_params)
del self.ongoing_prefetch[request]
completed_tokens, _ = self.cache_controller.terminate_prefetch(info.operation)
self._retire_ongoing_prefetch(
request,
info.host_indices[:completed_tokens],
(
[x for xfers in info.comp_xfers.values() for x in xfers]
if info.operation.pool_transfers_done
else None
),
)

def _retire_ongoing_prefetch(
self,
request: CacheRequestHandle,
host_indices: Optional[torch.Tensor],
extra_pools: Optional[list[PoolTransfer]],
) -> None:
"""Drop an in-flight prefetch, releasing ``host_indices`` and ``extra_pools``."""
info = self.ongoing_prefetch.pop(request)
if info.anchor_lock_params is not None:
self.dec_host_lock_ref(info.anchor_node_id, info.anchor_lock_params)
if self.buffer_pipeline is not None:
self.buffer_pipeline.pop_prefix_ctx(request)
self.buffer_pipeline.release_anchor_lock(request)
pool_transfers = [x for xfers in comp_xfers.values() for x in xfers]
self.cache_controller.append_host_mem_release(
host_indices=host_indices[:completed_tokens],
extra_pools=pool_transfers if operation.pool_transfers_done else None,
)
# Buffer mode granted occupancy at hit-alloc, sized to the bounce;
# cache mode reserved the requested span at enqueue.
self.cache_controller.prefetch_tokens_occupied -= self._prefetch_occupied_span(
prefetch_key, host_indices, operation=operation
cc = self.cache_controller
cc.append_host_mem_release(host_indices=host_indices, extra_pools=extra_pools)
cc.prefetch_tokens_occupied = max(
0,
cc.prefetch_tokens_occupied
- self._prefetch_occupied_span(
info.prefetch_key, info.host_indices, operation=info.operation
),
)

def _invalidate_absent_from_hit_query(self, operation) -> None:
Expand Down Expand Up @@ -2762,38 +2750,17 @@ def _trim_buffer_prefetch_full_head(
return info, hit_tokens - trim_tokens, original_hit_tokens

def revoke_pending_prefetch(self, request: CacheRequestHandle) -> None:
info = self.ongoing_prefetch.pop(request, None)
self._finish_storage_prefetch(request, fulfilled_tokens=0, reason="dropped")
info = self.ongoing_prefetch.get(request)
if info is None:
return
(
last_host_node_id,
prefetch_key,
_host_indices,
operation,
anchor_lock_params,
comp_xfers,
) = info
self._invalidate_absent_from_hit_query(operation)
if self.buffer_pipeline is not None:
self.buffer_pipeline.pop_prefix_ctx(request)
self.buffer_pipeline.release_anchor_lock(request)
cc = self.cache_controller
cc.append_host_mem_release(
extra_pools=[x for xfers in comp_xfers.values() for x in xfers]
)
if anchor_lock_params is not None:
self.dec_host_lock_ref(last_host_node_id, anchor_lock_params)
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
# release_aborted_request instead.
assert _host_indices is None or self.host_memory_mode != "buffer_only"
cc.prefetch_tokens_occupied = max(
0,
cc.prefetch_tokens_occupied
- self._prefetch_occupied_span(
prefetch_key, _host_indices, operation=operation
),
assert info.host_indices is None or self.host_memory_mode != "buffer_only"
self._retire_ongoing_prefetch(
request, None, [x for xfers in info.comp_xfers.values() for x in xfers]
)

def _drain_storage_control_queues_impl(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10799,6 +10799,9 @@ def test_storage_cleanup_releases_buffer_prefetch_anchor(self):
)
)
)
cache._retire_ongoing_prefetch.side_effect = lambda *args: (
UnifiedRadixCache._retire_ongoing_prefetch(cache, *args)
)
controller = cache.cache_controller
controller.terminate_prefetch.return_value = (4, None)
controller.prefetch_tokens_occupied = 12
Expand Down
Loading