[BugFix] fix(hicache): fix two slot-reuse races in DecodeKVCacheOffloadManager (#24226)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -21,7 +21,6 @@ register_cuda_ci(
|
||||
est_time=600,
|
||||
stage="base-b",
|
||||
runner_config="2-gpu-large",
|
||||
disabled="Temporarily disable the flaky test.",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user