refactor(hicache): simplify decode offload state bookkeeping (#37299)

This commit is contained in:
Liangsheng Yin
2026-08-31 23:09:41 -07:00
committed by GitHub
parent ae2bd5728b
commit 959ca033eb
3 changed files with 111 additions and 160 deletions
@@ -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()