477 lines
15 KiB
Python
477 lines
15 KiB
Python
import types
|
|
import unittest
|
|
from unittest.mock import patch
|
|
|
|
import torch
|
|
|
|
from sglang.srt.mem_cache.hicache_storage import (
|
|
HiCacheStorageConfig,
|
|
PoolName,
|
|
PoolTransfer,
|
|
)
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
from sglang.test.test_utils import CustomTestCase
|
|
|
|
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
|
|
|
|
|
class ReplicateConfigWithGroupIds:
|
|
def __init__(self):
|
|
self.group_ids = None
|
|
|
|
|
|
class ReplicateConfigWithoutGroupIds:
|
|
__slots__ = ()
|
|
|
|
|
|
class ReplicateConfigWithClassGroupIdsAndRequiredInit:
|
|
group_ids = None
|
|
|
|
def __init__(self, required):
|
|
self.group_ids = required
|
|
|
|
|
|
def _fake_mooncake_modules(fake_store_cls, replicate_config_cls):
|
|
mooncake = types.ModuleType("mooncake")
|
|
mooncake_store = types.ModuleType("mooncake.store")
|
|
mooncake_store.MooncakeDistributedStore = fake_store_cls
|
|
mooncake_store.ReplicateConfig = replicate_config_cls
|
|
return {
|
|
"mooncake": mooncake,
|
|
"mooncake.store": mooncake_store,
|
|
}
|
|
|
|
|
|
def _fake_pool_host_mla_module():
|
|
pool_host_mla = types.ModuleType("sglang.srt.mem_cache.pool_host.mla")
|
|
|
|
class MLATokenToKVPoolHost:
|
|
pass
|
|
|
|
pool_host_mla.MLATokenToKVPoolHost = MLATokenToKVPoolHost
|
|
return pool_host_mla
|
|
|
|
|
|
def _fake_pool_host_module():
|
|
pool_host = types.ModuleType("sglang.srt.mem_cache.pool_host")
|
|
|
|
class HostKVCache:
|
|
pass
|
|
|
|
class HostTensorAllocator:
|
|
pass
|
|
|
|
pool_host.HostKVCache = HostKVCache
|
|
pool_host.HostTensorAllocator = HostTensorAllocator
|
|
return pool_host
|
|
|
|
|
|
def _fake_host_pool_modules():
|
|
return {
|
|
"sglang.srt.mem_cache.pool_host": _fake_pool_host_module(),
|
|
"sglang.srt.mem_cache.pool_host.mla": _fake_pool_host_mla_module(),
|
|
}
|
|
|
|
|
|
def _fake_store_class():
|
|
class FakeMooncakeDistributedStore:
|
|
instances = []
|
|
|
|
def __init__(self):
|
|
self.batch_put_calls = []
|
|
self.existing_keys = set()
|
|
self.objects = {}
|
|
type(self).instances.append(self)
|
|
|
|
def setup(self, *args, **kwargs):
|
|
return 0
|
|
|
|
def register_buffer(self, *args, **kwargs):
|
|
return 0
|
|
|
|
def put(self, key, value, *args):
|
|
self.objects[key] = value
|
|
return 0
|
|
|
|
def is_exist(self, key):
|
|
return 1 if key in self.objects or key in self.existing_keys else 0
|
|
|
|
def get(self, key):
|
|
return self.objects.get(key)
|
|
|
|
def batch_is_exist(self, keys):
|
|
return [1 if key in self.existing_keys else 0 for key in keys]
|
|
|
|
def batch_put_from(self, keys, ptrs, sizes, *args):
|
|
self.batch_put_calls.append(
|
|
{
|
|
"method": "batch_put_from",
|
|
"keys": list(keys),
|
|
"ptrs": list(ptrs),
|
|
"sizes": list(sizes),
|
|
"args": args,
|
|
}
|
|
)
|
|
self.existing_keys.update(keys)
|
|
return [0] * len(keys)
|
|
|
|
def batch_put_from_multi_buffers(self, keys, ptrs, sizes, *args):
|
|
self.batch_put_calls.append(
|
|
{
|
|
"method": "batch_put_from_multi_buffers",
|
|
"keys": list(keys),
|
|
"ptrs": list(ptrs),
|
|
"sizes": list(sizes),
|
|
"args": args,
|
|
}
|
|
)
|
|
self.existing_keys.update(keys)
|
|
return [0] * len(keys)
|
|
|
|
return FakeMooncakeDistributedStore
|
|
|
|
|
|
class FakeHostKVCache:
|
|
def __init__(self, objects_per_page):
|
|
self.objects_per_page = objects_per_page
|
|
self.kv_buffer = torch.empty((1024,), dtype=torch.uint8)
|
|
self.layout = "page_first"
|
|
self.page_size = 1
|
|
|
|
def get_ksize_per_token(self):
|
|
return 1
|
|
|
|
def get_page_buffer_meta(self, indices):
|
|
page_count = len(indices) // self.page_size
|
|
ptrs = []
|
|
sizes = []
|
|
for page_idx in range(page_count):
|
|
for object_idx in range(self.objects_per_page):
|
|
ptrs.append(1000 + page_idx * 100 + object_idx)
|
|
sizes.append(8)
|
|
return ptrs, sizes
|
|
|
|
def get_split_heads_page_buffer_meta(self, indices, split_factor):
|
|
page_count = len(indices) // self.page_size
|
|
ptrs = []
|
|
sizes = []
|
|
for page_idx in range(page_count):
|
|
for object_idx in range(2 * split_factor):
|
|
ptrs.append(2000 + page_idx * 100 + object_idx)
|
|
sizes.append(8)
|
|
return ptrs, sizes
|
|
|
|
|
|
class FakeIndexerPool:
|
|
page_size = 1
|
|
|
|
def __init__(self):
|
|
self.buffer = torch.empty((128,), dtype=torch.uint8)
|
|
|
|
def get_hybrid_pool_buffer(self):
|
|
return [self.buffer]
|
|
|
|
def get_page_buffer_meta(self, indices):
|
|
return [3000 + i for i in range(len(indices))], [8] * len(indices)
|
|
|
|
|
|
class FakeMultiBufferPool:
|
|
page_size = 1
|
|
|
|
def __init__(self):
|
|
self.buffers = [
|
|
torch.empty((128,), dtype=torch.uint8),
|
|
torch.empty((128,), dtype=torch.uint8),
|
|
]
|
|
|
|
def get_hybrid_pool_buffer(self):
|
|
return self.buffers
|
|
|
|
def get_page_buffer_meta(self, indices):
|
|
ptrs = []
|
|
sizes = []
|
|
for i in range(len(indices)):
|
|
ptrs.extend([4000 + i * 10, 4001 + i * 10])
|
|
sizes.extend([8, 16])
|
|
return ptrs, sizes
|
|
|
|
|
|
def _make_config(
|
|
*,
|
|
enable_group_semantics=True,
|
|
extra_backend_tag=None,
|
|
model_name=None,
|
|
is_mla_model=False,
|
|
should_split_heads=False,
|
|
tp_rank=0,
|
|
tp_size=1,
|
|
tp_lcm_size=None,
|
|
):
|
|
extra_config = {
|
|
"master_server_address": "127.0.0.1:50051",
|
|
"check_server": False,
|
|
"global_segment_size": 1024 * 1024,
|
|
"enable_group_semantics": enable_group_semantics,
|
|
}
|
|
if extra_backend_tag is not None:
|
|
extra_config["extra_backend_tag"] = extra_backend_tag
|
|
|
|
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=model_name,
|
|
tp_lcm_size=tp_lcm_size,
|
|
should_split_heads=should_split_heads,
|
|
extra_config=extra_config,
|
|
)
|
|
|
|
|
|
def _make_store(
|
|
*,
|
|
enable_group_semantics=True,
|
|
replicate_config_cls=ReplicateConfigWithGroupIds,
|
|
extra_backend_tag=None,
|
|
model_name=None,
|
|
is_mla_model=False,
|
|
should_split_heads=False,
|
|
tp_rank=0,
|
|
tp_size=1,
|
|
tp_lcm_size=None,
|
|
):
|
|
fake_store_cls = _fake_store_class()
|
|
cfg = _make_config(
|
|
enable_group_semantics=enable_group_semantics,
|
|
extra_backend_tag=extra_backend_tag,
|
|
model_name=model_name,
|
|
is_mla_model=is_mla_model,
|
|
should_split_heads=should_split_heads,
|
|
tp_rank=tp_rank,
|
|
tp_size=tp_size,
|
|
tp_lcm_size=tp_lcm_size,
|
|
)
|
|
|
|
with patch.dict(
|
|
"sys.modules",
|
|
{
|
|
**_fake_mooncake_modules(fake_store_cls, replicate_config_cls),
|
|
**_fake_host_pool_modules(),
|
|
},
|
|
):
|
|
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
|
|
MooncakeStore,
|
|
)
|
|
|
|
store = MooncakeStore(cfg)
|
|
|
|
return store, fake_store_cls.instances[-1]
|
|
|
|
|
|
class TestMooncakeGroupSemantics(CustomTestCase):
|
|
def test_group_id_detection_uses_class_attribute_without_instantiating(self):
|
|
fake_store_cls = _fake_store_class()
|
|
with patch.dict(
|
|
"sys.modules",
|
|
{
|
|
**_fake_mooncake_modules(
|
|
fake_store_cls,
|
|
ReplicateConfigWithClassGroupIdsAndRequiredInit,
|
|
),
|
|
**_fake_host_pool_modules(),
|
|
},
|
|
):
|
|
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
|
|
MooncakeBaseStore,
|
|
)
|
|
|
|
replicate_config_cls, supports_group_ids = (
|
|
MooncakeBaseStore()._import_mooncake_group_semantics()
|
|
)
|
|
|
|
self.assertIs(
|
|
replicate_config_cls, ReplicateConfigWithClassGroupIdsAndRequiredInit
|
|
)
|
|
self.assertTrue(supports_group_ids)
|
|
|
|
def test_flag_off_uses_three_arg_batch_put(self):
|
|
store, fake_store = _make_store(enable_group_semantics=False)
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
|
|
result = store.batch_set_v1(["page0"], torch.tensor([0]))
|
|
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(len(fake_store.batch_put_calls), 1)
|
|
self.assertEqual(
|
|
fake_store.batch_put_calls[0]["keys"], ["page0_0_k", "page0_0_v"]
|
|
)
|
|
self.assertEqual(fake_store.batch_put_calls[0]["args"], ())
|
|
|
|
def test_mha_group_ids_use_tagged_logical_page_key(self):
|
|
store, fake_store = _make_store(extra_backend_tag="tag")
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
|
|
result = store.batch_set_v1(["page0", "page1"], torch.tensor([0, 1]))
|
|
|
|
self.assertEqual(result, [True, True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(
|
|
call["keys"],
|
|
[
|
|
"tag_page0_0_k",
|
|
"tag_page0_0_v",
|
|
"tag_page1_0_k",
|
|
"tag_page1_0_v",
|
|
],
|
|
)
|
|
self.assertEqual(len(call["args"]), 1)
|
|
self.assertEqual(
|
|
call["args"][0].group_ids,
|
|
[
|
|
"sglang-hicache:tag_page0",
|
|
"sglang-hicache:tag_page0",
|
|
"sglang-hicache:tag_page1",
|
|
"sglang-hicache:tag_page1",
|
|
],
|
|
)
|
|
|
|
def test_old_mooncake_falls_back_to_three_arg_batch_put(self):
|
|
store, fake_store = _make_store(
|
|
enable_group_semantics=True,
|
|
replicate_config_cls=ReplicateConfigWithoutGroupIds,
|
|
)
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
|
|
result = store.batch_set_v1(["page0"], torch.tensor([0]))
|
|
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(fake_store.batch_put_calls[0]["args"], ())
|
|
|
|
def test_mla_group_ids(self):
|
|
store, fake_store = _make_store(is_mla_model=True)
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=1))
|
|
|
|
result = store.batch_set_v1(["page0"], torch.tensor([0]))
|
|
|
|
self.assertEqual(result, [True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(call["keys"], ["page0__k"])
|
|
self.assertEqual(call["args"][0].group_ids, ["sglang-hicache:page0"])
|
|
|
|
def test_split_heads_group_ids(self):
|
|
store, fake_store = _make_store(
|
|
should_split_heads=True,
|
|
tp_rank=1,
|
|
tp_size=2,
|
|
tp_lcm_size=4,
|
|
)
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=4))
|
|
|
|
result = store.batch_set_v1(["page0"], torch.tensor([0]))
|
|
|
|
self.assertEqual(result, [True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(
|
|
call["keys"],
|
|
["page0_2_k", "page0_2_v", "page0_3_k", "page0_3_v"],
|
|
)
|
|
self.assertEqual(
|
|
call["args"][0].group_ids,
|
|
["sglang-hicache:page0"] * 4,
|
|
)
|
|
|
|
def test_existing_filter_keeps_group_ids_aligned_with_missing_keys(self):
|
|
store, fake_store = _make_store()
|
|
store.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
fake_store.existing_keys.update({"page0_0_k", "page1_0_v"})
|
|
|
|
result = store.batch_set_v1(["page0", "page1"], torch.tensor([0, 1]))
|
|
|
|
self.assertEqual(result, [True, True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(call["keys"], ["page0_0_v", "page1_0_k"])
|
|
self.assertEqual(
|
|
call["args"][0].group_ids,
|
|
["sglang-hicache:page0", "sglang-hicache:page1"],
|
|
)
|
|
|
|
def test_v2_indexer_group_ids_use_logical_page_key(self):
|
|
store, fake_store = _make_store(extra_backend_tag="tag", is_mla_model=True)
|
|
indexer_pool = FakeIndexerPool()
|
|
store.register_mem_host_pool_v2(indexer_pool, PoolName.INDEXER)
|
|
|
|
result = store.batch_set_v2(
|
|
[
|
|
PoolTransfer(
|
|
name=PoolName.INDEXER,
|
|
keys=["page0", "page1"],
|
|
host_indices=torch.tensor([0, 1]),
|
|
)
|
|
]
|
|
)
|
|
|
|
self.assertEqual(result[PoolName.INDEXER], [True, True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(call["keys"], ["tag_page0__indexer", "tag_page1__indexer"])
|
|
self.assertEqual(
|
|
call["args"][0].group_ids,
|
|
["sglang-hicache:tag_page0", "sglang-hicache:tag_page1"],
|
|
)
|
|
|
|
def test_v2_multi_buffer_put_passes_group_ids(self):
|
|
store, fake_store = _make_store(extra_backend_tag="tag", is_mla_model=True)
|
|
multi_buffer_pool = FakeMultiBufferPool()
|
|
store.register_mem_host_pool_v2(multi_buffer_pool, PoolName.DEEPSEEK_V4_C4)
|
|
|
|
result = store.batch_set_v2(
|
|
[
|
|
PoolTransfer(
|
|
name=PoolName.DEEPSEEK_V4_C4,
|
|
keys=["page0", "page1"],
|
|
host_indices=torch.tensor([0, 1]),
|
|
)
|
|
]
|
|
)
|
|
|
|
self.assertEqual(result[PoolName.DEEPSEEK_V4_C4], [True, True])
|
|
call = fake_store.batch_put_calls[0]
|
|
self.assertEqual(call["method"], "batch_put_from_multi_buffers")
|
|
self.assertEqual(
|
|
call["keys"],
|
|
["tag_page0__deepseek_v4_c4", "tag_page1__deepseek_v4_c4"],
|
|
)
|
|
self.assertEqual(call["ptrs"], [[4000, 4001], [4010, 4011]])
|
|
self.assertEqual(call["sizes"], [[8, 16], [8, 16]])
|
|
self.assertEqual(
|
|
call["args"][0].group_ids,
|
|
["sglang-hicache:tag_page0", "sglang-hicache:tag_page1"],
|
|
)
|
|
|
|
def test_model_names_isolate_the_same_logical_key(self):
|
|
store_a, fake_store_a = _make_store(
|
|
enable_group_semantics=False, model_name="org/model-a"
|
|
)
|
|
store_b, fake_store_b = _make_store(
|
|
enable_group_semantics=False, model_name="org/model-b"
|
|
)
|
|
store_a.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
store_b.register_mem_pool_host(FakeHostKVCache(objects_per_page=2))
|
|
|
|
self.assertEqual(store_a.batch_set_v1(["page0"], torch.tensor([0])), [True])
|
|
self.assertEqual(store_b.batch_set_v1(["page0"], torch.tensor([0])), [True])
|
|
|
|
keys_a = fake_store_a.batch_put_calls[0]["keys"]
|
|
keys_b = fake_store_b.batch_put_calls[0]["keys"]
|
|
self.assertEqual(keys_a, ["org-model-a_page0_0_k", "org-model-a_page0_0_v"])
|
|
self.assertEqual(keys_b, ["org-model-b_page0_0_k", "org-model-b_page0_0_v"])
|
|
self.assertTrue(set(keys_a).isdisjoint(keys_b))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main(verbosity=3)
|