diff --git a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py index 27b33846f..bccdd809b 100644 --- a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py +++ b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py @@ -104,8 +104,22 @@ class DecodeKVCacheOffloadManager: self.ongoing_offload = {} self.ongoing_backup = {} self.offloaded_state = {} + self.offload_inflight = {} logger.info("Enable offload kv cache for decode side") + def _mark_offload_started(self, rid): + self.offload_inflight[rid] = self.offload_inflight.get(rid, 0) + 1 + + def _mark_offload_finished(self, rid): + count = self.offload_inflight.get(rid, 0) + if count <= 1: + self.offload_inflight.pop(rid, None) + else: + self.offload_inflight[rid] = count - 1 + + def _has_inflight_offload(self, rid): + return self.offload_inflight.get(rid, 0) > 0 + def offload_kv_cache(self, req) -> bool: """Offload incremental KV cache for decode side.""" @@ -153,9 +167,11 @@ class DecodeKVCacheOffloadManager: incremental_tokens = all_tokens[start:end] incremental_indices = token_indices[start:end] - # Early free prefill-offloaded GPU memory - if state.prefill_len > 0 and state.inc_len == 0: - self.token_to_kv_pool_allocator.free(token_indices[: state.prefill_len]) + # Prefill-aligned GPU slots are freed at request finish in + # _release_finished_req, NOT here. The decoding request + # continues to attend to those slots via req_to_token; freeing + # them mid-decode races with concurrent admission, which can + # reuse the slots and produce cross-pollinated KV reads. # Asynchronously offload incremental KV cache from device to host self.request_counter += 1 @@ -168,6 +184,7 @@ class DecodeKVCacheOffloadManager: logger.error(f"Not enough host memory for request {req.rid}") return False + self._mark_offload_started(req.rid) self.ongoing_offload[ack_id] = ( req, host_indices, @@ -214,14 +231,7 @@ class DecodeKVCacheOffloadManager: end, ) = self.ongoing_offload.pop(ack_id) - if req.finished(): - self._release_finished_req(req, start) - else: - kv_indices = self.req_to_token_pool.req_to_token[ - req.req_pool_idx, start:end - ] - self.token_to_kv_pool_allocator.free(kv_indices) - + self._mark_offload_finished(req.rid) prior_hash = ( self.offloaded_state[req.rid].last_hash if req.rid in self.offloaded_state @@ -232,10 +242,34 @@ class DecodeKVCacheOffloadManager: ) if req.rid in self.offloaded_state: self.offloaded_state[req.rid].last_hash = last_hash + + if req.finished() and not self._has_inflight_offload(req.rid): + state = self.offloaded_state.get(req.rid) + start_offset = state.prefill_len if state is not None else start + self._release_finished_req(req, start_offset) finish_count -= 1 def _release_finished_req(self, req: Req, start_offset: int): + # Defensive guard: ReqToTokenPool.free sets req_pool_idx to None, + # so a previously-released request must be skipped here to avoid + # non-idempotent side effects (e.g. tree_cache.protected_size_ + # double-decrement, host pool double-free). + if req.req_pool_idx is None or req.req_pool_idx == -1: + return + kv_committed_len = req.pop_committed_kv_cache() + + # Free the prefill-aligned slots. Previously this was done + # eagerly in offload_kv_cache (mid-decode), which raced with + # concurrent admission. Now consolidated here at request + # finish, where the request is guaranteed to no longer attend + # to those slots. + state = self.offloaded_state.get(req.rid) + if state is not None and state.prefill_len > 0: + prefill_indices = self.req_to_token_pool.req_to_token[ + req.req_pool_idx, : state.prefill_len + ] + self.token_to_kv_pool_allocator.free(prefill_indices) start = start_offset end = kv_committed_len # Free the incremental part of the request (DSA-aware) @@ -296,7 +330,9 @@ class DecodeKVCacheOffloadManager: def finalize_release_on_finish(self, req: Req): """Free any remaining tail KV that was not offloaded due to non-aligned length.""" - if req.req_pool_idx == -1: + # ReqToTokenPool.free sets req_pool_idx to None on release, so + # guard against both sentinels here. + if req.req_pool_idx is None or req.req_pool_idx == -1: return state = self.offloaded_state.get(req.rid) if state is None: @@ -305,13 +341,13 @@ class DecodeKVCacheOffloadManager: else: prefill_len = state.prefill_len inc_len = state.inc_len - # If no incremental offload ever happened, the prefill-aligned part was never freed. - # Free the prefill portion on request finish to avoid leaks. - if prefill_len > 0 and inc_len == 0: - token_indices = self.req_to_token_pool.req_to_token[req.req_pool_idx] - self.token_to_kv_pool_allocator.free(token_indices[:prefill_len]) - logger.info( - f"Finalize release: freed prefill-aligned KV for req {req.rid}, len:{prefill_len}" + # Prefill-aligned slots are freed by _release_finished_req. Make + # sure state exists so it can find prefill_len. + if state is None: + self.offloaded_state[req.rid] = OffloadedState( + prefill_len=prefill_len, inc_len=0, last_hash=None ) - start_offset = prefill_len + inc_len + if self._has_inflight_offload(req.rid): + return + start_offset = prefill_len self._release_finished_req(req, start_offset) diff --git a/test/registered/disaggregation/test_disaggregation_decode_offload.py b/test/registered/disaggregation/test_disaggregation_decode_offload.py index 9a8baba68..058154021 100644 --- a/test/registered/disaggregation/test_disaggregation_decode_offload.py +++ b/test/registered/disaggregation/test_disaggregation_decode_offload.py @@ -21,7 +21,6 @@ register_cuda_ci( est_time=600, stage="base-b", runner_config="2-gpu-large", - disabled="Temporarily disable the flaky test.", ) diff --git a/test/registered/disaggregation/test_specv2_kvcache_offloading.py b/test/registered/disaggregation/test_specv2_kvcache_offloading.py index d9876dbcd..0cd5c77bd 100644 --- a/test/registered/disaggregation/test_specv2_kvcache_offloading.py +++ b/test/registered/disaggregation/test_specv2_kvcache_offloading.py @@ -15,6 +15,7 @@ import torch from sglang.srt.disaggregation.decode_kvcache_offload_manager import ( DecodeKVCacheOffloadManager, ) +from sglang.srt.disaggregation.kv_events import OffloadedState from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small") @@ -77,10 +78,18 @@ def _make_manager(pool_size: int, page_size: int = 1): manager.page_size = page_size manager.tree_cache = tree_cache manager.offloaded_state = {} + manager.ongoing_offload = {} + manager.ongoing_backup = {} + manager.offload_inflight = {} return manager, freed_indices +class _FinishedEvent: + def synchronize(self): + pass + + class TestReleaseFinishedReq(unittest.TestCase): """Tests for _release_finished_req overallocation cleanup.""" @@ -176,6 +185,232 @@ class TestReleaseFinishedReq(unittest.TestCase): self.assertEqual(manager.tree_cache.protected_size_, 5) + def test_release_finished_req_frees_prefill_when_state_present(self): + """ + When offloaded_state[rid].prefill_len > 0, _release_finished_req must + free the prefill-aligned slots in addition to the committed range. + + This is the consolidated free path that replaces the eager free that + previously happened in offload_kv_cache (which raced with concurrent + admission and produced cross-pollinated KV reads). + """ + manager, freed = _make_manager(pool_size=32) + rid = "req-prefill-present" + req = _make_mock_req( + req_pool_idx=0, + kv_committed_len=20, + kv_allocated_len=20, + rid=rid, + ) + manager.offloaded_state[rid] = OffloadedState( + prefill_len=8, inc_len=0, last_hash=None + ) + + manager._release_finished_req(req, start_offset=8) + + # Two frees in order: prefill [0:8] then committed [8:20]. + self.assertEqual(len(freed), 2) + expected_prefill = torch.arange(0, 8, dtype=torch.int64) + expected_committed = torch.arange(8, 20, dtype=torch.int64) + self.assertTrue(torch.equal(freed[0], expected_prefill)) + self.assertTrue(torch.equal(freed[1], expected_committed)) + # State entry is removed at the end of _release_finished_req. + self.assertNotIn(rid, manager.offloaded_state) + + def test_release_finished_req_skips_prefill_free_when_prefill_len_zero(self): + """ + When state exists but prefill_len == 0 (request shorter than page_size, + so no prefill chunk was ever offloaded), no prefill-aligned free is + emitted. + """ + manager, freed = _make_manager(pool_size=32) + rid = "req-prefill-zero" + req = _make_mock_req( + req_pool_idx=0, + kv_committed_len=10, + kv_allocated_len=10, + rid=rid, + ) + manager.offloaded_state[rid] = OffloadedState( + prefill_len=0, inc_len=0, last_hash=None + ) + + manager._release_finished_req(req, start_offset=0) + + # Only the committed range [0:10] is freed. + self.assertEqual(len(freed), 1) + expected_committed = torch.arange(0, 10, dtype=torch.int64) + self.assertTrue(torch.equal(freed[0], expected_committed)) + + def test_finalize_release_creates_state_so_prefill_is_freed(self): + """ + finalize_release_on_finish handles the case where no incremental + offload ever ran (offloaded_state is empty). It must materialize an + OffloadedState with the correct prefill_len so that the consolidated + free site in _release_finished_req can locate and free those slots. + """ + page_size = 4 + manager, freed = _make_manager(pool_size=32, page_size=page_size) + rid = "req-finalize-no-state" + req = _make_mock_req( + req_pool_idx=0, + kv_committed_len=13, + kv_allocated_len=13, + rid=rid, + ) + # 12 input tokens => prefill_len = 12 // 4 * 4 = 12 + req.origin_input_ids = list(range(12)) + + manager.finalize_release_on_finish(req) + + # finalize creates state, then _release_finished_req frees: + # prefill [0:12] then committed [12:13]. + self.assertEqual(len(freed), 2) + expected_prefill = torch.arange(0, 12, dtype=torch.int64) + expected_committed = torch.arange(12, 13, dtype=torch.int64) + self.assertTrue(torch.equal(freed[0], expected_prefill)) + self.assertTrue(torch.equal(freed[1], expected_committed)) + # State is deleted by _release_finished_req on the way out. + self.assertNotIn(rid, manager.offloaded_state) + + def test_unfinished_offload_ack_does_not_free_incremental_slots(self): + manager, freed = _make_manager(pool_size=32) + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=1 + ) + req.finished.return_value = False + manager.offloaded_state[req.rid] = OffloadedState( + prefill_len=4, inc_len=4, last_hash=None + ) + manager.offload_inflight[req.rid] = 1 + manager.ongoing_offload[7] = ( + req, + torch.arange(4, 8, dtype=torch.int64), + [10, 11, 12, 13], + 0.0, + 4, + 8, + ) + manager.cache_controller = MagicMock() + manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [7])] + manager._trigger_backup = MagicMock(return_value="last_hash") + + manager._check_offload_progress(1) + + self.assertEqual(freed, []) + manager.req_to_token_pool.free.assert_not_called() + self.assertNotIn(req.rid, manager.offload_inflight) + + def test_offload_kv_cache_tracks_inflight_write_until_ack(self): + manager, freed = _make_manager(pool_size=32, page_size=4) + manager.cache_controller = MagicMock() + manager.cache_controller.get_hash_str = MagicMock(return_value="prefill_hash") + manager.cache_controller.write = MagicMock( + return_value=torch.arange(4, 8, dtype=torch.int64) + ) + manager.decode_host_mem_pool = MagicMock() + manager.request_counter = 0 + manager.offload_stride = 4 + + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=5 + ) + req.origin_input_ids = [0, 1, 2, 3] + req.output_ids = [4, 5, 6, 7, 8] + req.finished.return_value = False + + did_offload = manager.offload_kv_cache(req) + + self.assertTrue(did_offload) + self.assertEqual(manager.offload_inflight[req.rid], 1) + self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4) + manager.cache_controller.write.assert_called_once() + + manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [1])] + manager._trigger_backup = MagicMock(return_value="last_hash") + + manager._check_offload_progress(1) + + self.assertEqual(freed, []) + self.assertNotIn(req.rid, manager.offload_inflight) + + def test_finalize_release_defers_while_offload_is_in_flight(self): + manager, freed = _make_manager(pool_size=32) + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=2 + ) + manager.offloaded_state[req.rid] = OffloadedState( + prefill_len=4, inc_len=8, last_hash=None + ) + manager.offload_inflight[req.rid] = 1 + + manager.finalize_release_on_finish(req) + + self.assertEqual(freed, []) + manager.req_to_token_pool.free.assert_not_called() + self.assertIn(req.rid, manager.offloaded_state) + + def test_finished_offload_ack_waits_for_other_inflight_writes(self): + manager, freed = _make_manager(pool_size=32) + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=3 + ) + req.finished.return_value = True + manager.offloaded_state[req.rid] = OffloadedState( + prefill_len=4, inc_len=8, last_hash=None + ) + manager.offload_inflight[req.rid] = 2 + manager.ongoing_offload[8] = ( + req, + torch.arange(4, 8, dtype=torch.int64), + [10, 11, 12, 13], + 0.0, + 4, + 8, + ) + manager.cache_controller = MagicMock() + manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [8])] + manager._trigger_backup = MagicMock(return_value="last_hash") + + manager._check_offload_progress(1) + + self.assertEqual(freed, []) + manager.req_to_token_pool.free.assert_not_called() + self.assertEqual(manager.offload_inflight[req.rid], 1) + + def test_finished_request_releases_all_committed_slots_after_last_offload_ack( + self, + ): + manager, freed = _make_manager(pool_size=32) + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=4 + ) + req.finished.return_value = True + manager.offloaded_state[req.rid] = OffloadedState( + prefill_len=4, inc_len=8, last_hash=None + ) + manager.offload_inflight[req.rid] = 1 + manager.ongoing_offload[9] = ( + req, + torch.arange(8, 12, dtype=torch.int64), + [14, 15, 16, 17], + 0.0, + 8, + 12, + ) + manager.cache_controller = MagicMock() + manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [9])] + manager._trigger_backup = MagicMock(return_value="last_hash") + + manager._check_offload_progress(1) + + self.assertEqual(len(freed), 2) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 4, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[1], torch.arange(4, 20, dtype=torch.int64))) + manager.req_to_token_pool.free.assert_called_once_with(req) + self.assertNotIn(req.rid, manager.offloaded_state) + self.assertNotIn(req.rid, manager.offload_inflight) + if __name__ == "__main__": unittest.main()