Update MooncakeStore batch tests to use v1 APIs (#25880)

Signed-off-by: fangchizheng <fangchizheng@mail.ustc.edu.cn>
This commit is contained in:
Chizheng Fang
2026-05-29 00:18:05 -07:00
committed by GitHub
parent ace730db48
commit a42a7654a2
@@ -2,9 +2,9 @@ import logging
import uuid import uuid
import torch import torch
from mooncake_store import MooncakeStore
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import MooncakeStore
logging.basicConfig( logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s" level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
@@ -12,44 +12,70 @@ logging.basicConfig(
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def generate_batch_query_keys(kv_num: int, config: HiCacheStorageConfig): def make_hicache_storage_config(
keys = [] *,
for _ in range(kv_num): is_mla_model: bool,
key = "test_" + str(uuid.uuid4()) tp_rank: int,
keys.append(key) tp_size: int,
set_keys = [] ) -> HiCacheStorageConfig:
for key in keys: return HiCacheStorageConfig(
if config.is_mla_model: tp_rank=tp_rank,
set_keys.append(key + "_k") tp_size=tp_size,
else: pp_rank=0,
set_keys.append(key + f"_{config.tp_rank}_k") pp_size=1,
set_keys.append(key + f"_{config.tp_rank}_v") attn_cp_rank=0,
get_keys = set_keys attn_cp_size=1,
exist_keys = keys is_mla_model=is_mla_model,
return set_keys, get_keys, exist_keys 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.""" """Create a mock HostKVCache-like object for testing."""
buffer = torch.randn(buffer_size, dtype=dtype) buffer = torch.randn(buffer_size, dtype=dtype)
class MockHostKVCache: class MockHostKVCache:
def __init__(self, buffer): def __init__(self, buffer, entries_per_page, page_elements):
self.kv_buffer = buffer self.kv_buffer = buffer
self.layout = "page_first" self.layout = "page_first"
self.page_size = 1 # Simple page size for testing 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): def get_page_buffer_meta(self, indices):
"""Mock implementation of get_page_buffer_meta.""" """Mock implementation of get_page_buffer_meta."""
ptr_list = [] ptr_list = []
element_size_list = [] element_size_list = []
for idx in indices: for idx in indices:
# Create mock pointers and sizes for each page page_idx = int(idx)
ptr_list.append(idx * self.page_size * self.kv_buffer.element_size()) page_offset = page_idx * self.entries_per_page * self.page_elements
element_size_list.append(self.page_size * self.kv_buffer.element_size()) 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 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(): def test_single_operation():
@@ -59,8 +85,14 @@ def test_single_operation():
buffer_size = 1024 * 1024 * 16 # 16MB buffer_size = 1024 * 1024 * 16 # 16MB
value_elements = 1024 value_elements = 1024
store = MooncakeStore() store = MooncakeStore(
mock_host_kv_cache, buffer = create_mock_host_kv_cache(buffer_size) 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 # Register the memory pool host - this is the proper workflow
store.register_mem_pool_host(mock_host_kv_cache) 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 buffer_size = 1024 * 1024 * 16 # 16MB
value_elements = 256 value_elements = 256
kv_num = 13 kv_num = 13
entries_per_page = 1 if config.is_mla_model else 2
store = MooncakeStore(config) 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) store.register_mem_pool_host(mock_host_kv_cache)
value_size = value_elements * buffer.element_size() keys = generate_batch_query_keys(kv_num)
set_keys, get_keys, exist_keys = generate_batch_query_keys(kv_num, config)
set_slices = [ set_slices = [
buffer[i * value_elements : (i + 1) * value_elements] 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] set_indices = torch.arange(kv_num)
target_sizes = [value_size for _ in range(len(set_keys))]
# Test batch set operation # Test batch set operation
result = store.batch_set( result = store.batch_set_v1(keys, set_indices)
set_keys, target_locations=set_locations, target_sizes=target_sizes assert all(result), "batch set operation failed"
)
assert result is True, f"❌batch set operation failed"
# Test batch exists operation # Test batch exists operation
assert store.batch_exists( assert (
exist_keys store.batch_exists(keys) == kv_num
), f"❌keys should exist after batch set operation" ), "keys should exist after batch set operation"
# Test batch get operation # Test batch get operation
get_slices = [ get_slices = [
buffer[ buffer[
(len(set_keys) + i) (kv_num * entries_per_page + i)
* value_elements : (len(set_keys) + i + 1) * value_elements : (kv_num * entries_per_page + i + 1)
* value_elements * 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] get_indices = torch.arange(kv_num, 2 * kv_num)
result = store.batch_get( result = store.batch_get_v1(keys, get_indices)
get_keys, target_locations=get_locations, target_sizes=target_sizes assert all(result), "❌batch get operation failed"
) for i in range(kv_num * entries_per_page):
assert result == kv_num, f"❌batch get operation failed"
for i in range(len(get_keys)):
assert torch.allclose( assert torch.allclose(
set_slices[i], get_slices[i], atol=1e-6 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") logger.info(f"✅ Batch operation passed")
@@ -151,39 +181,15 @@ def test_batch_operation(config: HiCacheStorageConfig):
if __name__ == "__main__": if __name__ == "__main__":
test_single_operation() test_single_operation()
test_batch_operation( test_batch_operation(
HiCacheStorageConfig( make_hicache_storage_config(is_mla_model=False, tp_rank=0, tp_size=1)
is_mla_model=False,
tp_rank=0,
tp_size=1,
model_name=None,
is_page_first_layout=True,
)
) )
test_batch_operation( test_batch_operation(
HiCacheStorageConfig( make_hicache_storage_config(is_mla_model=True, tp_rank=0, tp_size=1)
is_mla_model=True,
tp_rank=0,
tp_size=1,
model_name=None,
is_page_first_layout=True,
)
) )
test_batch_operation( test_batch_operation(
HiCacheStorageConfig( make_hicache_storage_config(is_mla_model=False, tp_rank=1, tp_size=4)
is_mla_model=False,
tp_rank=1,
tp_size=4,
model_name=None,
is_page_first_layout=True,
)
) )
test_batch_operation( test_batch_operation(
HiCacheStorageConfig( make_hicache_storage_config(is_mla_model=True, tp_rank=3, tp_size=8)
is_mla_model=True,
tp_rank=3,
tp_size=8,
model_name=None,
is_page_first_layout=True,
)
) )
logger.info(f"✅ All tests passed") logger.info(f"✅ All tests passed")