From a42a7654a261d4af4b847340cb9f16b1ef7c7c44 Mon Sep 17 00:00:00 2001 From: Chizheng Fang <93508110+fcczzz@users.noreply.github.com> Date: Fri, 29 May 2026 15:18:05 +0800 Subject: [PATCH] Update MooncakeStore batch tests to use v1 APIs (#25880) Signed-off-by: fangchizheng --- .../mooncake_store/test_mooncake_store.py | 158 +++++++++--------- 1 file changed, 82 insertions(+), 76 deletions(-) diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/test_mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/test_mooncake_store.py index ae4788cb5..9929ea228 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/test_mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/test_mooncake_store.py @@ -2,9 +2,9 @@ import logging import uuid import torch -from mooncake_store import MooncakeStore from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig +from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import MooncakeStore logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s" @@ -12,44 +12,70 @@ logging.basicConfig( logger = logging.getLogger(__name__) -def generate_batch_query_keys(kv_num: int, config: HiCacheStorageConfig): - keys = [] - for _ in range(kv_num): - key = "test_" + str(uuid.uuid4()) - keys.append(key) - set_keys = [] - for key in keys: - if config.is_mla_model: - set_keys.append(key + "_k") - else: - set_keys.append(key + f"_{config.tp_rank}_k") - set_keys.append(key + f"_{config.tp_rank}_v") - get_keys = set_keys - exist_keys = keys - return set_keys, get_keys, exist_keys +def make_hicache_storage_config( + *, + is_mla_model: bool, + tp_rank: int, + tp_size: int, +) -> HiCacheStorageConfig: + return HiCacheStorageConfig( + tp_rank=tp_rank, + tp_size=tp_size, + pp_rank=0, + pp_size=1, + attn_cp_rank=0, + attn_cp_size=1, + is_mla_model=is_mla_model, + enable_storage_metrics=False, + is_page_first_layout=True, + model_name=None, + ) -def create_mock_host_kv_cache(buffer_size, dtype=torch.float32): +def generate_batch_query_keys(kv_num: int): + return ["test_" + str(uuid.uuid4()) for _ in range(kv_num)] + + +def create_mock_host_kv_cache( + buffer_size, + entries_per_page=2, + page_elements=1, + dtype=torch.float32, +): """Create a mock HostKVCache-like object for testing.""" buffer = torch.randn(buffer_size, dtype=dtype) class MockHostKVCache: - def __init__(self, buffer): + def __init__(self, buffer, entries_per_page, page_elements): self.kv_buffer = buffer self.layout = "page_first" self.page_size = 1 # Simple page size for testing + self.entries_per_page = entries_per_page + self.page_elements = page_elements def get_page_buffer_meta(self, indices): """Mock implementation of get_page_buffer_meta.""" ptr_list = [] element_size_list = [] for idx in indices: - # Create mock pointers and sizes for each page - ptr_list.append(idx * self.page_size * self.kv_buffer.element_size()) - element_size_list.append(self.page_size * self.kv_buffer.element_size()) + page_idx = int(idx) + page_offset = page_idx * self.entries_per_page * self.page_elements + for entry_idx in range(self.entries_per_page): + offset = page_offset + entry_idx * self.page_elements + ptr_list.append(self.kv_buffer[offset:].data_ptr()) + element_size_list.append( + self.page_elements * self.kv_buffer.element_size() + ) return ptr_list, element_size_list - return MockHostKVCache(buffer), buffer + def get_ksize_per_token(self): + return ( + self.entries_per_page + * self.page_elements + * self.kv_buffer.element_size() + ) + + return MockHostKVCache(buffer, entries_per_page, page_elements), buffer def test_single_operation(): @@ -59,8 +85,14 @@ def test_single_operation(): buffer_size = 1024 * 1024 * 16 # 16MB value_elements = 1024 - store = MooncakeStore() - mock_host_kv_cache, buffer = create_mock_host_kv_cache(buffer_size) + store = MooncakeStore( + make_hicache_storage_config(is_mla_model=False, tp_rank=0, tp_size=1) + ) + mock_host_kv_cache, buffer = create_mock_host_kv_cache( + buffer_size, + entries_per_page=2, + page_elements=value_elements, + ) # Register the memory pool host - this is the proper workflow store.register_mem_pool_host(mock_host_kv_cache) @@ -100,50 +132,48 @@ def test_batch_operation(config: HiCacheStorageConfig): buffer_size = 1024 * 1024 * 16 # 16MB value_elements = 256 kv_num = 13 + entries_per_page = 1 if config.is_mla_model else 2 store = MooncakeStore(config) - mock_host_kv_cache, buffer = create_mock_host_kv_cache(buffer_size) + mock_host_kv_cache, buffer = create_mock_host_kv_cache( + buffer_size, + entries_per_page=entries_per_page, + page_elements=value_elements, + ) store.register_mem_pool_host(mock_host_kv_cache) - value_size = value_elements * buffer.element_size() - - set_keys, get_keys, exist_keys = generate_batch_query_keys(kv_num, config) + keys = generate_batch_query_keys(kv_num) set_slices = [ buffer[i * value_elements : (i + 1) * value_elements] - for i in range(len(set_keys)) + for i in range(kv_num * entries_per_page) ] - set_locations = [set_slice.data_ptr() for set_slice in set_slices] - target_sizes = [value_size for _ in range(len(set_keys))] + set_indices = torch.arange(kv_num) # Test batch set operation - result = store.batch_set( - set_keys, target_locations=set_locations, target_sizes=target_sizes - ) - assert result is True, f"❌batch set operation failed" + result = store.batch_set_v1(keys, set_indices) + assert all(result), "batch set operation failed" # Test batch exists operation - assert store.batch_exists( - exist_keys - ), f"❌keys should exist after batch set operation" + assert ( + store.batch_exists(keys) == kv_num + ), "keys should exist after batch set operation" # Test batch get operation get_slices = [ buffer[ - (len(set_keys) + i) - * value_elements : (len(set_keys) + i + 1) + (kv_num * entries_per_page + i) + * value_elements : (kv_num * entries_per_page + i + 1) * value_elements ] - for i in range(len(get_keys)) + for i in range(kv_num * entries_per_page) ] - get_locations = [get_slice.data_ptr() for get_slice in get_slices] - result = store.batch_get( - get_keys, target_locations=get_locations, target_sizes=target_sizes - ) - assert result == kv_num, f"❌batch get operation failed" - for i in range(len(get_keys)): + get_indices = torch.arange(kv_num, 2 * kv_num) + result = store.batch_get_v1(keys, get_indices) + assert all(result), "❌batch get operation failed" + for i in range(kv_num * entries_per_page): assert torch.allclose( set_slices[i], get_slices[i], atol=1e-6 - ), f"❌batch get operation failed for key: {get_keys[i]}" + ), f"❌batch get operation failed for key: {keys[i // entries_per_page]}" logger.info(f"✅ Batch operation passed") @@ -151,39 +181,15 @@ def test_batch_operation(config: HiCacheStorageConfig): if __name__ == "__main__": test_single_operation() test_batch_operation( - HiCacheStorageConfig( - is_mla_model=False, - tp_rank=0, - tp_size=1, - model_name=None, - is_page_first_layout=True, - ) + make_hicache_storage_config(is_mla_model=False, tp_rank=0, tp_size=1) ) test_batch_operation( - HiCacheStorageConfig( - is_mla_model=True, - tp_rank=0, - tp_size=1, - model_name=None, - is_page_first_layout=True, - ) + make_hicache_storage_config(is_mla_model=True, tp_rank=0, tp_size=1) ) test_batch_operation( - HiCacheStorageConfig( - is_mla_model=False, - tp_rank=1, - tp_size=4, - model_name=None, - is_page_first_layout=True, - ) + make_hicache_storage_config(is_mla_model=False, tp_rank=1, tp_size=4) ) test_batch_operation( - HiCacheStorageConfig( - is_mla_model=True, - tp_rank=3, - tp_size=8, - model_name=None, - is_page_first_layout=True, - ) + make_hicache_storage_config(is_mla_model=True, tp_rank=3, tp_size=8) ) logger.info(f"✅ All tests passed")