diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py index d25b7811a..e0d5726a7 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py @@ -296,6 +296,52 @@ class MooncakeBaseStore: class MooncakeStore(HiCacheStorage, MooncakeBaseStore): + @staticmethod + def _standalone_required_bytes(mem_pool: Any) -> int: + """Compute total bytes of host buffers that must be visible to the real client. + + In standalone (dummy client) mode, the real mooncake_client process needs + to map any host buffers we will later pass by pointer via register_buffer(). + For hybrid models, that includes KV + sidecar pools (e.g. Mamba temporal/conv). + """ + # Prefer a generic "hybrid pool" accessor when present. + total = 0 + seen_ptrs: set[int] = set() + + def _add_tensor(t: Optional[torch.Tensor]): + nonlocal total + if t is None: + return + try: + ptr = int(t.data_ptr()) + except Exception: + return + if ptr in seen_ptrs: + return + seen_ptrs.add(ptr) + total += int(t.numel() * t.element_size()) + + # Always include the anchor KV buffer if present. + _add_tensor(getattr(mem_pool, "kv_buffer", None)) + + # HostPoolGroup: include each pool's hybrid buffers when available. + entries = getattr(mem_pool, "entries", None) + if entries: + for entry in entries: + host_pool = getattr(entry, "host_pool", None) + if host_pool is None: + continue + # KV pool anchor memory is already covered, but harmless if added twice. + _add_tensor(getattr(host_pool, "kv_buffer", None)) + for buf in getattr(host_pool, "get_hybrid_pool_buffer", lambda: [])(): + _add_tensor(buf) + return total + + # Single HostKVCache-like pool: add its sidecar buffers if any. + for buf in getattr(mem_pool, "get_hybrid_pool_buffer", lambda: [])(): + _add_tensor(buf) + return total + def __init__( self, storage_config: HiCacheStorageConfig = None, mem_pool: HostKVCache = None ): @@ -351,8 +397,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): "Please set standalone_storage=False " "or upgrade Mooncake by 'pip install mooncake --upgrade'." ) + required_bytes = self._standalone_required_bytes(mem_pool) ret_code = self.store.setup_dummy( - mem_pool.size * mem_pool.size_per_token, + required_bytes, DEFAULT_LOCAL_BUFFER_SIZE, # Zero copy interface does not need local buffer self.config.client_server_address, ) diff --git a/test/registered/unit/mem_cache/test_mooncake_standalone_dummy_mamba.py b/test/registered/unit/mem_cache/test_mooncake_standalone_dummy_mamba.py new file mode 100644 index 000000000..2ce3e58e3 --- /dev/null +++ b/test/registered/unit/mem_cache/test_mooncake_standalone_dummy_mamba.py @@ -0,0 +1,147 @@ +import types +import unittest +from unittest.mock import patch + +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") + + +def _fake_mooncake_modules(fake_store_cls): + mooncake = types.ModuleType("mooncake") + mooncake_store = types.ModuleType("mooncake.store") + mooncake_store.MooncakeDistributedStore = fake_store_cls + return { + "mooncake": mooncake, + "mooncake.store": mooncake_store, + } + + +class TestMooncakeStandaloneDummyMamba(CustomTestCase): + def test_setup_dummy_includes_hybrid_buffers(self): + """Standalone(dummy) must size shared mapping for KV + Mamba buffers.""" + import torch + + captured = {} + + class FakeMooncakeDistributedStore: + def setup_dummy(self, required_bytes, local_buffer_bytes, addr): + captured["required_bytes"] = int(required_bytes) + captured["local_buffer_bytes"] = int(local_buffer_bytes) + captured["addr"] = addr + return 0 + + def setup(self, *args, **kwargs): + raise AssertionError("should not call setup() in standalone mode") + + def register_buffer(self, ptr, size): + return 0 + + def put(self, *args, **kwargs): + return 0 + + def is_exist(self, *args, **kwargs): + return 1 + + def get(self, *args, **kwargs): + return bytes(4 * 1024) + + with patch.dict( + "sys.modules", + _fake_mooncake_modules(FakeMooncakeDistributedStore), + ): + from sglang.srt.mem_cache.hicache_storage import ( + HiCacheStorageConfig, + PoolName, + ) + from sglang.srt.mem_cache.storage.mooncake_store import ( + mooncake_store as mc_mod, + ) + from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import ( + MooncakeStore, + ) + + class FakeAllocator: + pass + + class FakeKVPool: + def __init__(self): + # KV buffer (anchor). + self.kv_buffer = torch.empty((128,), dtype=torch.uint8) + self.size = 128 + self.size_per_token = 1 + self.allocator = FakeAllocator() + + class FakeMambaPool: + def __init__(self): + self.temporal_buffer = torch.empty((64,), dtype=torch.uint8) + self.conv_buffer = [torch.empty((32,), dtype=torch.uint8)] + + def get_hybrid_pool_buffer(self): + return [self.temporal_buffer, *self.conv_buffer] + + class FakeEntry: + def __init__(self, name, host_pool): + self.name = name + self.host_pool = host_pool + + class FakeHostPoolGroup: + def __init__(self): + self.kv = FakeKVPool() + self.mamba = FakeMambaPool() + self.entries = [ + FakeEntry(PoolName.KV, self.kv), + FakeEntry(PoolName.MAMBA, self.mamba), + ] + + # Anchor-like fields accessed by MooncakeStore. + @property + def kv_buffer(self): + return self.kv.kv_buffer + + @property + def allocator(self): + return self.kv.allocator + + @property + def size(self): + return self.kv.size + + @property + def size_per_token(self): + return self.kv.size_per_token + + mem_pool = FakeHostPoolGroup() + cfg = HiCacheStorageConfig( + tp_rank=0, + tp_size=1, + pp_rank=0, + pp_size=1, + attn_cp_rank=0, + attn_cp_size=1, + is_mla_model=False, + enable_storage_metrics=False, + is_page_first_layout=True, + model_name="test", + extra_config={ + "standalone_storage": True, + "client_server_address": "127.0.0.1:50052", + }, + ) + + with patch.object(mc_mod, "MooncakeHostTensorAllocator", FakeAllocator): + MooncakeStore(cfg, mem_pool) + + expected = ( + mem_pool.kv.kv_buffer.numel() * mem_pool.kv.kv_buffer.element_size() + + mem_pool.mamba.temporal_buffer.numel() + * mem_pool.mamba.temporal_buffer.element_size() + + mem_pool.mamba.conv_buffer[0].numel() + * mem_pool.mamba.conv_buffer[0].element_size() + ) + self.assertEqual(captured["required_bytes"], expected) + + +if __name__ == "__main__": + unittest.main(verbosity=3)