From 52384cab4f937ccdbc6562ac2ef6095a5865c956 Mon Sep 17 00:00:00 2001 From: alphabetc1 <2508695655@qq.com> Date: Sun, 27 Sep 2026 21:35:26 +0800 Subject: [PATCH] refactor: retire in-flight storage prefetches through one helper --- .../unified_cache/storage_attachment.py | 24 +--- .../srt/mem_cache/unified_radix_cache.py | 115 +++++++----------- .../test_unified_radix_cache_unittest.py | 3 + 3 files changed, 49 insertions(+), 93 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py b/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py index 82ee8bd21e01..14a404aefcfc 100644 --- a/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py +++ b/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py @@ -384,26 +384,10 @@ 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=[ - x for xfers in info.comp_xfers.values() for x in xfers - ], - ) - controller.prefetch_tokens_occupied = max( - 0, - controller.prefetch_tokens_occupied - - cache._prefetch_occupied_span( - info.prefetch_key, info.host_indices - ), + cache._retire_ongoing_prefetch( + handle, + info.host_indices[:completed_tokens], + [x for xfers in info.comp_xfers.values() for x in xfers], ) except Exception: logger.exception("Failed to release pending prefetch %s", handle.rid) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index f9c0fcc04f5b..20590b748d9a 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2263,14 +2263,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 @@ -2364,12 +2357,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. @@ -2432,23 +2421,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) + # 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) @@ -2596,35 +2577,42 @@ 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 + 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), ) def _invalidate_absent_from_hit_query(self, operation) -> None: @@ -2754,36 +2742,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), + 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( diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 3a2214379065..1d0e5b937182 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -10501,6 +10501,9 @@ def test_storage_cleanup_releases_buffer_prefetch_anchor(self): cache._prefetch_occupied_span.side_effect = lambda key, indices: ( UnifiedRadixCache._prefetch_occupied_span(cache, key, indices) ) + 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