diff --git a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py index b8ad4fc1a..354ae4039 100644 --- a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py +++ b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py @@ -5,6 +5,7 @@ import logging import threading import time from typing import TYPE_CHECKING +from weakref import WeakKeyDictionary as WeakKeyDict import torch @@ -93,12 +94,10 @@ class DecodeKVCacheOffloadManager: self.ongoing_offload = {} self.ongoing_backup = {} - # A caller may reuse a rid as soon as the previous response finishes, - # while that request's asynchronous D2H copy is still in flight. Key - # lifecycle state by the Req instance so a late ack cannot mutate the - # new request's state. - self.offloaded_state: dict[Req, OffloadedState] = {} - self.offload_inflight: dict[Req, int] = {} + # Keyed by Req identity (rids can be reused while a D2H copy is still + # in flight); weak keys so a dropped Req is never pinned here. + self.offloaded_state: WeakKeyDict[Req, OffloadedState] = WeakKeyDict() + self.offload_inflight: WeakKeyDict[Req, int] = WeakKeyDict() logger.info("Enable offload kv cache for decode side") def release_host_resources(self) -> None: @@ -117,6 +116,10 @@ class DecodeKVCacheOffloadManager: def _has_inflight_offload(self, req: Req): return self.offload_inflight.get(req, 0) > 0 + def _prefill_offloaded_len(self, req: Req) -> int: + # Page-aligned prompt length; the prefill instance offloaded this part. + return len(req.origin_input_ids) // self.page_size * self.page_size + def offload_kv_cache(self, req) -> bool: """Offload incremental KV cache for decode side.""" @@ -132,9 +135,7 @@ class DecodeKVCacheOffloadManager: # Prefill side offloads page-aligned origin_input_ids, decode side offloads the incremental part all_tokens = req.origin_input_ids + req.output_ids[:-1] - prefill_offloaded_len = ( - len(req.origin_input_ids) // self.page_size * self.page_size - ) + prefill_offloaded_len = self._prefill_offloaded_len(req) state = self.offloaded_state.get(req) if state is None: prefill_hashes = self._compute_prefix_hash( @@ -143,13 +144,9 @@ class DecodeKVCacheOffloadManager: last_prefill_hash = ( prefill_hashes[-1] if prefill_offloaded_len > 0 else None ) - state = OffloadedState( - prefill_len=prefill_offloaded_len, - inc_len=0, - last_hash=last_prefill_hash, - ) + state = OffloadedState(last_hash=last_prefill_hash) self.offloaded_state[req] = state - incremental_total = len(all_tokens) - state.prefill_len + incremental_total = len(all_tokens) - prefill_offloaded_len incremental_new = incremental_total - state.inc_len incremental_aligned_len = ( incremental_new // self.offload_stride * self.offload_stride @@ -159,7 +156,7 @@ class DecodeKVCacheOffloadManager: return False # Extract incremental tokens and indices for the newly available chunk - start = state.prefill_len + state.inc_len + start = prefill_offloaded_len + state.inc_len end = start + incremental_aligned_len incremental_tokens = all_tokens[start:end] incremental_indices = token_indices[start:end] @@ -187,8 +184,6 @@ class DecodeKVCacheOffloadManager: host_indices, incremental_tokens, time.time(), - start, - end, ) state.inc_len += incremental_aligned_len return True @@ -224,8 +219,6 @@ class DecodeKVCacheOffloadManager: host_indices, incremental_tokens, start_time, - start, - end, ) = self.ongoing_offload.pop(ack_id) self._mark_offload_finished(req) @@ -241,12 +234,10 @@ class DecodeKVCacheOffloadManager: self.offloaded_state[req].last_hash = last_hash if req.finished() and not self._has_inflight_offload(req): - state = self.offloaded_state.get(req) - start_offset = state.prefill_len if state is not None else start - self._release_finished_req(req, start_offset) + self._release_finished_req(req) finish_count -= 1 - def _release_finished_req(self, req: Req, start_offset: int): + def _release_finished_req(self, req: Req): # 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_ @@ -256,18 +247,15 @@ class DecodeKVCacheOffloadManager: kv_committed_len = req.effective_kv_committed_len() - # 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) - if state is not None and state.prefill_len > 0: + # Prefill-aligned slots are freed only here, at request finish; freeing + # them mid-decode races with concurrent admission over live slots. + prefill_len = self._prefill_offloaded_len(req) + if prefill_len > 0: prefill_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, : state.prefill_len + req.kv.req_pool_idx, :prefill_len ] self.token_to_kv_pool_allocator.free(prefill_indices) - start = start_offset + start = prefill_len end = kv_committed_len # Free the incremental part of the request (DSA-aware) kv_indices = self.req_to_token_pool.req_to_token[req.kv.req_pool_idx, start:end] @@ -327,24 +315,6 @@ class DecodeKVCacheOffloadManager: def finalize_release_on_finish(self, req: Req): """Free any remaining tail KV that was not offloaded due to non-aligned length.""" - # ReqToTokenPool.free sets req_pool_idx to None on release, so - # guard against both sentinels here. - if req.kv.req_pool_idx is None or req.kv.req_pool_idx == -1: - return - state = self.offloaded_state.get(req) - if state is None: - prefill_len = len(req.origin_input_ids) // self.page_size * self.page_size - inc_len = 0 - else: - prefill_len = state.prefill_len - inc_len = state.inc_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] = OffloadedState( - prefill_len=prefill_len, inc_len=0, last_hash=None - ) if self._has_inflight_offload(req): return - start_offset = prefill_len - self._release_finished_req(req, start_offset) + self._release_finished_req(req) diff --git a/python/sglang/srt/disaggregation/kv_events.py b/python/sglang/srt/disaggregation/kv_events.py index b03b1efd8..1672cbf8a 100644 --- a/python/sglang/srt/disaggregation/kv_events.py +++ b/python/sglang/srt/disaggregation/kv_events.py @@ -263,21 +263,13 @@ class BlockStoredMetadata(msgspec.Struct, omit_defaults=True, gc=False): cache_salt: str -class OffloadedState: - """ - OffloadedState represents the state of a KV cache block offloaded to the hicache. +class OffloadedState(msgspec.Struct): + """Decode-side offload progress for one request, keyed by Req in the manager.""" - - prefill_len (int): The length of the prefill part of the KV cache block. - - inc_len (int): The length of the incremental part of the KV cache block. - - last_hash (Optional[str]): The hash of the last token in the KV cache block. - """ - - def __init__( - self, prefill_len: int, inc_len: int = 0, last_hash: Optional[str] = None - ): - self.prefill_len = prefill_len - self.inc_len = inc_len - self.last_hash = last_hash + # Decode-incremental length already submitted for D2H offload. + inc_len: int = 0 + # Tail of the page hash chain, extended as each offloaded chunk is backed up. + last_hash: Optional[str] = None class BlockStored(KVCacheEvent): diff --git a/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py b/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py index 5700987a6..0268eef4f 100644 --- a/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py +++ b/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py @@ -7,8 +7,10 @@ are correctly freed when a request finishes, preventing GPU memory leaks. Requires: torch, sglang (run in an environment with sglang installed) """ +import gc import unittest from unittest.mock import MagicMock +from weakref import WeakKeyDictionary as WeakKeyDict import torch @@ -29,10 +31,12 @@ def _make_mock_req( kv_allocated_len: int, prefix_indices_len: int = 0, rid: int = 0, + origin_len: int = 0, ): """Create a mock Req with the KV cache state needed for testing.""" req = MagicMock() req.rid = rid + req.origin_input_ids = list(range(origin_len)) req.kv = ReqKvInfo( req_pool_idx=req_pool_idx, kv_committed_len=kv_committed_len, @@ -67,10 +71,10 @@ def _make_manager(pool_size: int, page_size: int = 1): manager.token_to_kv_pool_allocator = allocator manager.page_size = page_size manager.tree_cache = tree_cache - manager.offloaded_state = {} + manager.offloaded_state = WeakKeyDict() manager.ongoing_offload = {} manager.ongoing_backup = {} - manager.offload_inflight = {} + manager.offload_inflight = WeakKeyDict() return manager, freed_indices @@ -90,15 +94,15 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, # no overallocation + origin_len=8, ) - prefill_offloaded_len = 8 - manager._release_finished_req(req, prefill_offloaded_len) + manager._release_finished_req(req) - # Only one free call: the committed range [8:20] - self.assertEqual(len(freed), 1) - expected = torch.arange(8, 20, dtype=torch.int64) - self.assertTrue(torch.equal(freed[0], expected)) + # Prefill [0:8] and committed [8:20]; no overalloc free. + self.assertEqual(len(freed), 2) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[1], torch.arange(8, 20, dtype=torch.int64))) manager.req_to_token_pool.free.assert_called_once_with(req) def test_with_overallocation(self): @@ -108,17 +112,16 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=20, kv_allocated_len=28, # 8 over-allocated slots + origin_len=8, ) - prefill_offloaded_len = 8 - manager._release_finished_req(req, prefill_offloaded_len) + manager._release_finished_req(req) - # Two free calls: committed [8:20] and overallocated [20:28] - self.assertEqual(len(freed), 2) - expected_committed = torch.arange(8, 20, dtype=torch.int64) - expected_overalloc = torch.arange(20, 28, dtype=torch.int64) - self.assertTrue(torch.equal(freed[0], expected_committed)) - self.assertTrue(torch.equal(freed[1], expected_overalloc)) + # Prefill [0:8], committed [8:20], overallocated [20:28]. + self.assertEqual(len(freed), 3) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[1], torch.arange(8, 20, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[2], torch.arange(20, 28, dtype=torch.int64))) manager.req_to_token_pool.free.assert_called_once_with(req) def test_overallocation_with_page_alignment(self): @@ -129,18 +132,17 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=10, # not page-aligned kv_allocated_len=28, + origin_len=4, ) - prefill_offloaded_len = 4 - manager._release_finished_req(req, prefill_offloaded_len) + manager._release_finished_req(req) - # Committed range [4:10] - # Overallocated: start_p = ceil_align(10, 4) = 12, end_p = 28 => [12:28] - self.assertEqual(len(freed), 2) - expected_committed = torch.arange(4, 10, dtype=torch.int64) - expected_overalloc = torch.arange(12, 28, dtype=torch.int64) - self.assertTrue(torch.equal(freed[0], expected_committed)) - self.assertTrue(torch.equal(freed[1], expected_overalloc)) + # Prefill [0:4], committed [4:10], + # overallocated: start_p = ceil_align(10, 4) = 12, end_p = 28 => [12:28] + self.assertEqual(len(freed), 3) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 4, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[1], torch.arange(4, 10, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[2], torch.arange(12, 28, dtype=torch.int64))) def test_overallocation_page_aligned_noop(self): """When ceil_align(committed, page_size) >= allocated, no overalloc free.""" @@ -150,15 +152,15 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=10, # ceil_align(10, 4) = 12 kv_allocated_len=12, # same as aligned start + origin_len=4, ) - prefill_offloaded_len = 4 - manager._release_finished_req(req, prefill_offloaded_len) + manager._release_finished_req(req) - # Only committed [4:10], no overalloc because start_p == end_p - self.assertEqual(len(freed), 1) - expected_committed = torch.arange(4, 10, dtype=torch.int64) - self.assertTrue(torch.equal(freed[0], expected_committed)) + # Prefill [0:4] and committed [4:10]; no overalloc since start_p == end_p + 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, 10, dtype=torch.int64))) def test_prefix_indices_decremented(self): """protected_size_ is decremented by len(req.prefix_indices).""" @@ -171,96 +173,79 @@ class TestReleaseFinishedReq(unittest.TestCase): prefix_indices_len=5, ) - manager._release_finished_req(req, start_offset=0) + manager._release_finished_req(req) self.assertEqual(manager.tree_cache.protected_size_, 5) - def test_release_finished_req_frees_prefill_when_state_present(self): + def test_release_finished_req_frees_prefill_and_pops_state(self): """ - When offloaded_state[req].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). + _release_finished_req frees the prefill-aligned slots in addition to + the committed range; freeing them mid-decode instead races with + concurrent admission and cross-pollinates 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[req] = OffloadedState( - prefill_len=8, inc_len=0, last_hash=None + rid="req-prefill-present", + origin_len=8, ) + manager.offloaded_state[req] = OffloadedState(inc_len=4) - manager._release_finished_req(req, start_offset=8) + manager._release_finished_req(req) # 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)) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64))) + self.assertTrue(torch.equal(freed[1], torch.arange(8, 20, dtype=torch.int64))) # State entry is removed at the end of _release_finished_req. self.assertNotIn(req, manager.offloaded_state) - def test_release_finished_req_skips_prefill_free_when_prefill_len_zero(self): + def test_release_finished_req_skips_prefill_free_when_prompt_below_page(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. + When the prompt is shorter than page_size (no prefill chunk was ever + offloaded), no prefill-aligned free is emitted. """ - manager, freed = _make_manager(pool_size=32) - rid = "req-prefill-zero" + manager, freed = _make_manager(pool_size=32, page_size=4) req = _make_mock_req( req_pool_idx=0, kv_committed_len=10, kv_allocated_len=10, - rid=rid, - ) - manager.offloaded_state[req] = OffloadedState( - prefill_len=0, inc_len=0, last_hash=None + rid="req-prefill-zero", + origin_len=3, # 3 // 4 * 4 == 0 ) - manager._release_finished_req(req, start_offset=0) + manager._release_finished_req(req) # 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)) + self.assertTrue(torch.equal(freed[0], torch.arange(0, 10, dtype=torch.int64))) - def test_finalize_release_creates_state_so_prefill_is_freed(self): + def test_finalize_release_frees_prefill_without_prior_state(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. + offload ever ran: the prefill-aligned slots must still be freed by + the consolidated free site in _release_finished_req. """ - page_size = 4 - manager, freed = _make_manager(pool_size=32, page_size=page_size) - rid = "req-finalize-no-state" + manager, freed = _make_manager(pool_size=32, page_size=4) req = _make_mock_req( req_pool_idx=0, kv_committed_len=13, kv_allocated_len=13, - rid=rid, + rid="req-finalize-no-state", + origin_len=12, # prefill_len = 12 // 4 * 4 = 12 ) - # 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]. + # _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. + # No state entry is left behind. self.assertNotIn(req, manager.offloaded_state) def test_unfinished_offload_ack_does_not_free_incremental_slots(self): @@ -269,17 +254,13 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=1 ) req.finished.return_value = False - manager.offloaded_state[req] = OffloadedState( - prefill_len=4, inc_len=4, last_hash=None - ) + manager.offloaded_state[req] = OffloadedState(inc_len=4) manager.offload_inflight[req] = 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 = [ @@ -380,9 +361,7 @@ class TestReleaseFinishedReq(unittest.TestCase): req = _make_mock_req( req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=2 ) - manager.offloaded_state[req] = OffloadedState( - prefill_len=4, inc_len=8, last_hash=None - ) + manager.offloaded_state[req] = OffloadedState(inc_len=8) manager.offload_inflight[req] = 1 manager.finalize_release_on_finish(req) @@ -397,17 +376,13 @@ class TestReleaseFinishedReq(unittest.TestCase): req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=3 ) req.finished.return_value = True - manager.offloaded_state[req] = OffloadedState( - prefill_len=4, inc_len=8, last_hash=None - ) + manager.offloaded_state[req] = OffloadedState(inc_len=8) manager.offload_inflight[req] = 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 = [ @@ -426,20 +401,20 @@ class TestReleaseFinishedReq(unittest.TestCase): ): 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_pool_idx=0, + kv_committed_len=20, + kv_allocated_len=20, + rid=4, + origin_len=4, ) req.finished.return_value = True - manager.offloaded_state[req] = OffloadedState( - prefill_len=4, inc_len=8, last_hash=None - ) + manager.offloaded_state[req] = OffloadedState(inc_len=8) manager.offload_inflight[req] = 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 = [ @@ -456,6 +431,20 @@ class TestReleaseFinishedReq(unittest.TestCase): self.assertNotIn(req, manager.offloaded_state) self.assertNotIn(req, manager.offload_inflight) + def test_dropped_req_does_not_pin_offload_state(self): + manager, _ = _make_manager(pool_size=32) + req = _make_mock_req( + req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=6 + ) + manager.offloaded_state[req] = OffloadedState(inc_len=4) + manager.offload_inflight[req] = 1 + + del req + gc.collect() + + self.assertEqual(len(manager.offloaded_state), 0) + self.assertEqual(len(manager.offload_inflight), 0) + if __name__ == "__main__": unittest.main()