feat: add Mooncake group semantics (#26574)

Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
Co-authored-by: Teng Ma <sima.mt@alibaba-inc.com>
This commit is contained in:
Xingyuan Wu
2026-06-22 19:38:58 -07:00
committed by GitHub
co-authored by Zhiqiang Xie Teng Ma
parent 28d5627fd8
commit 62f7ffc492
3 changed files with 564 additions and 4 deletions
@@ -289,6 +289,22 @@ You can enable it in any of the three supported configuration methods:
> **Note:** `enable_ssd_offload` requires a Mooncake version that supports the `enable_ssd_offload` parameter in `MooncakeDistributedStore.setup()`. If the installed version does not support it, SGLang will automatically fall back to the old behavior and print a warning.
**Mooncake Group Semantics (`enable_group_semantics`):**
When `enable_group_semantics` is set to `true`, SGLang passes Mooncake `group_ids` for physical objects derived from the same logical HiCache page. This allows Mooncake to apply group-aware metadata routing, lease refresh, and eviction behavior to related KV objects such as MHA K/V pairs, split-head shards, MLA objects, and supported sidecar objects.
This option is disabled by default. It requires a Mooncake version that exposes `ReplicateConfig.group_ids`. If the installed Mooncake package does not support it, SGLang automatically falls back to the existing write path and prints a warning.
Example:
```bash
python -m sglang.launch_server \
--enable-hierarchical-cache \
--hicache-storage-backend mooncake \
--model-path [model_path] \
--hicache-storage-backend-extra-config '{"master_server_address": "127.0.0.1:50051", "enable_group_semantics": true}'
```
**HiCache Related Parameters for SGLang Server**
For a comprehensive overview of HiCache-related parameters, please refer to [this document](https://docs.sglang.io/advanced_features/hicache_design.html#related-parameters).
@@ -261,6 +261,20 @@ class MooncakeBaseStore:
"to run SGLang with MooncakeConnector."
) from e
def _import_mooncake_group_semantics(self):
try:
from mooncake.store import ReplicateConfig
except ImportError:
return None, False
supports_group_ids = hasattr(ReplicateConfig, "group_ids")
if not supports_group_ids:
try:
supports_group_ids = hasattr(ReplicateConfig(), "group_ids")
except Exception:
supports_group_ids = False
return ReplicateConfig, supports_group_ids
def _load_config(self, storage_config: Any = None):
extra_config = (
getattr(storage_config, "extra_config", None) if storage_config else None
@@ -349,6 +363,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
):
MooncakeBaseStore.__init__(self)
MooncakeDistributedStore = self._import_mooncake_store()
self._replicate_config_cls, self._supports_group_ids = (
self._import_mooncake_group_semantics()
)
try:
self.store = MooncakeDistributedStore()
@@ -358,6 +375,22 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
if storage_config
else None
)
self.enable_group_semantics = bool(
extra_config.get("enable_group_semantics", False)
if extra_config
else False
)
self._use_group_semantics = (
self.enable_group_semantics
and self._supports_group_ids
and self._replicate_config_cls is not None
)
if self.enable_group_semantics and not self._supports_group_ids:
logger.warning(
"Mooncake group semantics is enabled, but the installed "
"Mooncake package does not support ReplicateConfig.group_ids. "
"Falling back to the existing batch_put_from path."
)
tp_scale_factor = 1 if storage_config is None else storage_config.tp_size
per_tp_global_segment_size = (
@@ -648,6 +681,27 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
return keys
return [f"{self.extra_backend_tag}_{key}" for key in keys]
def _can_use_group_semantics(self) -> bool:
return self._use_group_semantics
def _make_group_id(self, logical_key: str) -> str:
return f"sglang-hicache:{logical_key}"
def _expand_group_ids(
self, logical_keys: List[str], key_multiplier: int
) -> List[str]:
group_ids = []
for key in logical_keys:
group_ids.extend([self._make_group_id(key)] * key_multiplier)
return group_ids
def _filter_group_ids(
self, group_ids: Optional[List[str]], indices: List[int]
) -> Optional[List[str]]:
if group_ids is None:
return None
return [group_ids[i] for i in indices]
def _get_hybrid_page_component_keys(
self, page_keys: List[str], transfer: PoolTransfer
) -> Tuple[List[str], int]:
@@ -779,6 +833,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
assert len(keys) > 0
assert len(keys) == len(host_indices) // page_size
tagged_keys = self._tag_keys(keys)
key_strs, key_multiplier = self._get_hybrid_page_component_keys(
keys, transfer
)
@@ -790,6 +845,11 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
)
if is_set:
group_ids = (
self._expand_group_ids(tagged_keys, key_multiplier)
if self._can_use_group_semantics()
else None
)
exist_result = self._batch_exist(key_strs)
io_results = [0 if state == 1 else -1 for state in exist_result]
missing_idx = [i for i, state in enumerate(exist_result) if state != 1]
@@ -798,6 +858,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
[key_strs[i] for i in missing_idx],
[ptr_list[i] for i in missing_idx],
[element_size_list[i] for i in missing_idx],
self._filter_group_ids(group_ids, missing_idx),
)
for i, res in zip(missing_idx, put_results):
io_results[i] = res
@@ -970,6 +1031,12 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
keys = self._tag_keys(keys)
key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices)
key_multiplier = len(key_strs) // len(keys)
group_ids = (
self._expand_group_ids(keys, key_multiplier)
if self._can_use_group_semantics()
else None
)
exist_result = self._batch_exist(key_strs)
set_keys = []
@@ -990,7 +1057,10 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
if len(set_keys) > 0:
start_time = time.perf_counter()
put_results = self._put_batch_zero_copy_impl(
set_keys, set_buffer_ptrs, set_buffer_sizes
set_keys,
set_buffer_ptrs,
set_buffer_sizes,
self._filter_group_ids(group_ids, set_indices),
)
end_time = time.perf_counter()
@@ -1166,13 +1236,33 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
self.store.remove_all()
def _put_batch_zero_copy_impl(
self, key_strs: List[str], buffer_ptrs: List[Any], buffer_sizes: List[Any]
self,
key_strs: List[str],
buffer_ptrs: List[Any],
buffer_sizes: List[Any],
group_ids: Optional[List[str]] = None,
) -> List[int]:
config = None
if self._can_use_group_semantics() and group_ids is not None:
if len(group_ids) != len(key_strs):
raise ValueError(
"Mooncake group_ids length must match key_strs length: "
f"{len(group_ids)} != {len(key_strs)}"
)
config = self._replicate_config_cls()
config.group_ids = group_ids
if self._uses_multi_buffer(buffer_ptrs):
config = config or self._replicate_config_cls()
return self.store.batch_put_from_multi_buffers(
key_strs, buffer_ptrs, buffer_sizes
key_strs, buffer_ptrs, buffer_sizes, config
)
return self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes)
elif config is not None:
return self.store.batch_put_from(
key_strs, buffer_ptrs, buffer_sizes, config
)
else:
return self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes)
def _get_batch_zero_copy_impl(
self, key_strs: List[str], buffer_ptrs: List[Any], buffer_sizes: List[Any]
@@ -0,0 +1,454 @@
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=1, 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_memory_pool_host_module():
memory_pool_host = types.ModuleType("sglang.srt.mem_cache.memory_pool_host")
class MLATokenToKVPoolHost:
pass
memory_pool_host.MLATokenToKVPoolHost = MLATokenToKVPoolHost
return memory_pool_host
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.memory_pool_host": _fake_memory_pool_host_module(),
"sglang.srt.mem_cache.pool_host": _fake_pool_host_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,
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="test",
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,
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,
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"],
)
if __name__ == "__main__":
unittest.main(verbosity=3)