Refactor cache init logic (#13800)

This commit is contained in:
Liangsheng Yin
2025-11-24 11:41:46 +08:00
committed by GitHub
parent 75222bfed9
commit b2f7b08c49
17 changed files with 307 additions and 419 deletions

View File

@@ -240,11 +240,9 @@ class TestRadixCache(unittest.TestCase):
with self.subTest(
page_size=page_size, disable=disable, enable_events=enable_events
):
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=page_size,
cache = RadixCache.create_simulated(
disable=disable,
page_size=page_size,
enable_kv_cache_events=enable_events,
)
@@ -257,9 +255,7 @@ class TestRadixCache(unittest.TestCase):
def test_reset(self):
"""Test reset method."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
# Insert some data
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
@@ -275,12 +271,7 @@ class TestRadixCache(unittest.TestCase):
"""Test basic insert and match operations."""
for disable_cache in [False, True]:
with self.subTest(disable_cache=disable_cache):
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=1,
disable=disable_cache,
)
cache = RadixCache.create_simulated(disable=disable_cache)
key = RadixKey([1, 2, 3])
value = torch.tensor([10, 20, 30], dtype=torch.int64)
@@ -309,9 +300,7 @@ class TestRadixCache(unittest.TestCase):
def test_insert_with_none_value(self):
"""Test insert with None value (should use token_ids as list)."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
key = RadixKey([1, 2, 3])
prefix_len = cache.insert(key, None)
@@ -322,9 +311,7 @@ class TestRadixCache(unittest.TestCase):
def test_total_size(self):
"""Test total_size calculation."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
self.assertEqual(cache.total_size(), 0)
@@ -344,11 +331,8 @@ class TestRadixCache(unittest.TestCase):
for page_size, enable_events in test_cases:
with self.subTest(page_size=page_size, enable_events=enable_events):
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=page_size,
enable_kv_cache_events=enable_events,
cache = RadixCache.create_simulated(
page_size=page_size, enable_kv_cache_events=enable_events
)
# Insert data
@@ -374,11 +358,8 @@ class TestRadixCache(unittest.TestCase):
mock_allocator = unittest.mock.Mock()
mock_allocator.device = torch.device("cpu")
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=mock_allocator,
page_size=1,
enable_kv_cache_events=True,
cache = RadixCache.create_simulated(
mock_allocator=mock_allocator, enable_kv_cache_events=True
)
# Insert and then evict data
@@ -400,9 +381,7 @@ class TestRadixCache(unittest.TestCase):
def test_extra_key_isolation(self):
"""Test that keys with different extra_key values are isolated."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
# Insert same token sequence with different extra keys
cache.insert(
@@ -442,9 +421,7 @@ class TestRadixCache(unittest.TestCase):
def test_lock_ref_operations(self):
"""Test lock reference counting operations."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
# Insert sequence
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
@@ -471,11 +448,7 @@ class TestRadixCache(unittest.TestCase):
mock_allocator = unittest.mock.Mock()
mock_allocator.device = torch.device("cpu")
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=mock_allocator,
page_size=1,
)
cache = RadixCache.create_simulated(mock_allocator=mock_allocator)
# Insert sequences
cache.insert(RadixKey([1, 2]), torch.tensor([10, 20], dtype=torch.int64))
@@ -500,11 +473,7 @@ class TestRadixCache(unittest.TestCase):
for page_size, sequence_length in test_cases:
with self.subTest(page_size=page_size, sequence_length=sequence_length):
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=page_size,
)
cache = RadixCache.create_simulated(page_size=page_size)
tokens = list(range(sequence_length))
cache.insert(RadixKey(tokens), torch.tensor(tokens, dtype=torch.int64))
@@ -518,9 +487,7 @@ class TestRadixCache(unittest.TestCase):
def test_pretty_print_basic(self):
"""Test pretty_print produces output."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
@@ -532,9 +499,7 @@ class TestRadixCache(unittest.TestCase):
def test_all_values_flatten(self):
"""Test all_values_flatten method."""
cache = RadixCache(
req_to_token_pool=None, token_to_kv_pool_allocator=None, page_size=1
)
cache = RadixCache.create_simulated()
cache.insert(RadixKey([1, 2]), torch.tensor([10, 20], dtype=torch.int64))
cache.insert(RadixKey([3, 4]), torch.tensor([30, 40], dtype=torch.int64))
@@ -549,11 +514,7 @@ class TestRadixCache(unittest.TestCase):
"""Advanced prefix matching: splits inside nodes and across pages."""
for page_size in [1, 2]:
with self.subTest(page_size=page_size):
cache = RadixCache(
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=page_size,
)
cache = RadixCache.create_simulated(page_size=page_size)
# Insert a long sequence that will be split later.
seq1 = [1, 2, 3, 4, 5, 6, 7, 8]