diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 41eb16383..556fb6cd2 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -361,6 +361,9 @@ class HiCacheController: self.load_queue: List[CacheOperation] = [] self.write_queue: List[CacheOperation] = [] self.ack_load_queue: List[HiCacheAck] = [] + # Set by the scheduler to the forward stream; gates load-back H2D + # behind in-flight forwards (see start_loading). + self.load_fence_stream = None self.ack_write_queue: List[HiCacheAck] = [] self.l2_transfer_engine = L2TransferEngine(io_backend) @@ -922,6 +925,14 @@ class HiCacheController: producer_event = self.layer_done_counter.events[producer_id] producer_event.start_event.record() + if self.load_fence_stream is not None: + # in overlap scheduling, reclaimed pages might still be written by the forward thread + # therefore a fence is needed for loading thread to prevent memory corruption + # todo: it's possible to use a finer-grained fence + self.l2_transfer_engine.host_to_device_stream.wait_stream( + self.load_fence_stream + ) + completion = self.l2_transfer_engine.submit_host_to_device( self._l2_load_transfers(host_indices, device_indices, pool_transfers), start_event=producer_event.start_event, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index d6f547328..5b65d4426 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -577,6 +577,12 @@ class Scheduler( self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator self.disable_radix_cache = result.disable_radix_cache self.tree_cache = result.tree_cache + if self.enable_hierarchical_cache: + cache_controller = self.tree_cache.cache_controller + if cache_controller is not None: + cache_controller.load_fence_stream = ( + self.tp_worker.model_runner.forward_stream + ) self.emit_metrics_constants() self.maybe_init_hccl_dp_prewarm() diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index f5711f2f6..ff6f65593 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -230,6 +230,7 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): controller._num_tokens_by_pool.return_value = {} controller._transfer_num_bytes.return_value = 0 controller.l2_transfer_engine = mock.Mock() + controller.load_fence_stream = None completion = SimpleNamespace( start_event=object(), finish_event=object(), timing_enabled=False )