Let cache backend do not couple with owned committed kv details and avoid kv_committed_freed/kv_overallocated_freed fields (#29428)
This commit is contained in:
@@ -36,22 +36,8 @@ def _make_mock_req(
|
||||
req.req_pool_idx = req_pool_idx
|
||||
req.kv_committed_len = kv_committed_len
|
||||
req.kv = SimpleNamespace(kv_allocated_len=kv_allocated_len)
|
||||
req.kv_committed_freed = False
|
||||
req.kv_overallocated_freed = False
|
||||
req.prefix_indices = list(range(prefix_indices_len))
|
||||
|
||||
def pop_committed():
|
||||
assert not req.kv_committed_freed
|
||||
req.kv_committed_freed = True
|
||||
return req.kv_committed_len
|
||||
|
||||
def pop_overallocated():
|
||||
assert not req.kv_overallocated_freed
|
||||
req.kv_overallocated_freed = True
|
||||
return req.kv_committed_len, req.kv.kv_allocated_len
|
||||
|
||||
req.pop_committed_kv_cache = pop_committed
|
||||
req.pop_overallocated_kv_cache = pop_overallocated
|
||||
req.effective_kv_committed_len = lambda: req.kv_committed_len
|
||||
return req
|
||||
|
||||
|
||||
|
||||
@@ -83,18 +83,10 @@ class MockReq:
|
||||
self.priority = 0
|
||||
self.kv_committed_len = len(fill_ids)
|
||||
self.kv = SimpleNamespace(kv_allocated_len=len(fill_ids))
|
||||
self.kv_committed_freed = False
|
||||
|
||||
def get_fill_ids(self):
|
||||
return self.full_untruncated_fill_ids[: self.extend_range.end]
|
||||
|
||||
def pop_committed_kv_cache(self):
|
||||
self.kv_committed_freed = True
|
||||
return self.kv_committed_len
|
||||
|
||||
def pop_overallocated_kv_cache(self):
|
||||
return (self.kv_committed_len, self.kv.kv_allocated_len)
|
||||
|
||||
|
||||
def _make_req(fill_ids, req_pool_idx=0, cache_protected_len=0, last_node=None):
|
||||
return MockReq(fill_ids, req_pool_idx, cache_protected_len, last_node)
|
||||
@@ -152,7 +144,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
cache.cache_unfinished_req(req)
|
||||
|
||||
# Step 3: cache_finished_req with is_insert=True (dec lock)
|
||||
cache.cache_finished_req(req)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv_committed_len)
|
||||
|
||||
# Verify: all non-root nodes should have lock_ref == 0
|
||||
# (root always has lock_ref == 1)
|
||||
@@ -201,7 +193,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
cache.cache_unfinished_req(req)
|
||||
|
||||
# Step 3: cache_finished_req (dec leaf)
|
||||
cache.cache_finished_req(req)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv_committed_len)
|
||||
|
||||
# Root lock unchanged, all nodes unlocked
|
||||
self.assertEqual(cache.root_node.lock_ref, root_lock_before)
|
||||
@@ -244,7 +236,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
|
||||
# Transfer fails -> cache_finished_req with is_insert=False
|
||||
# This frees delta tokens and dec_lock_ref on last_node
|
||||
cache.cache_finished_req(req, is_insert=False)
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=False, kv_len_to_handle=req.kv_committed_len
|
||||
)
|
||||
|
||||
# The prefix node should be unlocked (back to evictable)
|
||||
self.assertEqual(cache.root_node.lock_ref, 1)
|
||||
@@ -289,7 +283,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
|
||||
# Transfer fails -> cache_finished_req with is_insert=False
|
||||
# dec_lock_ref(root) is a no-op
|
||||
cache.cache_finished_req(req, is_insert=False)
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=False, kv_len_to_handle=req.kv_committed_len
|
||||
)
|
||||
|
||||
# Root lock unchanged, nothing protected or evictable
|
||||
self.assertEqual(cache.root_node.lock_ref, root_lock_before)
|
||||
@@ -400,7 +396,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
)
|
||||
|
||||
cache.cache_unfinished_req(req)
|
||||
cache.cache_finished_req(req)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv_committed_len)
|
||||
|
||||
# After all iterations, root lock should be 1, no protected nodes
|
||||
self.assertEqual(cache.root_node.lock_ref, 1)
|
||||
|
||||
@@ -37,7 +37,7 @@ class TestPureSWAChunkCache(CustomTestCase):
|
||||
)
|
||||
cache.token_to_kv_pool_allocator = _FakeAllocator()
|
||||
|
||||
cache.cache_finished_req(_FakeReq())
|
||||
cache.cache_finished_req(_FakeReq(), kv_len_to_handle=8)
|
||||
|
||||
self.assertEqual(len(cache.token_to_kv_pool_allocator.freed), 1)
|
||||
freed = cache.token_to_kv_pool_allocator.freed[0]
|
||||
|
||||
@@ -61,8 +61,6 @@ class _FakeReq:
|
||||
kv_allocated_len=allocated,
|
||||
swa_evicted_seqlen=0,
|
||||
)
|
||||
self.kv_committed_freed = False
|
||||
self.kv_overallocated_freed = False
|
||||
self.origin_input_ids = list(range(committed))
|
||||
self.output_ids = []
|
||||
self.extra_key = None
|
||||
@@ -74,22 +72,10 @@ class _FakeReq:
|
||||
self.mamba_next_track_idx = None
|
||||
self.mamba_last_track_seqlen = None
|
||||
self.mamba_branching_seqlen = None
|
||||
self.pop_overallocated_calls = 0
|
||||
self.to_finish = None
|
||||
self.finished_reason = None
|
||||
self.finished_len = None
|
||||
|
||||
def pop_committed_kv_cache(self):
|
||||
assert not self.kv_committed_freed
|
||||
self.kv_committed_freed = True
|
||||
return self.kv_committed_len
|
||||
|
||||
def pop_overallocated_kv_cache(self):
|
||||
assert not self.kv_overallocated_freed
|
||||
self.pop_overallocated_calls += 1
|
||||
self.kv_overallocated_freed = True
|
||||
return self.kv_committed_len, self.kv.kv_allocated_len
|
||||
|
||||
|
||||
def test_preabort_detaches_session_and_preserves_slot():
|
||||
"""Pre-aborted req (to_finish set before match_prefix) is detached from
|
||||
@@ -161,9 +147,6 @@ def test_first_mid_abort_nukes_ephemeral_slot():
|
||||
assert req_to_token_pool.free_slots == [0]
|
||||
assert len(allocator.freed) == 1
|
||||
assert allocator.freed[0].tolist() == list(range(20))
|
||||
# Bookkeeping flags set.
|
||||
assert req.kv_committed_freed is True
|
||||
assert req.kv_overallocated_freed is True
|
||||
|
||||
|
||||
def test_nth_mid_abort_nukes_session_slot():
|
||||
@@ -200,9 +183,6 @@ def test_nth_mid_abort_nukes_session_slot():
|
||||
# Pool slot returned.
|
||||
assert req_to_token_pool.free_slots == [0]
|
||||
assert req.req_pool_idx is None
|
||||
# Bookkeeping flags set.
|
||||
assert req.kv_committed_freed is True
|
||||
assert req.kv_overallocated_freed is True
|
||||
|
||||
|
||||
# Shrink tests removed: streaming sessions are append-only after the
|
||||
|
||||
@@ -116,7 +116,6 @@ def _make_req(req_pool_idx, token_ids, cache_protected_len, tree):
|
||||
prefix_indices=torch.tensor([], dtype=torch.int64, device=tree.device),
|
||||
_kv_committed_len=len(token_ids),
|
||||
)
|
||||
req.pop_committed_kv_cache = lambda: req._kv_committed_len
|
||||
return req
|
||||
|
||||
|
||||
@@ -191,7 +190,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
insert_len = seq_len // page_size * page_size
|
||||
self.assertLess(req.kv.swa_evicted_seqlen, insert_len)
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
# -- Eviction formula: page_size == 1 --
|
||||
@@ -216,7 +217,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
req.kv.swa_evicted_seqlen, max(0, seq_len - 1 - max(window, page_size))
|
||||
)
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
# -- Eviction formula: no-op when seq too short --
|
||||
@@ -253,7 +256,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
kv1 = _swa_alloc(allocator, first_len)
|
||||
pool.write((0, slice(0, first_len)), kv1)
|
||||
req1 = _make_req(0, list(range(first_len)), 0, tree)
|
||||
tree.cache_finished_req(req1, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req1, is_insert=True, kv_len_to_handle=req1._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
# Second request: 24 tokens, first 16 overlap with tree
|
||||
@@ -269,7 +274,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
self.assertLessEqual(req2.kv.swa_evicted_seqlen, first_len)
|
||||
|
||||
swa_evictable_before = tree.swa_evictable_size_
|
||||
tree.cache_finished_req(req2, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req2, is_insert=True, kv_len_to_handle=req2._kv_committed_len
|
||||
)
|
||||
|
||||
# New tokens [16, 24) should all be non-tombstone
|
||||
new_tokens = second_len // page_size * page_size - first_len
|
||||
@@ -300,7 +307,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
self.assertGreater(req.kv.swa_evicted_seqlen, 0, "Should have some eviction")
|
||||
self.assertLess(req.kv.swa_evicted_seqlen, insert_len, "Should be partial")
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
|
||||
non_tombstone = insert_len - req.kv.swa_evicted_seqlen
|
||||
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + non_tombstone)
|
||||
@@ -338,7 +347,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
req.kv.swa_evicted_seqlen = old_evicted
|
||||
swa_evictable_before = tree.swa_evictable_size_
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
|
||||
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before)
|
||||
|
||||
@@ -367,7 +378,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
insert_len = seq_len // page_size * page_size
|
||||
self.assertLess(req.kv.swa_evicted_seqlen, insert_len, f"turn {turn}")
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
# -- Integration: page_size=1 full flow --
|
||||
@@ -391,7 +404,9 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
req.kv.swa_evicted_seqlen, max(0, seq_len - 1 - max(window, page_size))
|
||||
)
|
||||
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
|
||||
|
||||
@@ -38,9 +38,6 @@ class _DummyReq:
|
||||
self.swa_prefix_lock_released = False
|
||||
self.kv = SimpleNamespace(swa_evicted_seqlen=0)
|
||||
|
||||
def pop_committed_kv_cache(self):
|
||||
return self._kv_committed_len
|
||||
|
||||
|
||||
def _build_swa_tree(
|
||||
is_eagle: bool,
|
||||
@@ -592,7 +589,9 @@ class TestSWA(unittest.TestCase):
|
||||
return original_insert(params)
|
||||
|
||||
tree.insert = wrapped_insert
|
||||
tree.cache_finished_req(req, is_insert=True)
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
|
||||
self.assertEqual(captured["prev_prefix_len"], req.cache_protected_len)
|
||||
self.assertTrue(captured["is_bigram"])
|
||||
@@ -624,7 +623,9 @@ class TestSWA(unittest.TestCase):
|
||||
return original_free(indices)
|
||||
|
||||
allocator.free = wrapped_free
|
||||
tree.cache_finished_req(req2, is_insert=False)
|
||||
tree.cache_finished_req(
|
||||
req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len
|
||||
)
|
||||
|
||||
# EAGLE + page_size=1 => page_aligned_len = committed_len - 1 = 5
|
||||
# Expected frees:
|
||||
|
||||
@@ -643,7 +643,6 @@ def bench_cache_finished(
|
||||
req.last_node = node
|
||||
req.cache_protected_len = matched_len
|
||||
req.kv_committed_len = len(seq)
|
||||
req.kv_committed_freed = False
|
||||
if hasattr(lr, "swa_uuid_for_lock"):
|
||||
req.swa_uuid_for_lock = lr.swa_uuid_for_lock
|
||||
env.rtp.req_to_token[req.req_pool_idx, : len(kv_indices)] = kv_indices
|
||||
@@ -656,7 +655,9 @@ def bench_cache_finished(
|
||||
return bench_api(
|
||||
"cache_finished",
|
||||
lambda: req_items,
|
||||
lambda req: env.tree.cache_finished_req(req, is_insert=True),
|
||||
lambda req: env.tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.kv_committed_len
|
||||
),
|
||||
len(req_items) - warmup,
|
||||
env.avg_tokens,
|
||||
warmup,
|
||||
|
||||
@@ -854,7 +854,9 @@ class UnifiedRadixCacheSuite:
|
||||
if self.cfg.has_mamba:
|
||||
req.mamba_last_track_seqlen = kv_len
|
||||
|
||||
cache.cache_finished_req(req, is_insert=True)
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
|
||||
all_ids = input_ids + output_ids
|
||||
aligned_len = (len(all_ids) // ps) * ps
|
||||
@@ -893,8 +895,10 @@ class UnifiedRadixCacheSuite:
|
||||
get_server_args().strip_thinking_cache = True
|
||||
try:
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(req, is_insert=True)
|
||||
start_p, end_p = req.pop_overallocated_kv_cache()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
start_p, end_p = req.effective_kv_committed_len(), req.kv.kv_allocated_len
|
||||
finally:
|
||||
get_server_args().strip_thinking_cache = False
|
||||
if ps > 1:
|
||||
@@ -936,7 +940,9 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(req, is_insert=False)
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=False, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
|
||||
self.assertEqual(allocator.available_size(), avail_before + kv_len)
|
||||
m = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
|
||||
@@ -1108,7 +1114,9 @@ class UnifiedRadixCacheSuite:
|
||||
req.mamba_last_track_seqlen = kv_len
|
||||
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(req, is_insert=True)
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
|
||||
self.assertEqual(allocator.available_size(), avail_before + tail_extra)
|
||||
aligned = input_ids[: (len(input_ids) // ps) * ps]
|
||||
|
||||
Reference in New Issue
Block a user