[HiCache] fix: Mooncake Dummy Client mode for hybrid Mamba models (#25278)
Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: Teng Ma <stmatengss@users.noreply.github.com>
This commit is contained in:
co-authored by
Cursor Agent
Teng Ma
parent
1051a8456f
commit
d45ee3f6c5
@@ -296,6 +296,52 @@ class MooncakeBaseStore:
|
|||||||
|
|
||||||
class MooncakeStore(HiCacheStorage, 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__(
|
def __init__(
|
||||||
self, storage_config: HiCacheStorageConfig = None, mem_pool: HostKVCache = None
|
self, storage_config: HiCacheStorageConfig = None, mem_pool: HostKVCache = None
|
||||||
):
|
):
|
||||||
@@ -351,8 +397,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
"Please set standalone_storage=False "
|
"Please set standalone_storage=False "
|
||||||
"or upgrade Mooncake by 'pip install mooncake --upgrade'."
|
"or upgrade Mooncake by 'pip install mooncake --upgrade'."
|
||||||
)
|
)
|
||||||
|
required_bytes = self._standalone_required_bytes(mem_pool)
|
||||||
ret_code = self.store.setup_dummy(
|
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
|
DEFAULT_LOCAL_BUFFER_SIZE, # Zero copy interface does not need local buffer
|
||||||
self.config.client_server_address,
|
self.config.client_server_address,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user