[mem_cache] Release up to owned_kv_len on radix cache insert (#40075)
This commit is contained in:
@@ -234,7 +234,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
cache.cache_unfinished_req(req)
|
||||
|
||||
# Step 3: cache_finished_req with is_insert=True (dec lock)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len)
|
||||
cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len)
|
||||
|
||||
# Verify: all non-root nodes should have lock_ref == 0
|
||||
# (root always has lock_ref == 1)
|
||||
@@ -283,7 +283,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
cache.cache_unfinished_req(req)
|
||||
|
||||
# Step 3: cache_finished_req (dec leaf)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len)
|
||||
cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len)
|
||||
|
||||
# Root lock unchanged, all nodes unlocked
|
||||
self.assertEqual(cache.root_node.lock_ref, root_lock_before)
|
||||
@@ -331,7 +331,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
# Transfer fails -> cache_finished_req with is_insert=False
|
||||
cache.token_to_kv_pool_allocator.reset_mock()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=False, kv_len_to_handle=req.kv.kv_committed_len
|
||||
req, is_insert=False, owned_kv_len=req.kv.kv_committed_len
|
||||
)
|
||||
|
||||
free_call = cache.token_to_kv_pool_allocator.free_segment.call_args
|
||||
@@ -346,6 +346,53 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
# Prefix tokens should still be in tree and evictable
|
||||
self.assertEqual(cache.evictable_size(), len(prefix))
|
||||
|
||||
def test_insert_releases_committed_slot_without_token_id(self):
|
||||
"""The insert path releases up to owned_kv_len, not len(token_ids).
|
||||
|
||||
Pins the ownership contract on BasePrefixCache.cache_finished_req.
|
||||
"""
|
||||
cache, req_to_token = _make_cache_with_pools()
|
||||
|
||||
prefix = [1, 2, 3]
|
||||
prefix_vals = [10, 20, 30]
|
||||
self._populate_prefix(cache, prefix, prefix_vals)
|
||||
|
||||
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", prefix))))
|
||||
matched_node = result.last_device_node
|
||||
prefix_len = len(result.device_indices)
|
||||
cache.inc_lock_ref(matched_node)
|
||||
|
||||
# Token sequence is 5 long; a 6th KV slot is committed with no token id.
|
||||
full_ids = [1, 2, 3, 4, 5]
|
||||
row_vals = [10, 20, 30, 40, 50, 60]
|
||||
req_to_token[0, : len(row_vals)] = torch.tensor(row_vals, dtype=torch.int64)
|
||||
|
||||
req = _make_req(
|
||||
fill_ids=full_ids,
|
||||
req_pool_idx=0,
|
||||
cache_protected_len=prefix_len,
|
||||
last_node=matched_node,
|
||||
)
|
||||
req.kv.kv_committed_len = len(row_vals)
|
||||
req.kv.kv_allocated_len = len(row_vals)
|
||||
|
||||
cache.token_to_kv_pool_allocator.reset_mock()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, owned_kv_len=req.kv.kv_committed_len
|
||||
)
|
||||
|
||||
# The unnamed tail slot is freed as the segment past the radix key.
|
||||
segments = cache.token_to_kv_pool_allocator.free_segments.call_args.args[0]
|
||||
freed = {
|
||||
int(start_pos) + i: int(v)
|
||||
for indices, start_pos in segments
|
||||
for i, v in enumerate(indices.tolist())
|
||||
}
|
||||
self.assertIn(
|
||||
len(full_ids), freed, "committed slot with no token id was not freed"
|
||||
)
|
||||
self.assertEqual(freed[len(full_ids)], row_vals[-1])
|
||||
|
||||
def test_full_transfer_failure(self):
|
||||
"""Scenario 4: no prefix match, transfer fails.
|
||||
|
||||
@@ -384,7 +431,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
# 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, kv_len_to_handle=req.kv.kv_committed_len
|
||||
req, is_insert=False, owned_kv_len=req.kv.kv_committed_len
|
||||
)
|
||||
|
||||
# Root lock unchanged, nothing protected or evictable
|
||||
@@ -572,7 +619,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
)
|
||||
|
||||
cache.cache_unfinished_req(req)
|
||||
cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len)
|
||||
cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len)
|
||||
|
||||
# After all iterations, root lock should be 1, no protected nodes
|
||||
self.assertEqual(cache.root_node.lock_ref, 1)
|
||||
|
||||
@@ -768,7 +768,7 @@ class TestRotationGraftDecline(CustomTestCase):
|
||||
req = _GraftReq(list(range(8)) + [90, 91, 92, 93])
|
||||
req.kv_rotation_base = 3
|
||||
own_locs = self._own_row(tree, req, 12)
|
||||
tree.cache_finished_req(req, kv_len_to_handle=12)
|
||||
tree.cache_finished_req(req, owned_kv_len=12)
|
||||
released = torch.cat(freed)
|
||||
# Everything past the protected prefix is released: the duplicates of
|
||||
# the matched region AND the declined tail (nothing leaks, nothing is
|
||||
@@ -784,7 +784,7 @@ class TestRotationGraftDecline(CustomTestCase):
|
||||
req = _GraftReq(list(range(8)) + [90, 91, 92, 93])
|
||||
req.kv_rotation_base = 1
|
||||
own_locs = self._own_row(tree, req, 12)
|
||||
tree.cache_finished_req(req, kv_len_to_handle=12)
|
||||
tree.cache_finished_req(req, owned_kv_len=12)
|
||||
self.assertEqual(_match_len(tree, req.fill_ids), 12)
|
||||
released = torch.cat(freed) if freed else torch.empty(0, dtype=torch.int64)
|
||||
# Only the 8 duplicate rows go back; the tail stays live in the tree.
|
||||
|
||||
@@ -45,7 +45,7 @@ class TestPureSWAChunkCache(CustomTestCase):
|
||||
def test_finished_req_skips_already_evicted_swa_range(self):
|
||||
cache = self._make_cache()
|
||||
|
||||
cache.cache_finished_req(_FakeReq(), kv_len_to_handle=8)
|
||||
cache.cache_finished_req(_FakeReq(), owned_kv_len=8)
|
||||
|
||||
self.assertEqual(len(cache.token_to_kv_pool_allocator.freed), 1)
|
||||
freed = cache.token_to_kv_pool_allocator.freed[0]
|
||||
@@ -56,7 +56,7 @@ class TestPureSWAChunkCache(CustomTestCase):
|
||||
req = _FakeReq()
|
||||
req.kv.cache_protected_len = 2
|
||||
|
||||
cache.cache_finished_req(req, kv_len_to_handle=8)
|
||||
cache.cache_finished_req(req, owned_kv_len=8)
|
||||
|
||||
freed = cache.token_to_kv_pool_allocator.freed[0]
|
||||
self.assertTrue(torch.equal(freed, torch.tensor([2, 6, 7])))
|
||||
|
||||
@@ -26,7 +26,6 @@ class TestPureSWARadixCache(CustomTestCase):
|
||||
def test_no_insert_frees_window_after_evict_floor_before_swa_eviction(self):
|
||||
allocator = _FakeAllocator()
|
||||
cache = PureSWARadixCache.__new__(PureSWARadixCache)
|
||||
cache.disable_finished_insert = False
|
||||
cache.disable = False
|
||||
cache.is_eagle = False
|
||||
cache.page_size = 1
|
||||
@@ -48,7 +47,7 @@ class TestPureSWARadixCache(CustomTestCase):
|
||||
),
|
||||
)
|
||||
|
||||
cache.cache_finished_req(req, is_insert=False, kv_len_to_handle=8)
|
||||
cache.cache_finished_req(req, is_insert=False, owned_kv_len=8)
|
||||
|
||||
self.assertEqual(allocator.freed, [[0, 1, 2, 3], [4, 5, 6, 7]])
|
||||
|
||||
|
||||
@@ -572,7 +572,7 @@ class TestRadixCache(CustomTestCase):
|
||||
cache.cache_finished_req(
|
||||
req,
|
||||
is_insert=True,
|
||||
kv_len_to_handle=len(prompt_ids) + len(output_ids),
|
||||
owned_kv_len=len(prompt_ids) + len(output_ids),
|
||||
)
|
||||
|
||||
(prompt_node,) = cache.root_node.children.values()
|
||||
|
||||
@@ -191,7 +191,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
self.assertLess(req.kv.swa_evicted_seqlen, insert_len)
|
||||
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
req, is_insert=True, owned_kv_len=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
@@ -358,7 +358,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
)
|
||||
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
req, is_insert=True, owned_kv_len=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
@@ -397,7 +397,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
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, kv_len_to_handle=req1._kv_committed_len
|
||||
req1, is_insert=True, owned_kv_len=req1._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
@@ -415,7 +415,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
|
||||
swa_evictable_before = tree.swa_evictable_size_
|
||||
tree.cache_finished_req(
|
||||
req2, is_insert=True, kv_len_to_handle=req2._kv_committed_len
|
||||
req2, is_insert=True, owned_kv_len=req2._kv_committed_len
|
||||
)
|
||||
|
||||
# New tokens [16, 24) should all be non-tombstone
|
||||
@@ -447,9 +447,7 @@ 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, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
|
||||
|
||||
non_tombstone = insert_len - req.kv.swa_evicted_seqlen
|
||||
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + non_tombstone)
|
||||
@@ -487,9 +485,7 @@ 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, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
|
||||
|
||||
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before)
|
||||
|
||||
@@ -519,7 +515,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
self.assertLess(req.kv.swa_evicted_seqlen, insert_len, f"turn {turn}")
|
||||
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
req, is_insert=True, owned_kv_len=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
@@ -545,7 +541,7 @@ class TestSWAEvictionBoundary(unittest.TestCase):
|
||||
)
|
||||
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
req, is_insert=True, owned_kv_len=req._kv_committed_len
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
|
||||
@@ -813,9 +813,7 @@ class TestSWA(unittest.TestCase):
|
||||
return original_insert(params)
|
||||
|
||||
tree.insert = wrapped_insert
|
||||
tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req._kv_committed_len
|
||||
)
|
||||
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
|
||||
|
||||
self.assertEqual(captured["prev_prefix_len"], req.kv.cache_protected_len)
|
||||
self.assertTrue(captured["is_bigram"])
|
||||
@@ -849,7 +847,7 @@ class TestSWA(unittest.TestCase):
|
||||
|
||||
allocator.free_segment = wrapped_free_segment
|
||||
tree.cache_finished_req(
|
||||
req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len
|
||||
req2, is_insert=False, owned_kv_len=req2._kv_committed_len
|
||||
)
|
||||
|
||||
# EAGLE + page_size=1 => page_aligned_len = committed_len - 1 = 5
|
||||
@@ -1465,7 +1463,7 @@ class TestCacheUnfinishedReqEvictedPrefix(CustomTestCase):
|
||||
|
||||
# Finishing drops the locks, which sanity_check needs; the accounting
|
||||
# must survive the re-walk.
|
||||
tree.cache_finished_req(req, kv_len_to_handle=num_tokens)
|
||||
tree.cache_finished_req(req, owned_kv_len=num_tokens)
|
||||
self.assertEqual(allocator.swa_available_size(), swa_before)
|
||||
self.assertEqual(
|
||||
tree.swa_evictable_size_ + tree.swa_protected_size_,
|
||||
|
||||
@@ -662,7 +662,7 @@ def bench_cache_finished(
|
||||
"cache_finished",
|
||||
lambda: req_items,
|
||||
lambda req: env.tree.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.kv.kv_committed_len
|
||||
req, is_insert=True, owned_kv_len=req.kv.kv_committed_len
|
||||
),
|
||||
len(req_items) - warmup,
|
||||
env.avg_tokens,
|
||||
|
||||
@@ -1668,9 +1668,7 @@ class UnifiedRadixCacheSuite:
|
||||
if self.cfg.has_mamba:
|
||||
req.kv.mamba_last_track_seqlen = kv_len
|
||||
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len())
|
||||
|
||||
all_ids = input_ids + output_ids
|
||||
aligned_len = (len(all_ids) // ps) * ps
|
||||
@@ -1732,9 +1730,9 @@ class UnifiedRadixCacheSuite:
|
||||
with get_serving().override(strip_thinking_cache=True):
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
req, is_insert=True, owned_kv_len=req.owned_kv_len()
|
||||
)
|
||||
start_p, end_p = req.effective_kv_committed_len(), req.kv.kv_allocated_len
|
||||
start_p, end_p = req.owned_kv_len(), req.kv.kv_allocated_len
|
||||
if ps > 1:
|
||||
start_p = ((start_p + ps - 1) // ps) * ps
|
||||
if start_p < end_p:
|
||||
@@ -1775,9 +1773,7 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=False, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
cache.cache_finished_req(req, is_insert=False, owned_kv_len=req.owned_kv_len())
|
||||
|
||||
self.assertEqual(allocator.available_size(), avail_before + kv_len)
|
||||
m = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
|
||||
@@ -1952,9 +1948,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.mamba_last_track_seqlen = kv_len
|
||||
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len())
|
||||
|
||||
self.assertEqual(allocator.available_size(), avail_before + tail_extra)
|
||||
aligned = input_ids[: (len(input_ids) // ps) * ps]
|
||||
@@ -8408,9 +8402,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
||||
)
|
||||
req.last_node = cache.root_node_handle()
|
||||
|
||||
cache.cache_finished_req(
|
||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||
)
|
||||
cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len())
|
||||
|
||||
def test_finished_req_stores_radix_mamba_state_in_int8_pool(self):
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
@@ -35,7 +35,7 @@ class TestUnifiedRadixLockRefScenarios(unittest.TestCase):
|
||||
swa_prefix_lock_released=False,
|
||||
)
|
||||
|
||||
cache.cache_finished_req(req, is_insert=False, kv_len_to_handle=3)
|
||||
cache.cache_finished_req(req, is_insert=False, owned_kv_len=3)
|
||||
|
||||
cache.free_kv_row.assert_called_once_with(kv, [(0, 3)])
|
||||
cache._dec_req_lock.assert_not_called()
|
||||
|
||||
Reference in New Issue
Block a user