refactor(hicache): simplify decode offload state bookkeeping (#37299)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user