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
@@ -5,6 +5,7 @@ import logging
import threading import threading
import time import time
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from weakref import WeakKeyDictionary as WeakKeyDict
import torch import torch
@@ -93,12 +94,10 @@ class DecodeKVCacheOffloadManager:
self.ongoing_offload = {} self.ongoing_offload = {}
self.ongoing_backup = {} self.ongoing_backup = {}
# A caller may reuse a rid as soon as the previous response finishes, # Keyed by Req identity (rids can be reused while a D2H copy is still
# while that request's asynchronous D2H copy is still in flight. Key # in flight); weak keys so a dropped Req is never pinned here.
# lifecycle state by the Req instance so a late ack cannot mutate the self.offloaded_state: WeakKeyDict[Req, OffloadedState] = WeakKeyDict()
# new request's state. self.offload_inflight: WeakKeyDict[Req, int] = WeakKeyDict()
self.offloaded_state: dict[Req, OffloadedState] = {}
self.offload_inflight: dict[Req, int] = {}
logger.info("Enable offload kv cache for decode side") logger.info("Enable offload kv cache for decode side")
def release_host_resources(self) -> None: def release_host_resources(self) -> None:
@@ -117,6 +116,10 @@ class DecodeKVCacheOffloadManager:
def _has_inflight_offload(self, req: Req): def _has_inflight_offload(self, req: Req):
return self.offload_inflight.get(req, 0) > 0 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: def offload_kv_cache(self, req) -> bool:
"""Offload incremental KV cache for decode side.""" """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 # Prefill side offloads page-aligned origin_input_ids, decode side offloads the incremental part
all_tokens = req.origin_input_ids + req.output_ids[:-1] all_tokens = req.origin_input_ids + req.output_ids[:-1]
prefill_offloaded_len = ( prefill_offloaded_len = self._prefill_offloaded_len(req)
len(req.origin_input_ids) // self.page_size * self.page_size
)
state = self.offloaded_state.get(req) state = self.offloaded_state.get(req)
if state is None: if state is None:
prefill_hashes = self._compute_prefix_hash( prefill_hashes = self._compute_prefix_hash(
@@ -143,13 +144,9 @@ class DecodeKVCacheOffloadManager:
last_prefill_hash = ( last_prefill_hash = (
prefill_hashes[-1] if prefill_offloaded_len > 0 else None prefill_hashes[-1] if prefill_offloaded_len > 0 else None
) )
state = OffloadedState( state = OffloadedState(last_hash=last_prefill_hash)
prefill_len=prefill_offloaded_len,
inc_len=0,
last_hash=last_prefill_hash,
)
self.offloaded_state[req] = state 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_new = incremental_total - state.inc_len
incremental_aligned_len = ( incremental_aligned_len = (
incremental_new // self.offload_stride * self.offload_stride incremental_new // self.offload_stride * self.offload_stride
@@ -159,7 +156,7 @@ class DecodeKVCacheOffloadManager:
return False return False
# Extract incremental tokens and indices for the newly available chunk # 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 end = start + incremental_aligned_len
incremental_tokens = all_tokens[start:end] incremental_tokens = all_tokens[start:end]
incremental_indices = token_indices[start:end] incremental_indices = token_indices[start:end]
@@ -187,8 +184,6 @@ class DecodeKVCacheOffloadManager:
host_indices, host_indices,
incremental_tokens, incremental_tokens,
time.time(), time.time(),
start,
end,
) )
state.inc_len += incremental_aligned_len state.inc_len += incremental_aligned_len
return True return True
@@ -224,8 +219,6 @@ class DecodeKVCacheOffloadManager:
host_indices, host_indices,
incremental_tokens, incremental_tokens,
start_time, start_time,
start,
end,
) = self.ongoing_offload.pop(ack_id) ) = self.ongoing_offload.pop(ack_id)
self._mark_offload_finished(req) self._mark_offload_finished(req)
@@ -241,12 +234,10 @@ class DecodeKVCacheOffloadManager:
self.offloaded_state[req].last_hash = last_hash self.offloaded_state[req].last_hash = last_hash
if req.finished() and not self._has_inflight_offload(req): if req.finished() and not self._has_inflight_offload(req):
state = self.offloaded_state.get(req) self._release_finished_req(req)
start_offset = state.prefill_len if state is not None else start
self._release_finished_req(req, start_offset)
finish_count -= 1 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, # Defensive guard: ReqToTokenPool.free sets req_pool_idx to None,
# so a previously-released request must be skipped here to avoid # so a previously-released request must be skipped here to avoid
# non-idempotent side effects (e.g. tree_cache.protected_size_ # non-idempotent side effects (e.g. tree_cache.protected_size_
@@ -256,18 +247,15 @@ class DecodeKVCacheOffloadManager:
kv_committed_len = req.effective_kv_committed_len() kv_committed_len = req.effective_kv_committed_len()
# Free the prefill-aligned slots. Previously this was done # Prefill-aligned slots are freed only here, at request finish; freeing
# eagerly in offload_kv_cache (mid-decode), which raced with # them mid-decode races with concurrent admission over live slots.
# concurrent admission. Now consolidated here at request prefill_len = self._prefill_offloaded_len(req)
# finish, where the request is guaranteed to no longer attend if prefill_len > 0:
# to those slots.
state = self.offloaded_state.get(req)
if state is not None and state.prefill_len > 0:
prefill_indices = self.req_to_token_pool.req_to_token[ 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) self.token_to_kv_pool_allocator.free(prefill_indices)
start = start_offset start = prefill_len
end = kv_committed_len end = kv_committed_len
# Free the incremental part of the request (DSA-aware) # 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] 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): def finalize_release_on_finish(self, req: Req):
"""Free any remaining tail KV that was not offloaded due to non-aligned length.""" """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): if self._has_inflight_offload(req):
return return
start_offset = prefill_len self._release_finished_req(req)
self._release_finished_req(req, start_offset)
+6 -14
View File
@@ -263,21 +263,13 @@ class BlockStoredMetadata(msgspec.Struct, omit_defaults=True, gc=False):
cache_salt: str cache_salt: str
class OffloadedState: class OffloadedState(msgspec.Struct):
""" """Decode-side offload progress for one request, keyed by Req in the manager."""
OffloadedState represents the state of a KV cache block offloaded to the hicache.
- prefill_len (int): The length of the prefill part of the KV cache block. # Decode-incremental length already submitted for D2H offload.
- inc_len (int): The length of the incremental part of the KV cache block. inc_len: int = 0
- last_hash (Optional[str]): The hash of the last token in the KV cache block. # Tail of the page hash chain, extended as each offloaded chunk is backed up.
""" last_hash: Optional[str] = None
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
class BlockStored(KVCacheEvent): class BlockStored(KVCacheEvent):
@@ -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) Requires: torch, sglang (run in an environment with sglang installed)
""" """
import gc
import unittest import unittest
from unittest.mock import MagicMock from unittest.mock import MagicMock
from weakref import WeakKeyDictionary as WeakKeyDict
import torch import torch
@@ -29,10 +31,12 @@ def _make_mock_req(
kv_allocated_len: int, kv_allocated_len: int,
prefix_indices_len: int = 0, prefix_indices_len: int = 0,
rid: int = 0, rid: int = 0,
origin_len: int = 0,
): ):
"""Create a mock Req with the KV cache state needed for testing.""" """Create a mock Req with the KV cache state needed for testing."""
req = MagicMock() req = MagicMock()
req.rid = rid req.rid = rid
req.origin_input_ids = list(range(origin_len))
req.kv = ReqKvInfo( req.kv = ReqKvInfo(
req_pool_idx=req_pool_idx, req_pool_idx=req_pool_idx,
kv_committed_len=kv_committed_len, 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.token_to_kv_pool_allocator = allocator
manager.page_size = page_size manager.page_size = page_size
manager.tree_cache = tree_cache manager.tree_cache = tree_cache
manager.offloaded_state = {} manager.offloaded_state = WeakKeyDict()
manager.ongoing_offload = {} manager.ongoing_offload = {}
manager.ongoing_backup = {} manager.ongoing_backup = {}
manager.offload_inflight = {} manager.offload_inflight = WeakKeyDict()
return manager, freed_indices return manager, freed_indices
@@ -90,15 +94,15 @@ class TestReleaseFinishedReq(unittest.TestCase):
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=20, kv_committed_len=20,
kv_allocated_len=20, # no overallocation 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] # Prefill [0:8] and committed [8:20]; no overalloc free.
self.assertEqual(len(freed), 1) self.assertEqual(len(freed), 2)
expected = torch.arange(8, 20, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[0], expected)) self.assertTrue(torch.equal(freed[1], torch.arange(8, 20, dtype=torch.int64)))
manager.req_to_token_pool.free.assert_called_once_with(req) manager.req_to_token_pool.free.assert_called_once_with(req)
def test_with_overallocation(self): def test_with_overallocation(self):
@@ -108,17 +112,16 @@ class TestReleaseFinishedReq(unittest.TestCase):
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=20, kv_committed_len=20,
kv_allocated_len=28, # 8 over-allocated slots 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] # Prefill [0:8], committed [8:20], overallocated [20:28].
self.assertEqual(len(freed), 2) self.assertEqual(len(freed), 3)
expected_committed = torch.arange(8, 20, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64)))
expected_overalloc = torch.arange(20, 28, dtype=torch.int64) self.assertTrue(torch.equal(freed[1], torch.arange(8, 20, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[0], expected_committed)) self.assertTrue(torch.equal(freed[2], torch.arange(20, 28, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[1], expected_overalloc))
manager.req_to_token_pool.free.assert_called_once_with(req) manager.req_to_token_pool.free.assert_called_once_with(req)
def test_overallocation_with_page_alignment(self): def test_overallocation_with_page_alignment(self):
@@ -129,18 +132,17 @@ class TestReleaseFinishedReq(unittest.TestCase):
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=10, # not page-aligned kv_committed_len=10, # not page-aligned
kv_allocated_len=28, 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] # Prefill [0:4], committed [4:10],
# Overallocated: start_p = ceil_align(10, 4) = 12, end_p = 28 => [12:28] # overallocated: start_p = ceil_align(10, 4) = 12, end_p = 28 => [12:28]
self.assertEqual(len(freed), 2) self.assertEqual(len(freed), 3)
expected_committed = torch.arange(4, 10, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], torch.arange(0, 4, dtype=torch.int64)))
expected_overalloc = torch.arange(12, 28, dtype=torch.int64) self.assertTrue(torch.equal(freed[1], torch.arange(4, 10, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[0], expected_committed)) self.assertTrue(torch.equal(freed[2], torch.arange(12, 28, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[1], expected_overalloc))
def test_overallocation_page_aligned_noop(self): def test_overallocation_page_aligned_noop(self):
"""When ceil_align(committed, page_size) >= allocated, no overalloc free.""" """When ceil_align(committed, page_size) >= allocated, no overalloc free."""
@@ -150,15 +152,15 @@ class TestReleaseFinishedReq(unittest.TestCase):
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=10, # ceil_align(10, 4) = 12 kv_committed_len=10, # ceil_align(10, 4) = 12
kv_allocated_len=12, # same as aligned start 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 # Prefill [0:4] and committed [4:10]; no overalloc since start_p == end_p
self.assertEqual(len(freed), 1) self.assertEqual(len(freed), 2)
expected_committed = torch.arange(4, 10, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], torch.arange(0, 4, dtype=torch.int64)))
self.assertTrue(torch.equal(freed[0], expected_committed)) self.assertTrue(torch.equal(freed[1], torch.arange(4, 10, dtype=torch.int64)))
def test_prefix_indices_decremented(self): def test_prefix_indices_decremented(self):
"""protected_size_ is decremented by len(req.prefix_indices).""" """protected_size_ is decremented by len(req.prefix_indices)."""
@@ -171,96 +173,79 @@ class TestReleaseFinishedReq(unittest.TestCase):
prefix_indices_len=5, 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) 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 _release_finished_req frees the prefill-aligned slots in addition to
free the prefill-aligned slots in addition to the committed range. the committed range; freeing them mid-decode instead races with
concurrent admission and cross-pollinates KV reads.
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) manager, freed = _make_manager(pool_size=32)
rid = "req-prefill-present"
req = _make_mock_req( req = _make_mock_req(
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=20, kv_committed_len=20,
kv_allocated_len=20, kv_allocated_len=20,
rid=rid, rid="req-prefill-present",
) origin_len=8,
manager.offloaded_state[req] = OffloadedState(
prefill_len=8, inc_len=0, last_hash=None
) )
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]. # Two frees in order: prefill [0:8] then committed [8:20].
self.assertEqual(len(freed), 2) self.assertEqual(len(freed), 2)
expected_prefill = torch.arange(0, 8, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], torch.arange(0, 8, dtype=torch.int64)))
expected_committed = torch.arange(8, 20, dtype=torch.int64) self.assertTrue(torch.equal(freed[1], 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. # State entry is removed at the end of _release_finished_req.
self.assertNotIn(req, manager.offloaded_state) 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, When the prompt is shorter than page_size (no prefill chunk was ever
so no prefill chunk was ever offloaded), no prefill-aligned free is offloaded), no prefill-aligned free is emitted.
emitted.
""" """
manager, freed = _make_manager(pool_size=32) manager, freed = _make_manager(pool_size=32, page_size=4)
rid = "req-prefill-zero"
req = _make_mock_req( req = _make_mock_req(
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=10, kv_committed_len=10,
kv_allocated_len=10, kv_allocated_len=10,
rid=rid, rid="req-prefill-zero",
) origin_len=3, # 3 // 4 * 4 == 0
manager.offloaded_state[req] = OffloadedState(
prefill_len=0, inc_len=0, last_hash=None
) )
manager._release_finished_req(req, start_offset=0) manager._release_finished_req(req)
# Only the committed range [0:10] is freed. # Only the committed range [0:10] is freed.
self.assertEqual(len(freed), 1) self.assertEqual(len(freed), 1)
expected_committed = torch.arange(0, 10, dtype=torch.int64) self.assertTrue(torch.equal(freed[0], 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): def test_finalize_release_frees_prefill_without_prior_state(self):
""" """
finalize_release_on_finish handles the case where no incremental finalize_release_on_finish handles the case where no incremental
offload ever ran (offloaded_state is empty). It must materialize an offload ever ran: the prefill-aligned slots must still be freed by
OffloadedState with the correct prefill_len so that the consolidated the consolidated free site in _release_finished_req.
free site in _release_finished_req can locate and free those slots.
""" """
page_size = 4 manager, freed = _make_manager(pool_size=32, page_size=4)
manager, freed = _make_manager(pool_size=32, page_size=page_size)
rid = "req-finalize-no-state"
req = _make_mock_req( req = _make_mock_req(
req_pool_idx=0, req_pool_idx=0,
kv_committed_len=13, kv_committed_len=13,
kv_allocated_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) manager.finalize_release_on_finish(req)
# finalize creates state, then _release_finished_req frees: # _release_finished_req frees prefill [0:12] then committed [12:13].
# prefill [0:12] then committed [12:13].
self.assertEqual(len(freed), 2) self.assertEqual(len(freed), 2)
expected_prefill = torch.arange(0, 12, dtype=torch.int64) expected_prefill = torch.arange(0, 12, dtype=torch.int64)
expected_committed = torch.arange(12, 13, 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[0], expected_prefill))
self.assertTrue(torch.equal(freed[1], expected_committed)) 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) self.assertNotIn(req, manager.offloaded_state)
def test_unfinished_offload_ack_does_not_free_incremental_slots(self): 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_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=1
) )
req.finished.return_value = False req.finished.return_value = False
manager.offloaded_state[req] = OffloadedState( manager.offloaded_state[req] = OffloadedState(inc_len=4)
prefill_len=4, inc_len=4, last_hash=None
)
manager.offload_inflight[req] = 1 manager.offload_inflight[req] = 1
manager.ongoing_offload[7] = ( manager.ongoing_offload[7] = (
req, req,
torch.arange(4, 8, dtype=torch.int64), torch.arange(4, 8, dtype=torch.int64),
[10, 11, 12, 13], [10, 11, 12, 13],
0.0, 0.0,
4,
8,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [ manager.cache_controller.ack_write_queue = [
@@ -380,9 +361,7 @@ class TestReleaseFinishedReq(unittest.TestCase):
req = _make_mock_req( req = _make_mock_req(
req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=2 req_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=2
) )
manager.offloaded_state[req] = OffloadedState( manager.offloaded_state[req] = OffloadedState(inc_len=8)
prefill_len=4, inc_len=8, last_hash=None
)
manager.offload_inflight[req] = 1 manager.offload_inflight[req] = 1
manager.finalize_release_on_finish(req) 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_pool_idx=0, kv_committed_len=20, kv_allocated_len=20, rid=3
) )
req.finished.return_value = True req.finished.return_value = True
manager.offloaded_state[req] = OffloadedState( manager.offloaded_state[req] = OffloadedState(inc_len=8)
prefill_len=4, inc_len=8, last_hash=None
)
manager.offload_inflight[req] = 2 manager.offload_inflight[req] = 2
manager.ongoing_offload[8] = ( manager.ongoing_offload[8] = (
req, req,
torch.arange(4, 8, dtype=torch.int64), torch.arange(4, 8, dtype=torch.int64),
[10, 11, 12, 13], [10, 11, 12, 13],
0.0, 0.0,
4,
8,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [ manager.cache_controller.ack_write_queue = [
@@ -426,20 +401,20 @@ class TestReleaseFinishedReq(unittest.TestCase):
): ):
manager, freed = _make_manager(pool_size=32) manager, freed = _make_manager(pool_size=32)
req = _make_mock_req( 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 req.finished.return_value = True
manager.offloaded_state[req] = OffloadedState( manager.offloaded_state[req] = OffloadedState(inc_len=8)
prefill_len=4, inc_len=8, last_hash=None
)
manager.offload_inflight[req] = 1 manager.offload_inflight[req] = 1
manager.ongoing_offload[9] = ( manager.ongoing_offload[9] = (
req, req,
torch.arange(8, 12, dtype=torch.int64), torch.arange(8, 12, dtype=torch.int64),
[14, 15, 16, 17], [14, 15, 16, 17],
0.0, 0.0,
8,
12,
) )
manager.cache_controller = MagicMock() manager.cache_controller = MagicMock()
manager.cache_controller.ack_write_queue = [ manager.cache_controller.ack_write_queue = [
@@ -456,6 +431,20 @@ class TestReleaseFinishedReq(unittest.TestCase):
self.assertNotIn(req, manager.offloaded_state) self.assertNotIn(req, manager.offloaded_state)
self.assertNotIn(req, manager.offload_inflight) 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__": if __name__ == "__main__":
unittest.main() unittest.main()