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 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")