[CCF Archive] vLLM EPD hidden connector submission #4

Closed
kancel wants to merge 1 commits from kancel:ccf-archive-vllm-epd into main
14 changed files with 8375 additions and 0 deletions

View File

@ -0,0 +1,28 @@
# CCF vLLM EPD 代码归档
本目录用于 CCF Mooncake 赛题提交归档。由于本赛题 GitLink 侧只提供
Mooncake 仓库,而本部分工作实现于 vLLM因此将相关 vLLM 改动源码快照
和 patch 一并放入 Mooncake 仓库根目录,便于比赛评审在 GitLink 上查看。
对应的 vLLM 实现分支为:
- 仓库https://github.com/kanceler/vllm
- 分支:`mooncake-store-ec-hidden`
- 归档基准提交:`b4482f0a10a25ddbf34b687129c344594fab4613`
## 目录内容
- `code/`vLLM EPD Hidden State connector 相关源码快照,包括 connector
实现、factory 注册入口和单元测试。
- `patches/vllm-epd-hidden-ec-connector-b4482f0a1-full-feature.patch`:归档
vLLM 实现的完整功能 patch。
- `patches/vllm-epd-hidden-ec-connector-b4482f0a1-latest-refinement.patch`
归档实现过程中的 refinement patch。
- `patches/vllm-epd-hidden-ec-connector-review-fixes.patch`:初赛提交前代码
审查阶段补充的增量修正,主要包括异步保存失败可观测性和 connector
资源释放逻辑。
## 归档范围
本目录不作为 Mooncake 主工程的可合入代码,仅用于满足比赛要求中“上游软件
托管至 GitLink”的提交要求。实际 vLLM 实现请以上述 vLLM 仓库和分支为准。

View File

@ -0,0 +1,419 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from types import SimpleNamespace
import torch
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden import (
connector as connector_module,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.connector import (
MooncakeStoreECConnector,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HIDDEN_TENSOR_LAYOUT,
HiddenKeyMetadata,
HiddenPoolKey,
LoadSpec,
MMMeta,
MooncakeStoreConnectorMetadata,
)
from vllm.multimodal.inputs import MultiModalFeatureSpec, PlaceholderRange
class FakeWorker:
def __init__(self):
self.requests = []
self.key_metadata = HiddenKeyMetadata(
cache_prefix="",
kind="encoder_output",
model_name="qwen",
encoder="encoder-config-a",
storage="replicated_object",
parallel="tp:1@pp:1@pcp:1@dcp:1@mm_tp:weights",
tensor_layout=HIDDEN_TENSOR_LAYOUT,
)
def make_pool_key(self, identifier: str) -> HiddenPoolKey:
return HiddenPoolKey(self.key_metadata, identifier)
def enqueue_save(self, request):
self.requests.append(request)
def get_finished_sending(self):
return set()
def get_failed_sending(self):
return {}
def make_connector(*, soft_pin_video_hidden: bool = False):
connector = MooncakeStoreECConnector.__new__(MooncakeStoreECConnector)
connector._is_producer = True
connector._is_consumer = False
connector.lookup_client = None
connector.lookup_async = True
connector.worker = FakeWorker()
connector._connector_metadata = None
connector.soft_pin_video_hidden = soft_pin_video_hidden
connector.load_specs = {}
connector.lookup_result_cache = {}
connector.identifier_waiters = {}
connector._candidate_consumes = {}
connector._candidate_loads = {}
connector._candidate_saves = {}
connector._load_modalities = {}
connector._save_modalities = {}
return connector
class FakeLookupClient:
def __init__(self, results):
self.results = list(results)
self.calls = []
self.discarded = []
def lookup_batch(self, identifiers, non_block=True):
self.calls.append((tuple(identifiers), non_block))
return self.results.pop(0)
def discard(self, identifier):
self.discarded.append(identifier)
def make_request(request_id, features):
mm_features = [
MultiModalFeatureSpec(
data=None,
modality=modality,
identifier=identifier,
mm_position=PlaceholderRange(offset=offset, length=length),
)
for identifier, offset, length, modality in features
]
return SimpleNamespace(
request_id=request_id,
mm_features=mm_features,
num_tokens=1000,
)
def make_scheduler_output(*, finished_req_ids=None, preempted_req_ids=None):
return SimpleNamespace(
finished_req_ids=finished_req_ids or set(),
preempted_req_ids=preempted_req_ids,
)
def test_build_hidden_key_metadata_uses_structured_key_fields(monkeypatch):
monkeypatch.setattr(
connector_module,
"get_tensor_model_parallel_world_size",
lambda: 4,
)
monkeypatch.setattr(
connector_module,
"get_pcp_group",
lambda: SimpleNamespace(world_size=1),
)
monkeypatch.setattr(
connector_module,
"get_dcp_group",
lambda: SimpleNamespace(world_size=1),
)
multimodal_config = SimpleNamespace(
compute_hash=lambda: "encoder-config-a",
mm_encoder_tp_mode="data",
)
vllm_config = SimpleNamespace(
model_config=SimpleNamespace(
model="/models/qwen",
multimodal_config=multimodal_config,
),
parallel_config=SimpleNamespace(pipeline_parallel_size=2),
ec_transfer_config=SimpleNamespace(
ec_connector_extra_config={
"cache_prefix": "shared-prefix",
"hidden_cache_prefix": "hidden-prefix",
}
),
)
metadata = connector_module.build_hidden_key_metadata(vllm_config)
assert metadata.cache_prefix == "hidden-prefix"
assert metadata.kind == "encoder_output"
assert metadata.model_name == "qwen"
assert metadata.encoder == "encoder-config-a"
assert metadata.storage == "replicated_object"
assert metadata.parallel == "tp:4@pp:2@pcp:1@dcp:1@mm_tp:data"
assert "storage" not in metadata.parallel
assert metadata.tensor_layout == "tensor"
def test_ensure_cache_available_defers_pending_batch_lookup():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_client = FakeLookupClient([None])
request = make_request(
"req-1",
[
("image-1", 20, 60, "image"),
("image-2", 500, 60, "image"),
],
)
assert not connector.ensure_cache_available(request, num_computed_tokens=0)
assert connector.lookup_client.calls == [
(("image-1", "image-2"), True),
]
assert connector.identifier_waiters == {
"image-1": {"req-1"},
"image-2": {"req-1"},
}
def test_ensure_cache_available_deduplicates_request_waiters_and_lookup_results():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_client = FakeLookupClient(
[
{"image-1": True, "image-2": False},
]
)
request = make_request(
"req-1",
[
("image-1", 20, 60, "image"),
("image-2", 500, 60, "image"),
],
)
assert connector.ensure_cache_available(request, num_computed_tokens=0)
assert connector.ensure_cache_available(request, num_computed_tokens=0)
assert connector.lookup_client.calls == [
(("image-1", "image-2"), True),
]
assert connector.identifier_waiters == {
"image-1": {"req-1"},
"image-2": {"req-1"},
}
assert connector.lookup_result_cache == {
"image-1": True,
"image-2": False,
}
def test_has_cache_item_is_local_only():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_client = SimpleNamespace(lookup=lambda identifier: True)
connector.lookup_result_cache = {"image-1": True, "image-2": False}
assert connector.has_cache_item("image-1")
assert not connector.has_cache_item("image-2")
assert not connector.has_cache_item("unknown")
def test_build_connector_meta_commits_waiter_consumes_and_keeps_unreached_image():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_result_cache = {"image-1": True, "image-2": True}
connector.identifier_waiters = {
"image-1": {"req-1"},
"image-2": {"req-1"},
}
connector.load_specs["image-1"] = LoadSpec(can_load=False)
request = make_request(
"req-1",
[
("image-1", 20, 60, "image"),
("image-2", 500, 60, "image"),
],
)
connector.update_state_after_alloc(request, 0)
meta = connector.build_connector_meta(make_scheduler_output())
assert [item.identifier for item in meta.items] == ["image-1"]
assert "image-1" not in connector.identifier_waiters
assert "image-1" not in connector.lookup_result_cache
assert connector.identifier_waiters == {"image-2": {"req-1"}}
assert connector.lookup_result_cache == {"image-2": True}
def test_build_connector_meta_rolls_back_preempted_candidate_state():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_result_cache = {"image-1": True}
connector.identifier_waiters = {"image-1": {"req-1"}}
connector.load_specs["image-1"] = LoadSpec(can_load=False)
request = make_request("req-1", [("image-1", 20, 60, "image")])
connector.update_state_after_alloc(request, 0)
meta = connector.build_connector_meta(
make_scheduler_output(preempted_req_ids={"req-1"})
)
assert meta.items == []
assert connector.identifier_waiters == {"image-1": {"req-1"}}
assert connector.lookup_result_cache == {"image-1": True}
assert "image-1" in connector.load_specs
def test_build_connector_meta_cleans_finished_waiters():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_client = FakeLookupClient([])
connector.lookup_result_cache = {"image-1": True, "image-2": True}
connector.identifier_waiters = {
"image-1": {"req-1"},
"image-2": {"req-1", "req-2"},
}
connector.build_connector_meta(make_scheduler_output(finished_req_ids={"req-1"}))
assert "image-1" not in connector.identifier_waiters
assert "image-1" not in connector.lookup_result_cache
assert connector.identifier_waiters == {"image-2": {"req-2"}}
assert connector.lookup_result_cache == {"image-2": True}
assert connector.lookup_client.discarded == ["image-1"]
def test_cleanup_lookup_results_discards_inflight_lookup_without_waiters():
connector = make_connector()
connector._is_consumer = True
connector._is_producer = False
connector.lookup_client = FakeLookupClient([])
connector.identifier_waiters = {"image-1": set()}
connector.lookup_result_cache = {"image-1": True}
connector.load_specs["image-1"] = LoadSpec(can_load=False)
connector._cleanup_lookup_results_without_waiters()
assert connector.identifier_waiters == {}
assert connector.lookup_result_cache == {}
assert connector.load_specs == {}
assert connector.lookup_client.discarded == ["image-1"]
def test_build_connector_meta_merges_load_and_save_item_by_identifier():
connector = make_connector()
connector._is_consumer = True
connector.load_specs["video-hash"] = LoadSpec(can_load=False)
connector.lookup_result_cache["video-hash"] = True
connector.identifier_waiters["video-hash"] = {"req-1"}
request = SimpleNamespace(
request_id="req-1",
mm_features=[
SimpleNamespace(
identifier="video-hash",
modality="video",
)
],
)
connector.update_state_after_alloc(request, 0)
meta = connector.build_connector_meta(make_scheduler_output())
assert len(meta.items) == 1
item = meta.items[0]
assert item.identifier == "video-hash"
assert item.modality == "video"
assert item.can_save
assert item.load_spec is not None
assert item.load_spec.can_load
assert connector.load_specs == {}
def test_save_caches_skips_items_without_save_plan():
connector = make_connector()
connector.bind_connector_metadata(
MooncakeStoreConnectorMetadata(
items=[
MMMeta(
identifier="image-hash",
modality="image",
can_save=False,
)
]
)
)
connector.save_caches({"image-hash": torch.zeros((1, 2))}, "image-hash")
assert connector.worker.requests == []
def test_save_caches_enqueues_video_hidden_with_soft_pin():
connector = make_connector(soft_pin_video_hidden=True)
tensor = torch.zeros((1, 2))
connector.bind_connector_metadata(
MooncakeStoreConnectorMetadata(
items=[
MMMeta(
identifier="video-hash",
modality="video",
can_save=True,
load_spec=LoadSpec(can_load=False),
)
]
)
)
connector.save_caches({"video-hash": tensor}, "video-hash")
assert len(connector.worker.requests) == 1
request = connector.worker.requests[0]
assert request.identifier == "video-hash"
assert request.tensor is tensor
assert request.with_soft_pin
def test_save_caches_does_not_soft_pin_image_hidden():
connector = make_connector(soft_pin_video_hidden=True)
connector.bind_connector_metadata(
MooncakeStoreConnectorMetadata(
items=[
MMMeta(
identifier="image-hash",
modality="image",
can_save=True,
)
]
)
)
connector.save_caches({"image-hash": torch.zeros((1, 2))}, "image-hash")
assert len(connector.worker.requests) == 1
assert not connector.worker.requests[0].with_soft_pin
def test_get_finished_logs_failed_hidden_saves(caplog):
class FailedWorker(FakeWorker):
def get_finished_sending(self):
return {"image-ok"}
def get_failed_sending(self):
return {"image-failed": "batch put failed"}
connector = make_connector()
connector.worker = FailedWorker()
finished_sending, finished_recving = connector.get_finished({"req-1"})
assert finished_sending == {"image-ok"}
assert finished_recving is None
assert "hidden_store_save_failed" in caplog.text
assert "image-failed" in caplog.text
assert "batch put failed" in caplog.text

View File

@ -0,0 +1,143 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HIDDEN_TENSOR_LAYOUT,
HiddenKeyMetadata,
HiddenPoolKey,
LoadSpec,
MMMeta,
MooncakeStoreConnectorMetadata,
build_tensor_meta,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.keys import (
make_hidden_data_key,
)
def make_pool_key(
identifier: str = "image-hash",
*,
cache_prefix: str = "",
kind: str = "encoder_output",
model_name: str = "qwen",
encoder: str = "encoder-config-a",
storage: str = "replicated_object",
parallel: str = "tp:1@pp:1@pcp:1@dcp:1@mm_tp:weights",
tensor_layout: str = HIDDEN_TENSOR_LAYOUT,
) -> HiddenPoolKey:
return HiddenPoolKey(
key_metadata=HiddenKeyMetadata(
cache_prefix=cache_prefix,
kind=kind,
model_name=model_name,
encoder=encoder,
storage=storage,
parallel=parallel,
tensor_layout=tensor_layout,
),
identifier=identifier,
)
def test_hidden_pool_key_is_the_single_tensor_object_key():
pool_key = make_pool_key()
data_key = make_hidden_data_key(pool_key)
assert data_key == pool_key.to_string()
assert data_key.startswith("hidden@")
assert "kind:encoder_output" in data_key
assert "model:qwen" in data_key
assert "encoder:encoder-config-a" in data_key
assert "storage:replicated_object" in data_key
assert (
"parallel:tp%3A1%40pp%3A1%40pcp%3A1%40dcp%3A1%40mm_tp%3Aweights"
in data_key
)
assert "tensor_layout:tensor" in data_key
assert "storage%3Areplicated" not in data_key
assert "writer" not in data_key
assert "adapter:" not in data_key
assert "modality:" not in data_key
assert "image-hash" in data_key
def test_same_identifier_with_different_encoder_config_uses_different_keys():
pool_key_a = make_pool_key(encoder="encoder-config-a")
pool_key_b = make_pool_key(encoder="encoder-config-b")
assert make_hidden_data_key(pool_key_a) != make_hidden_data_key(pool_key_b)
def test_cache_prefix_namespaces_hidden_pool_key():
pool_key_a = make_pool_key(cache_prefix="deployment-a")
pool_key_b = make_pool_key(cache_prefix="deployment-b")
data_key_a = make_hidden_data_key(pool_key_a)
data_key_b = make_hidden_data_key(pool_key_b)
assert data_key_a.startswith("deployment-a@hidden@")
assert data_key_b.startswith("deployment-b@hidden@")
assert data_key_a != data_key_b
def test_request_id_and_modality_are_not_part_of_hidden_pool_key():
pool_key = make_pool_key(identifier="image-hash")
assert "req-1" not in make_hidden_data_key(pool_key)
assert "request" not in make_hidden_data_key(pool_key)
assert "image@" not in make_hidden_data_key(pool_key)
assert "modality" not in make_hidden_data_key(pool_key)
def test_mm_meta_carries_hidden_item_plan():
item = MMMeta(
identifier="image-hash",
modality="video",
can_save=True,
load_spec=LoadSpec(can_load=True),
)
meta = MooncakeStoreConnectorMetadata(items=[item])
assert meta.items == [item]
assert meta.items[0].identifier == "image-hash"
assert meta.items[0].modality == "video"
assert meta.items[0].can_save
assert meta.items[0].load_spec is not None
assert meta.items[0].load_spec.can_load
def test_tensor_meta_describes_canonical_contiguous_tensor():
pool_key = make_pool_key()
source = torch.zeros((4, 8), dtype=torch.float16).t()
stored = source.contiguous()
tensor_meta = build_tensor_meta(pool_key, stored)
assert tensor_meta.pool_key == pool_key
assert tensor_meta.layout == HIDDEN_TENSOR_LAYOUT
assert tensor_meta.shape == tuple(stored.shape)
assert tensor_meta.dtype == "torch.float16"
assert tensor_meta.nbytes == stored.numel() * stored.element_size()
def test_tensor_meta_rejects_non_contiguous_tensor():
pool_key = make_pool_key()
source = torch.zeros((4, 8), dtype=torch.float16).t()
try:
build_tensor_meta(pool_key, source)
except ValueError as exc:
assert "contiguous" in str(exc)
else:
raise AssertionError("non-contiguous tensor descriptor should fail")
def test_pool_key_namespace_carries_reuse_compatibility():
pool_key_a = make_pool_key(encoder="encoder-config-a")
pool_key_b = make_pool_key(encoder="encoder-config-b")
assert pool_key_a != pool_key_b
assert make_hidden_data_key(pool_key_a) != make_hidden_data_key(pool_key_b)

View File

@ -0,0 +1,850 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import ctypes
import sys
import struct
import types
from concurrent.futures import Future
import torch
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HIDDEN_TENSOR_LAYOUT,
HiddenKeyMetadata,
HiddenPoolKey,
HiddenSaveRequest,
HiddenTensorDatabase,
LoadSpec,
MMMeta,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.keys import (
make_hidden_data_key,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.store_client import (
HiddenStoreError,
HiddenStoreLoadError,
HiddenStoreSaveError,
MooncakeHiddenStoreClient,
_get_hidden_state_object_data_type,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.worker import (
HiddenLookupClient,
HiddenLookupServer,
HiddenStoreSendingThread,
HiddenStoreWorker,
)
TENSOR_METADATA_SIZE = 304
TENSOR_OBJECT_MAGIC = 0x4D4F4F4E
TENSOR_OBJECT_VERSION = 1
TORCH_DTYPE_TO_MOONCAKE_DTYPE = {
torch.float32: 0,
torch.float16: 11,
torch.bfloat16: 12,
}
class FakeStore:
def __init__(self):
self.objects = {}
self.batch_is_exist_calls = []
self.registered = []
self.unregistered = []
self.pub_tensors = []
self.range_gets = []
self.fail_register_addrs = set()
self.raise_on_batch_put = False
self.batch_put_results = [0]
def batch_is_exist(self, keys):
self.batch_is_exist_calls.append(list(keys))
return [1 if key in self.objects else 0 for key in keys]
def register_buffer(self, addr, size):
if addr in self.fail_register_addrs:
return -1
self.registered.append((addr, size))
return 0
def unregister_buffer(self, addr):
self.unregistered.append(addr)
return 0
def pub_tensor(self, key, tensor, replicate_config=None):
self.pub_tensors.append((key, tensor, replicate_config))
self.objects[key] = _serialize_tensor_object(tensor)
return 0
def put_tensor(self, key, tensor):
return self.pub_tensor(key, tensor)
def get_into_ranges(
self,
buffer_ptrs,
all_keys,
all_dst_offsets,
all_src_offsets,
all_sizes,
):
self.range_gets.append(
(buffer_ptrs, all_keys, all_dst_offsets, all_src_offsets, all_sizes)
)
results = []
for buffer_ptr, keys, dst_offsets, src_offsets, sizes in zip(
buffer_ptrs,
all_keys,
all_dst_offsets,
all_src_offsets,
all_sizes,
strict=True,
):
key_results = []
for key, key_dst_offsets, key_src_offsets, key_sizes in zip(
keys,
dst_offsets,
src_offsets,
sizes,
strict=True,
):
payload = self.objects.get(key)
fragment_results = []
for dst_offset, src_offset, size in zip(
key_dst_offsets,
key_src_offsets,
key_sizes,
strict=True,
):
if payload is None or src_offset + size > len(payload):
fragment_results.append(-1)
continue
ctypes.memmove(
buffer_ptr + dst_offset,
payload[src_offset : src_offset + size],
size,
)
fragment_results.append(size)
key_results.append(fragment_results)
results.append(key_results)
return results
class FakeClosableStore(FakeStore):
def __init__(self):
super().__init__()
self.closed = False
def close(self):
self.closed = True
class FakeStoreWithTeardown(FakeStore):
def __init__(self):
super().__init__()
self.teardown_called = False
def teardown(self):
self.teardown_called = True
class FakeSocket:
def __init__(self):
self.closed = False
self.linger = None
def close(self, linger=0):
self.closed = True
self.linger = linger
class FakeContext:
def __init__(self):
self.destroy_called = False
self.term_called = False
def destroy(self, linger=0):
self.destroy_called = True
def term(self):
self.term_called = True
class FakeThread:
def __init__(self):
self.join_called = False
self.timeout = None
def join(self, timeout=None):
self.join_called = True
self.timeout = timeout
def is_alive(self):
return False
class FakeBufferStore(FakeStore):
def batch_put_from_multi_buffers(
self,
keys,
buffer_ptrs,
buffer_sizes,
replicate_config=None,
):
if self.raise_on_batch_put:
raise RuntimeError("batch put failed")
self.objects[keys[0]] = b"tensor-object"
return self.batch_put_results
class FakeReplicateConfig:
def __init__(self):
self.replica_num = 1
self.nof_replica_num = 0
self.with_soft_pin = False
self.with_hard_pin = False
self.preferred_segments = []
self.preferred_nof_segments = []
self.preferred_segment = ""
self.prefer_alloc_in_same_node = False
self.data_type = None
self.group_ids = None
class FakeReplicateConfigWithoutGroups:
def __init__(self):
self.replica_num = 1
class FakeObjectDataTypeWithHidden:
HIDDEN_STATE = 10
TENSOR = 2
class FakeObjectDataTypeOnlyTensor:
TENSOR = 2
class FakeObjectDataTypeNoTensor:
UNKNOWN = 0
def make_pool_key(identifier: str = "image-hash") -> HiddenPoolKey:
return HiddenPoolKey(
key_metadata=HiddenKeyMetadata(
cache_prefix="",
kind="encoder_output",
model_name="qwen",
encoder="encoder-config-a",
storage="replicated_object",
parallel="tp:1@pp:1@pcp:1@dcp:1@mm_tp:weights",
tensor_layout=HIDDEN_TENSOR_LAYOUT,
),
identifier=identifier,
)
def test_hidden_tensor_database_prepares_data_key_addrs_and_sizes():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
key, addrs, sizes = HiddenTensorDatabase().prepare_value(pool_key, tensor)
assert key == make_hidden_data_key(pool_key)
assert addrs == [tensor.data_ptr()]
assert sizes == [tensor.numel() * tensor.element_size()]
def test_store_client_checks_single_tensor_object_exists():
pool_key = make_pool_key()
store = FakeStore()
client = MooncakeHiddenStoreClient(store)
assert not client.exists(pool_key)
store.objects[make_hidden_data_key(pool_key)] = b"tensor-object"
assert client.exists(pool_key)
def test_worker_lookup_checks_existence_without_reading_tensor_metadata():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key.key_metadata,
)
worker.save_tensor(pool_key, tensor)
assert worker.lookup(pool_key.identifier)
assert not worker.lookup("missing-image-hash")
assert store.range_gets == []
def test_worker_batch_lookup_checks_existence_in_one_store_call():
pool_key_a = make_pool_key("image-a")
pool_key_b = make_pool_key("image-b")
store = FakeBufferStore()
store.objects[make_hidden_data_key(pool_key_a)] = b"tensor-object"
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key_a.key_metadata,
)
results = worker.lookup_batch(["image-a", "image-b"])
assert results == {"image-a": True, "image-b": False}
assert store.batch_is_exist_calls == [
[make_hidden_data_key(pool_key_a), make_hidden_data_key(pool_key_b)]
]
assert store.range_gets == []
def test_lookup_client_discard_removes_identifier_future_mapping():
client = HiddenLookupClient.__new__(HiddenLookupClient)
future: Future[dict[str, bool]] = Future()
client.futures = {
"image-a": future,
"image-b": future,
}
client.discard("image-a")
assert "image-a" not in client.futures
assert client.futures == {"image-b": future}
assert not future.cancelled()
client.discard("image-b")
assert client.futures == {}
assert future.cancelled()
def test_worker_lookup_records_minimal_operation_stats():
pool_key = make_pool_key()
store = FakeBufferStore()
store.objects[make_hidden_data_key(pool_key)] = b"tensor-object"
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key.key_metadata,
)
assert worker.lookup_batch(["image-hash", "missing-image-hash"]) == {
"image-hash": True,
"missing-image-hash": False,
}
stats = worker.get_operation_stats()
records = stats.data["lookup_exists"]
assert len(records) == 1
assert records[0]["num_keys"] == 2
assert records[0]["num_bytes"] == 0
assert records[0]["status"] == "miss"
assert records[0]["num_failed_keys"] == 1
assert worker.get_operation_stats() is None
def test_worker_save_stores_hidden_as_single_tensor_object():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(
store,
replicate_config=FakeReplicateConfig(),
),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
assert store.pub_tensors[0][0] == make_hidden_data_key(pool_key)
assert store.pub_tensors[0][2] is not None
assert make_hidden_data_key(pool_key) in store.objects
def test_worker_save_rejects_dtype_that_load_cannot_decode():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float64)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
try:
worker.save_tensor(pool_key, tensor)
except HiddenStoreSaveError as exc:
assert "unsupported hidden tensor dtype" in str(exc)
else:
raise AssertionError("unsupported hidden dtype should fail before store put")
assert store.pub_tensors == []
def test_buffer_put_unregisters_payload_and_metadata_buffers():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeBufferStore()
client = MooncakeHiddenStoreClient(store, replicate_config=FakeReplicateConfig())
client.put_tensor(pool_key, tensor)
payload_addr = tensor.data_ptr()
metadata_addr = next(
addr for addr, size in store.registered if size == TENSOR_METADATA_SIZE
)
assert payload_addr in store.unregistered
assert metadata_addr in store.unregistered
assert store.unregistered[-2:] == [metadata_addr, payload_addr]
def test_buffer_put_unregisters_payload_and_metadata_when_put_raises():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeBufferStore()
store.raise_on_batch_put = True
client = MooncakeHiddenStoreClient(store, replicate_config=FakeReplicateConfig())
try:
client.put_tensor(pool_key, tensor)
except RuntimeError as exc:
assert "batch put failed" in str(exc)
else:
raise AssertionError("batch put exception should propagate")
payload_addr = tensor.data_ptr()
metadata_addr = next(
addr for addr, size in store.registered if size == TENSOR_METADATA_SIZE
)
assert payload_addr in store.unregistered
assert metadata_addr in store.unregistered
def test_buffer_put_unregisters_payload_when_metadata_registration_fails():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeBufferStore()
original_register = store.register_buffer
def register_buffer(addr, size):
if size == TENSOR_METADATA_SIZE:
store.fail_register_addrs.add(addr)
return original_register(addr, size)
store.register_buffer = register_buffer
client = MooncakeHiddenStoreClient(store, replicate_config=FakeReplicateConfig())
try:
client.put_tensor(pool_key, tensor)
except HiddenStoreError:
pass
else:
raise AssertionError("metadata registration failure should raise")
assert tensor.data_ptr() in store.unregistered
def test_worker_save_marks_hidden_state_data_type(monkeypatch):
fake_mooncake = types.ModuleType("mooncake")
fake_store = types.ModuleType("mooncake.store")
fake_store.ObjectDataType = FakeObjectDataTypeWithHidden
monkeypatch.setitem(sys.modules, "mooncake", fake_mooncake)
monkeypatch.setitem(sys.modules, "mooncake.store", fake_store)
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
replicate_config = FakeReplicateConfig()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(
store,
replicate_config=replicate_config,
),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
used_config = store.pub_tensors[0][2]
assert used_config is not replicate_config
assert int(used_config.data_type) == 10
def test_hidden_state_data_type_falls_back_to_tensor(monkeypatch):
fake_mooncake = types.ModuleType("mooncake")
fake_store = types.ModuleType("mooncake.store")
fake_store.ObjectDataType = FakeObjectDataTypeOnlyTensor
monkeypatch.setitem(sys.modules, "mooncake", fake_mooncake)
monkeypatch.setitem(sys.modules, "mooncake.store", fake_store)
assert _get_hidden_state_object_data_type() == FakeObjectDataTypeOnlyTensor.TENSOR
def test_hidden_state_data_type_missing_type_returns_none(monkeypatch):
fake_mooncake = types.ModuleType("mooncake")
fake_store = types.ModuleType("mooncake.store")
fake_store.ObjectDataType = FakeObjectDataTypeNoTensor
monkeypatch.setitem(sys.modules, "mooncake", fake_mooncake)
monkeypatch.setitem(sys.modules, "mooncake.store", fake_store)
assert _get_hidden_state_object_data_type() is None
def test_worker_save_does_not_require_mooncake_object_group_support():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(
store,
replicate_config=FakeReplicateConfigWithoutGroups(),
),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
assert store.pub_tensors[0][0] == make_hidden_data_key(pool_key)
def test_worker_save_skips_existing_tensor_object():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
worker.save_tensor(pool_key, tensor)
assert len(store.pub_tensors) == 1
def test_worker_save_records_exists_and_put_operation_stats():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
stats = worker.get_operation_stats()
assert stats.data["save_exists"][0]["status"] == "miss"
assert stats.data["save_exists"][0]["num_keys"] == 1
assert stats.data["save_put"][0]["status"] == "ok"
assert stats.data["save_put"][0]["num_keys"] == 1
assert stats.data["save_put"][0]["num_bytes"] == (
tensor.numel() * tensor.element_size()
)
def test_worker_save_existing_records_only_save_exists():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
store.objects[make_hidden_data_key(pool_key)] = b"tensor-object"
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
worker.save_tensor(pool_key, tensor)
stats = worker.get_operation_stats()
assert stats.data["save_exists"][0]["status"] == "ok"
assert "save_put" not in stats.data
def test_sending_thread_stores_hidden_tensor_asynchronously():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
sending_thread = HiddenStoreSendingThread(worker)
sending_thread.start()
sending_thread.add_request(
HiddenSaveRequest(pool_key=pool_key, tensor=tensor)
)
sending_thread.request_queue.join()
assert store.pub_tensors[0][0] == make_hidden_data_key(pool_key)
assert sending_thread.get_and_clear_finished_identifiers() == {pool_key.identifier}
sending_thread.close()
def test_sending_thread_records_failed_identifier_without_finishing():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeBufferStore()
store.raise_on_batch_put = True
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
sending_thread = HiddenStoreSendingThread(worker)
sending_thread.start()
sending_thread.add_request(
HiddenSaveRequest(pool_key=pool_key, tensor=tensor)
)
sending_thread.request_queue.join()
assert sending_thread.get_and_clear_finished_identifiers() == set()
assert sending_thread.get_and_clear_failed_identifiers() == {pool_key.identifier}
assert pool_key.identifier in sending_thread.failure_reasons
assert worker.get_operation_stats().data["save_put"][0]["status"] == "error"
sending_thread.close()
def test_worker_drains_failed_sending_reasons():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeBufferStore()
store.raise_on_batch_put = True
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
worker.start_sending_thread()
assert worker.sending_thread is not None
worker.enqueue_save(HiddenSaveRequest(pool_key=pool_key, tensor=tensor))
worker.sending_thread.request_queue.join()
failed = worker.get_failed_sending()
assert set(failed) == {pool_key.identifier}
assert "batch put failed" in failed[pool_key.identifier]
assert worker.get_failed_sending() == {}
worker.shutdown()
def test_sending_thread_close_joins_worker_thread():
pool_key = make_pool_key()
tensor = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
sending_thread = HiddenStoreSendingThread(worker)
sending_thread.start()
sending_thread.add_request(HiddenSaveRequest(pool_key=pool_key, tensor=tensor))
sending_thread.request_queue.join()
sending_thread.close()
assert not sending_thread.is_alive()
def test_worker_shutdown_closes_store_client():
store = FakeClosableStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
)
worker.shutdown()
assert store.closed
def test_store_client_close_uses_fallback_close_method():
store = FakeStoreWithTeardown()
client = MooncakeHiddenStoreClient(store)
client.close()
assert store.teardown_called
def test_lookup_server_close_joins_thread_and_closes_context(tmp_path):
socket = FakeSocket()
ctx = FakeContext()
thread = FakeThread()
ipc_path = tmp_path / "hidden_lookup.ipc"
ipc_path.write_text("socket")
server = HiddenLookupServer.__new__(HiddenLookupServer)
server.running = True
server.socket = socket
server.ctx = ctx
server.thread = thread
server._ipc_path = str(ipc_path)
server.close()
assert not server.running
assert socket.closed
assert socket.linger == 0
assert thread.join_called
assert ctx.destroy_called or ctx.term_called
assert not ipc_path.exists()
def test_lookup_client_close_shuts_down_executor_socket_and_context():
socket = FakeSocket()
ctx = FakeContext()
executor = types.SimpleNamespace(
shutdown_called=False,
shutdown=lambda wait=False, cancel_futures=True: setattr(
executor, "shutdown_called", True
),
)
client = HiddenLookupClient.__new__(HiddenLookupClient)
client.executor = executor
client.futures = {"image-hash": Future()}
client.socket = socket
client.ctx = ctx
client.close()
assert executor.shutdown_called
assert client.futures == {}
assert socket.closed
assert socket.linger == 0
assert ctx.destroy_called or ctx.term_called
def test_worker_load_gets_tensor_data_into_encoder_cache_before_returning():
pool_key = make_pool_key()
stored = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key.key_metadata,
)
worker.save_tensor(pool_key, stored)
encoder_cache = {}
worker.load(
[MMMeta(identifier=pool_key.identifier, load_spec=LoadSpec(can_load=True))],
encoder_cache,
device="cpu",
)
assert pool_key.identifier in encoder_cache
assert tuple(encoder_cache[pool_key.identifier].shape) == tuple(stored.shape)
assert str(encoder_cache[pool_key.identifier].dtype) == str(stored.dtype)
assert store.range_gets[0][1] == [[make_hidden_data_key(pool_key)]]
assert store.range_gets[0][3] == [[[0]]]
assert store.range_gets[0][4] == [[[TENSOR_METADATA_SIZE]]]
assert store.range_gets[-1][1] == [[make_hidden_data_key(pool_key)]]
assert store.range_gets[-1][3] == [[[TENSOR_METADATA_SIZE]]]
stats = worker.get_operation_stats()
assert stats.data["load_get"][0]["status"] == "ok"
assert stats.data["load_get"][0]["num_keys"] == 1
assert stats.data["load_get"][0]["num_bytes"] == (
stored.numel() * stored.element_size()
)
def test_worker_load_records_error_without_writing_encoder_cache():
pool_key = make_pool_key()
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key.key_metadata,
)
encoder_cache = {}
try:
worker.load(
[MMMeta(identifier=pool_key.identifier, load_spec=LoadSpec(can_load=True))],
encoder_cache,
device="cpu",
)
except HiddenStoreLoadError:
pass
else:
raise AssertionError("missing hidden tensor should fail fast")
assert pool_key.identifier not in encoder_cache
stats = worker.get_operation_stats()
assert stats.data["load_get"][0]["status"] == "error"
assert stats.data["load_get"][0]["num_failed_keys"] == 1
def test_get_tensor_payload_unregisters_target_buffer_after_success():
pool_key = make_pool_key()
stored = torch.zeros((2, 4), dtype=torch.float16)
store = FakeStore()
worker = HiddenStoreWorker(
store_client=MooncakeHiddenStoreClient(store),
tensor_database=HiddenTensorDatabase(),
key_metadata=pool_key.key_metadata,
)
worker.save_tensor(pool_key, stored)
target = torch.empty_like(stored)
worker.store_client.get_tensor_payload(
pool_key,
target.data_ptr(),
target.numel() * target.element_size(),
TENSOR_METADATA_SIZE,
)
assert target.data_ptr() in store.unregistered
def test_get_tensor_payload_unregisters_target_buffer_after_load_error():
pool_key = make_pool_key()
target = torch.empty((2, 4), dtype=torch.float16)
store = FakeStore()
client = MooncakeHiddenStoreClient(store)
try:
client.get_tensor_payload(
pool_key,
target.data_ptr(),
target.numel() * target.element_size(),
TENSOR_METADATA_SIZE,
)
except HiddenStoreLoadError:
pass
else:
raise AssertionError("missing payload should raise")
assert target.data_ptr() in store.unregistered
def _serialize_tensor_object(tensor: torch.Tensor) -> bytes:
tensor = tensor.detach().cpu().contiguous()
nbytes = tensor.numel() * tensor.element_size()
header = struct.pack(
"<IHHiiIIQQ",
TENSOR_OBJECT_MAGIC,
TENSOR_OBJECT_VERSION,
TENSOR_METADATA_SIZE,
TORCH_DTYPE_TO_MOONCAKE_DTYPE[tensor.dtype],
tensor.dim(),
0,
0,
TENSOR_METADATA_SIZE,
nbytes,
)
global_shape = _pack_shape(tuple(tensor.shape))
local_shape = _pack_shape(tuple(tensor.shape))
axes = b"\0" * (32 * 4)
metadata = header + global_shape + local_shape + struct.pack("<II", 0, 0) + axes
assert len(metadata) == TENSOR_METADATA_SIZE
return metadata + tensor.view(torch.uint8).numpy().tobytes()
def _pack_shape(shape: tuple[int, ...]) -> bytes:
dims = list(shape) + [-1] * (8 - len(shape))
return struct.pack("<8q", *dims)

View File

@ -0,0 +1,93 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib
from collections.abc import Callable
from typing import TYPE_CHECKING
from vllm.distributed.ec_transfer.ec_connector.base import (
ECConnectorBase,
ECConnectorRole,
)
from vllm.logger import init_logger
if TYPE_CHECKING:
from vllm.config import ECTransferConfig, VllmConfig
logger = init_logger(__name__)
class ECConnectorFactory:
_registry: dict[str, Callable[[], type[ECConnectorBase]]] = {}
@classmethod
def register_connector(cls, name: str, module_path: str, class_name: str) -> None:
"""Register a connector with a lazy-loading module and class name."""
if name in cls._registry:
raise ValueError(f"Connector '{name}' is already registered.")
def loader() -> type[ECConnectorBase]:
module = importlib.import_module(module_path)
return getattr(module, class_name)
cls._registry[name] = loader
@classmethod
def create_connector(
cls,
config: "VllmConfig",
role: ECConnectorRole,
) -> ECConnectorBase:
ec_transfer_config = config.ec_transfer_config
if ec_transfer_config is None:
raise ValueError("ec_transfer_config must be set to create a connector")
connector_cls = cls.get_connector_class(ec_transfer_config)
logger.info(
"Creating connector with name: %s and engine_id: %s",
connector_cls.__name__,
ec_transfer_config.engine_id,
)
# Connector is explicitly separated into two roles.
# Scheduler connector:
# - Co-locate with scheduler process
# - Should only be used inside the Scheduler class
# Worker connector:
# - Co-locate with worker process
return connector_cls(config, role)
@classmethod
def get_connector_class(
cls, ec_transfer_config: "ECTransferConfig"
) -> type[ECConnectorBase]:
"""Get the connector class by name."""
connector_name = ec_transfer_config.ec_connector
if connector_name is None:
raise ValueError("EC connect must not be None")
connector_module_path = ec_transfer_config.ec_connector_module_path
if connector_module_path is not None and not connector_module_path:
raise ValueError("ec_connector_module_path cannot be an empty string.")
if connector_module_path:
connector_module = importlib.import_module(connector_module_path)
connector_cls = getattr(connector_module, connector_name)
elif connector_name in cls._registry:
connector_cls = cls._registry[connector_name]()
else:
raise ValueError(f"Unsupported connector type: {connector_name}")
return connector_cls
# Register various connectors here.
# The registration should not be done in each individual file, as we want to
# only load the files corresponding to the current connector.
ECConnectorFactory.register_connector(
"ECExampleConnector",
"vllm.distributed.ec_transfer.ec_connector.example_connector",
"ECExampleConnector",
)
ECConnectorFactory.register_connector(
"MooncakeStoreECConnector",
"vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden",
"MooncakeStoreECConnector",
)

View File

@ -0,0 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Hidden-state Mooncake Store EC connector support."""
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.connector import (
MooncakeStoreECConnector,
)
__all__ = ["MooncakeStoreECConnector"]

View File

@ -0,0 +1,377 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""EC connector backed by Mooncake Store for hidden-state tensors."""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from vllm.distributed import (
get_dcp_group,
get_pcp_group,
get_tensor_model_parallel_world_size,
)
from vllm.distributed.ec_transfer.ec_connector.base import (
ECConnectorBase,
ECConnectorMetadata,
ECConnectorRole,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HIDDEN_OBJECT_KIND,
HIDDEN_STORAGE_LAYOUT,
HIDDEN_TENSOR_LAYOUT,
HiddenKeyMetadata,
HiddenSaveRequest,
LoadSpec,
MMMeta,
MooncakeStoreConnectorMetadata,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.store_client import (
MooncakeHiddenStoreClient,
create_mooncake_hidden_store_client,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.worker import (
HiddenLookupClient,
HiddenLookupServer,
HiddenStoreWorker,
)
from vllm.logger import init_logger
from vllm.multimodal.utils import get_mm_features_in_window
from vllm.v1.core.sched.output import SchedulerOutput
if TYPE_CHECKING:
from vllm.config import VllmConfig
from vllm.v1.request import Request
logger = init_logger(__name__)
class MooncakeStoreECConnector(ECConnectorBase):
"""Hidden-state EC connector that stores tensors in Mooncake Store."""
def __init__(
self,
vllm_config: VllmConfig,
role: ECConnectorRole,
store_client: MooncakeHiddenStoreClient | None = None,
):
super().__init__(vllm_config=vllm_config, role=role)
self.lookup_client: HiddenLookupClient | None = None
self.lookup_server: HiddenLookupServer | None = None
self.store_client: MooncakeHiddenStoreClient | None = None
self.worker: HiddenStoreWorker | None = None
assert vllm_config.ec_transfer_config is not None
extra_config = vllm_config.ec_transfer_config.ec_connector_extra_config
self.soft_pin_video_hidden = bool(
extra_config.get("soft_pin_video_hidden", False)
)
self.lookup_async = bool(extra_config.get("lookup_async", True))
if role == ECConnectorRole.SCHEDULER:
if self.is_consumer:
self.lookup_client = HiddenLookupClient(vllm_config)
else:
if not (self.is_producer or self.is_consumer):
return
hidden_key_metadata = build_hidden_key_metadata(vllm_config)
self.store_client = store_client or create_mooncake_hidden_store_client()
self.worker = HiddenStoreWorker(
store_client=self.store_client,
key_metadata=hidden_key_metadata,
)
if self.is_producer:
self.worker.start_sending_thread()
if self.is_consumer and vllm_config.parallel_config.rank == 0:
self.lookup_server = HiddenLookupServer(self.worker, vllm_config)
self.load_specs: dict[str, LoadSpec] = {}
self.lookup_result_cache: dict[str, bool] = {}
self.identifier_waiters: dict[str, set[str]] = {}
self._candidate_consumes: dict[str, set[str]] = {}
self._candidate_loads: dict[str, set[str]] = {}
self._candidate_saves: dict[str, set[str]] = {}
self._load_modalities: dict[str, str | None] = {}
self._save_modalities: dict[str, str | None] = {}
def shutdown(self) -> None:
if self.lookup_client is not None:
self.lookup_client.close()
if self.lookup_server is not None:
self.lookup_server.close()
if self.worker is not None:
self.worker.shutdown()
def has_cache_item(self, identifier: str) -> bool:
if not self.is_consumer:
return False
if not self.lookup_result_cache.get(identifier, False):
self.load_specs.pop(identifier, None)
logger.info(
"hidden_store_scheduler_miss identifier=%s "
"reason=local_lookup_result_miss",
identifier,
)
return False
self.load_specs.setdefault(identifier, LoadSpec(can_load=False))
logger.info(
"hidden_store_scheduler_hit identifier=%s",
identifier,
)
return True
def ensure_cache_available(
self,
request: Request,
num_computed_tokens: int,
) -> bool:
if not self.is_consumer:
return True
if not request.mm_features:
return True
assert self.lookup_client is not None
start = num_computed_tokens
end = request.num_tokens
lo, hi = get_mm_features_in_window(request.mm_features, start, end)
identifiers = list(
dict.fromkeys(
feature.identifier for feature in request.mm_features[lo:hi]
)
)
if not identifiers:
return True
request_id = request.request_id
for identifier in identifiers:
self.identifier_waiters.setdefault(identifier, set()).add(request_id)
unknown_identifiers = [
identifier
for identifier in identifiers
if identifier not in self.lookup_result_cache
]
if not unknown_identifiers:
return True
lookup_results = self.lookup_client.lookup_batch(
unknown_identifiers,
non_block=self.lookup_async,
)
if lookup_results is None:
return False
for identifier in unknown_identifiers:
self.lookup_result_cache[identifier] = lookup_results.get(
identifier,
False,
)
return True
def update_state_after_alloc(self, request: Request, index: int) -> None:
mm_feature = request.mm_features[index]
identifier = mm_feature.identifier
modality = mm_feature.modality
request_id = request.request_id
self._candidate_consumes.setdefault(request_id, set()).add(identifier)
if self.is_consumer and identifier in self.load_specs:
self._candidate_loads.setdefault(request_id, set()).add(identifier)
self._load_modalities[identifier] = modality
if self.is_producer:
self._save_modalities[identifier] = modality
self._candidate_saves.setdefault(request_id, set()).add(identifier)
def build_connector_meta(
self,
scheduler_output: SchedulerOutput,
) -> ECConnectorMetadata:
items_by_identifier: dict[str, MMMeta] = {}
preempted_ids = getattr(scheduler_output, "preempted_req_ids", None) or set()
for request_id, identifiers in self._candidate_consumes.items():
if request_id in preempted_ids:
continue
for identifier in identifiers:
waiters = self.identifier_waiters.get(identifier)
if waiters is not None:
waiters.discard(request_id)
for request_id, identifiers in self._candidate_loads.items():
if request_id in preempted_ids:
continue
for identifier in identifiers:
load_spec = self.load_specs.pop(identifier, None)
if load_spec is None:
continue
load_spec.can_load = True
items_by_identifier[identifier] = MMMeta(
identifier=identifier,
modality=self._load_modalities.get(identifier),
load_spec=load_spec,
)
for request_id, identifiers in self._candidate_saves.items():
if request_id in preempted_ids:
continue
for identifier in identifiers:
item = items_by_identifier.get(identifier)
if item is None:
item = MMMeta(
identifier=identifier,
modality=self._save_modalities.get(identifier),
)
items_by_identifier[identifier] = item
item.can_save = True
if item.modality is None:
item.modality = self._save_modalities.get(identifier)
finished_req_ids = getattr(scheduler_output, "finished_req_ids", set())
for finished_req_id in finished_req_ids:
for waiters in self.identifier_waiters.values():
waiters.discard(finished_req_id)
self._cleanup_lookup_results_without_waiters()
metadata = MooncakeStoreConnectorMetadata(
items=list(items_by_identifier.values()),
)
self._candidate_consumes.clear()
self._candidate_loads.clear()
self._candidate_saves.clear()
self._load_modalities.clear()
self._save_modalities.clear()
return metadata
def _cleanup_lookup_results_without_waiters(self) -> None:
for identifier, waiters in list(self.identifier_waiters.items()):
if waiters:
continue
del self.identifier_waiters[identifier]
self.lookup_result_cache.pop(identifier, None)
self.load_specs.pop(identifier, None)
if self.lookup_client is not None:
self.lookup_client.discard(identifier)
def start_load_caches(
self,
encoder_cache: dict[str, torch.Tensor],
**kwargs,
) -> None:
metadata = self._get_connector_metadata()
assert isinstance(metadata, MooncakeStoreConnectorMetadata)
assert self.worker is not None
self.worker.load(
metadata.items,
encoder_cache,
device=kwargs.get("device"),
)
def save_caches(
self,
encoder_cache: dict[str, torch.Tensor],
mm_hash: str,
**kwargs,
) -> None:
if not self.is_producer:
return
assert self.worker is not None
identifier = mm_hash
if identifier not in encoder_cache:
logger.warning(
"Skip hidden store save; identifier %s is missing",
identifier,
)
return
item = self._find_metadata_item(identifier)
if item is None or not item.can_save:
logger.debug(
"Skip hidden store save; identifier %s has no save plan",
identifier,
)
return
pool_key = self.worker.make_pool_key(identifier)
self.worker.enqueue_save(
HiddenSaveRequest(
pool_key=pool_key,
tensor=encoder_cache[identifier],
with_soft_pin=self._should_soft_pin(item),
)
)
def get_finished(
self, finished_req_ids: set[str]
) -> tuple[set[str] | None, set[str] | None]:
if self.worker is None or not self.is_producer:
return None, None
finished_sending = self.worker.get_finished_sending()
failed_sending = self.worker.get_failed_sending()
for identifier, reason in failed_sending.items():
logger.error(
"hidden_store_save_failed identifier=%s reason=%s",
identifier,
reason,
)
return finished_sending or None, None
def _find_metadata_item(self, identifier: str) -> MMMeta | None:
metadata = self._get_connector_metadata()
assert isinstance(metadata, MooncakeStoreConnectorMetadata)
for item in metadata.items:
if item.identifier == identifier:
return item
return None
def _should_soft_pin(self, item: MMMeta) -> bool:
return self.soft_pin_video_hidden and item.modality == "video"
def build_hidden_key_metadata(vllm_config: VllmConfig) -> HiddenKeyMetadata:
model_config = vllm_config.model_config
parallel_config = vllm_config.parallel_config
assert vllm_config.ec_transfer_config is not None
extra_config = vllm_config.ec_transfer_config.ec_connector_extra_config
multimodal_config = getattr(model_config, "multimodal_config", None)
compute_hash = getattr(multimodal_config, "compute_hash", None)
mm_encoder_config_hash = (
compute_hash() if callable(compute_hash) else "encoder:default"
)
tp_size = get_tensor_model_parallel_world_size()
pp_size = parallel_config.pipeline_parallel_size
pcp_size = get_pcp_group().world_size
dcp_size = get_dcp_group().world_size
mm_encoder_tp_mode = getattr(
multimodal_config,
"mm_encoder_tp_mode",
"unknown",
)
parallel = (
f"tp:{tp_size}"
f"@pp:{pp_size}"
f"@pcp:{pcp_size}"
f"@dcp:{dcp_size}"
f"@mm_tp:{mm_encoder_tp_mode}"
)
return HiddenKeyMetadata(
cache_prefix=str(
extra_config.get(
"hidden_cache_prefix",
extra_config.get("cache_prefix", ""),
)
),
kind=HIDDEN_OBJECT_KIND,
model_name=model_config.model.rstrip("/").split("/")[-1],
encoder=str(mm_encoder_config_hash),
storage=HIDDEN_STORAGE_LAYOUT,
parallel=parallel,
tensor_layout=HIDDEN_TENSOR_LAYOUT,
)

View File

@ -0,0 +1,202 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Data classes for the hidden-state Mooncake Store EC connector."""
from __future__ import annotations
from dataclasses import dataclass, field
import torch
from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorMetadata
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.keys import (
escape_key_part,
make_hidden_data_key,
)
HIDDEN_OBJECT_KIND = "encoder_output"
HIDDEN_STORAGE_LAYOUT = "replicated_object"
HIDDEN_TENSOR_LAYOUT = "tensor"
HIDDEN_PROTOCOL_VERSION = "v1"
MOONCAKE_TENSOR_METADATA_NBYTES = 304
@dataclass(frozen=True)
class HiddenKeyMetadata:
"""Metadata that defines the semantic namespace for hidden reuse."""
cache_prefix: str
kind: str
model_name: str
encoder: str
storage: str
parallel: str
tensor_layout: str
@dataclass(frozen=True, order=True)
class HiddenPoolKey:
"""Key for addressing one hidden tensor in the distributed store."""
key_metadata: HiddenKeyMetadata
identifier: str
def to_string(self) -> str:
meta = self.key_metadata
prefix = (
f"{escape_key_part(meta.cache_prefix)}@" if meta.cache_prefix else ""
)
return (
f"{prefix}hidden"
f"@kind:{escape_key_part(meta.kind)}"
f"@model:{escape_key_part(meta.model_name)}"
f"@encoder:{escape_key_part(meta.encoder)}"
f"@storage:{escape_key_part(meta.storage)}"
f"@parallel:{escape_key_part(meta.parallel)}"
f"@tensor_layout:{escape_key_part(meta.tensor_layout)}"
f"@id:{escape_key_part(self.identifier)}"
)
@dataclass
class MMMeta:
"""Per hidden object metadata passed from scheduler to worker."""
identifier: str
modality: str | None = None
can_save: bool = False
load_spec: LoadSpec | None = None
@dataclass(frozen=True)
class TensorMeta:
"""Canonical contiguous tensor descriptor for one hidden store object."""
pool_key: HiddenPoolKey
protocol_version: str
layout: str
shape: tuple[int, ...]
dtype: str
nbytes: int
device_type: str
data_offset: int = MOONCAKE_TENSOR_METADATA_NBYTES
producer_stage: str = "encoder"
@dataclass
class LoadSpec:
"""Specification for loading a hidden tensor from external store."""
can_load: bool = False
@dataclass
class HiddenSaveRequest:
"""Specification for asynchronously storing one hidden tensor."""
pool_key: HiddenPoolKey
tensor: torch.Tensor
with_soft_pin: bool = False
@property
def identifier(self) -> str:
return self.pool_key.identifier
@dataclass
class MooncakeStoreConnectorMetadata(ECConnectorMetadata):
"""Metadata passed from scheduler to worker for hidden store operations."""
items: list[MMMeta] = field(default_factory=list)
def add_item(self, item: MMMeta) -> None:
self.items.append(item)
@dataclass
class HiddenStoreOperationStats:
"""Minimal per-operation telemetry aligned with Mooncake KV store stats."""
data: dict[str, list[dict[str, int | float | str]]] = field(default_factory=dict)
def is_empty(self) -> bool:
return not self.data
def record_operation(
self,
operation: str,
duration_seconds: float,
num_keys: int,
*,
num_bytes: int = 0,
status: str = "ok",
num_failed_keys: int = 0,
) -> None:
self.data.setdefault(operation, []).append(
{
"duration_seconds": duration_seconds,
"num_keys": num_keys,
"num_bytes": num_bytes,
"status": status,
"num_failed_keys": num_failed_keys,
}
)
class HiddenTensorDatabase:
"""Maps hidden tensors to store keys and GPU memory descriptors."""
def prepare_value(
self,
pool_key: HiddenPoolKey,
tensor: torch.Tensor,
) -> tuple[str, list[int], list[int]]:
return (
make_hidden_data_key(pool_key),
[tensor.data_ptr()],
[tensor.numel() * tensor.element_size()],
)
def build_tensor_meta(
pool_key: HiddenPoolKey,
tensor: torch.Tensor,
) -> TensorMeta:
"""Build metadata for the canonical stored hidden tensor layout."""
if not tensor.is_contiguous():
raise ValueError("Hidden tensor descriptor requires a contiguous tensor")
return TensorMeta(
pool_key=pool_key,
protocol_version=HIDDEN_PROTOCOL_VERSION,
layout=HIDDEN_TENSOR_LAYOUT,
shape=tuple(tensor.shape),
dtype=str(tensor.dtype),
nbytes=tensor.numel() * tensor.element_size(),
device_type=tensor.device.type,
data_offset=MOONCAKE_TENSOR_METADATA_NBYTES,
)
def validate_loaded_tensor(tensor: torch.Tensor, meta: TensorMeta) -> None:
if tuple(tensor.shape) != tuple(meta.shape):
raise ValueError(
"Hidden tensor shape mismatch: "
f"actual={tuple(tensor.shape)} expected={meta.shape}"
)
if str(tensor.dtype) != meta.dtype:
raise ValueError(
"Hidden tensor dtype mismatch: "
f"actual={tensor.dtype} expected={meta.dtype}"
)
actual_nbytes = tensor.numel() * tensor.element_size()
if actual_nbytes != meta.nbytes:
raise ValueError(
"Hidden tensor nbytes mismatch: "
f"actual={actual_nbytes} expected={meta.nbytes}"
)
if meta.layout != HIDDEN_TENSOR_LAYOUT:
raise ValueError(f"Unsupported hidden tensor layout: {meta.layout}")

View File

@ -0,0 +1,22 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Store key helpers for the hidden-state Mooncake connector."""
from __future__ import annotations
from typing import TYPE_CHECKING
from urllib.parse import quote
if TYPE_CHECKING:
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HiddenPoolKey,
)
def escape_key_part(value: str) -> str:
"""Escape one key component while keeping simple values readable."""
return quote(str(value), safe="-_.~")
def make_hidden_data_key(pool_key: "HiddenPoolKey") -> str:
return pool_key.to_string()

View File

@ -0,0 +1,567 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Thin Mooncake Store client for hidden-state objects."""
from __future__ import annotations
import copy
import ctypes
import json
import os
import re
import struct
from dataclasses import dataclass
from typing import Any
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HIDDEN_PROTOCOL_VERSION,
HIDDEN_TENSOR_LAYOUT,
MOONCAKE_TENSOR_METADATA_NBYTES,
HiddenPoolKey,
TensorMeta,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.keys import (
make_hidden_data_key,
)
from vllm.logger import init_logger
from vllm.utils.network_utils import get_ip
logger = init_logger(__name__)
DEFAULT_GLOBAL_SEGMENT_SIZE = 4 * 1024 * 1024 * 1024
DEFAULT_LOCAL_BUFFER_SIZE = 4 * 1024 * 1024 * 1024
_MOONCAKE_TENSOR_OBJECT_MAGIC = 0x4D4F4F4E
_MOONCAKE_TENSOR_OBJECT_VERSION = 1
_MOONCAKE_TENSOR_HEADER_FORMAT = "<IHHiiIIQQ"
_MOONCAKE_TENSOR_HEADER_NBYTES = struct.calcsize(_MOONCAKE_TENSOR_HEADER_FORMAT)
_MOONCAKE_TENSOR_LOCAL_SHAPE_OFFSET = _MOONCAKE_TENSOR_HEADER_NBYTES + 64
_MOONCAKE_DTYPE_TO_TORCH_DTYPE = {
0: "torch.float32",
1: "torch.float64",
2: "torch.int8",
3: "torch.uint8",
4: "torch.int16",
5: "torch.uint16",
6: "torch.int32",
7: "torch.uint32",
8: "torch.int64",
9: "torch.uint64",
10: "torch.bool",
11: "torch.float16",
12: "torch.bfloat16",
13: "torch.float8_e4m3fn",
14: "torch.float8_e5m2",
}
_TORCH_DTYPE_TO_MOONCAKE_DTYPE = {
"torch.float32": 0,
"torch.float64": 1,
"torch.int8": 2,
"torch.uint8": 3,
"torch.int16": 4,
"torch.uint16": 5,
"torch.int32": 6,
"torch.uint32": 7,
"torch.int64": 8,
"torch.uint64": 9,
"torch.bool": 10,
"torch.float16": 11,
"torch.bfloat16": 12,
"torch.float8_e4m3fn": 13,
"torch.float8_e5m2": 14,
}
_SUPPORTED_HIDDEN_TORCH_DTYPES = {
"torch.float16",
"torch.bfloat16",
"torch.float32",
}
@dataclass
class MooncakeHiddenStoreConfig:
metadata_server: str
master_server_address: str
protocol: str
device_name: str
mode: str = "embedded"
global_segment_size: int = DEFAULT_GLOBAL_SEGMENT_SIZE
local_buffer_size: int = DEFAULT_LOCAL_BUFFER_SIZE
@staticmethod
def from_file(file_path: str) -> MooncakeHiddenStoreConfig:
with open(file_path, encoding="utf-8") as file:
config = json.load(file)
mode = config.get("mode", "embedded")
return MooncakeHiddenStoreConfig(
metadata_server=config.get("metadata_server", ""),
master_server_address=config.get("master_server_address", ""),
protocol=config.get("protocol", "rdma"),
device_name=config.get("device_name", ""),
mode=mode,
global_segment_size=_parse_size(
config.get(
"global_segment_size",
0 if mode == "standalone-store" else DEFAULT_GLOBAL_SEGMENT_SIZE,
)
),
local_buffer_size=_parse_size(
config.get("local_buffer_size", DEFAULT_LOCAL_BUFFER_SIZE)
),
)
@staticmethod
def load_from_env() -> MooncakeHiddenStoreConfig:
config_path = os.getenv("MOONCAKE_CONFIG_PATH")
if not config_path:
raise ValueError(
"The environment variable 'MOONCAKE_CONFIG_PATH' is not set."
)
return MooncakeHiddenStoreConfig.from_file(config_path)
def _parse_size(value: Any) -> int:
if isinstance(value, int):
return value
if not isinstance(value, str):
return int(value)
cleaned = value.strip().lower()
match = re.match(r"^\s*([\d.]+)\s*(gb|mb|kb|b)?\s*$", cleaned)
if not match:
raise ValueError(f"Invalid size format: {value!r}")
multipliers = {
"gb": 1024**3,
"mb": 1024**2,
"kb": 1024,
"b": 1,
None: 1,
}
return int(float(match.group(1)) * multipliers[match.group(2)])
def create_mooncake_hidden_store_client() -> MooncakeHiddenStoreClient:
try:
from mooncake.store import ( # type: ignore
MooncakeDistributedStore,
ReplicateConfig,
)
except ImportError as e:
raise ImportError(
"Please install mooncake to run vLLM with " "MooncakeStoreECConnector."
) from e
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake import rdma_utils
config = MooncakeHiddenStoreConfig.load_from_env()
config.device_name = rdma_utils.get_configured_worker_rnic(
protocol=config.protocol,
configured_device=config.device_name,
)
store = MooncakeDistributedStore()
local_ip = get_ip()
local_hostname = rdma_utils.get_requester_local_hostname(local_ip)
ret = store.setup(
local_hostname,
config.metadata_server,
config.global_segment_size,
config.local_buffer_size,
config.protocol,
config.device_name,
config.master_server_address,
)
if ret != 0:
raise RuntimeError("Initialize MooncakeDistributedStore failed.")
logger.info(
"Initialized hidden Mooncake store mode=%s global_segment_size=%d "
"local_buffer_size=%d",
config.mode,
config.global_segment_size,
config.local_buffer_size,
)
return MooncakeHiddenStoreClient(store, replicate_config=ReplicateConfig())
class HiddenStoreError(RuntimeError):
pass
class HiddenStoreLoadError(HiddenStoreError):
pass
class HiddenStoreSaveError(HiddenStoreError):
pass
class MooncakeHiddenStoreClient:
"""Wraps Mooncake object and buffer APIs used by hidden transfer."""
def __init__(self, store: Any, replicate_config: Any | None = None):
self.store = store
self.replicate_config = replicate_config
def close(self) -> None:
"""Best-effort shutdown for Mooncake store implementations."""
for method_name in ("close", "teardown", "disconnect", "finalize"):
close_fn = getattr(self.store, method_name, None)
if close_fn is None:
continue
try:
close_fn()
except Exception:
logger.warning(
"failed to close hidden Mooncake store with %s()",
method_name,
exc_info=True,
)
return
def exists(self, pool_key: HiddenPoolKey) -> bool:
data_key = make_hidden_data_key(pool_key)
states = self.store.batch_is_exist([data_key])
return len(states) == 1 and states[0] == 1
def batch_exists(self, pool_keys: list[HiddenPoolKey]) -> list[bool]:
if not pool_keys:
return []
keys = [make_hidden_data_key(pool_key) for pool_key in pool_keys]
states = self.store.batch_is_exist(keys)
return [state == 1 for state in states]
def get_tensor_meta(self, pool_key: HiddenPoolKey) -> TensorMeta | None:
metadata = self._read_range(
pool_key,
src_offset=0,
size=MOONCAKE_TENSOR_METADATA_NBYTES,
)
if metadata is None:
return None
try:
return _decode_mooncake_tensor_metadata(pool_key, metadata)
except HiddenStoreLoadError:
logger.exception(
"failed to decode hidden Mooncake tensor metadata for %s",
pool_key.to_string(),
)
return None
def put_tensor(
self,
pool_key: HiddenPoolKey,
tensor: Any,
*,
with_soft_pin: bool = False,
) -> None:
_validate_supported_hidden_tensor_dtype(tensor)
key = make_hidden_data_key(pool_key)
replicate_config = _make_hidden_replicate_config(
self.replicate_config,
with_soft_pin=with_soft_pin,
)
batch_put_from_multi_buffers = getattr(
self.store,
"batch_put_from_multi_buffers",
None,
)
if batch_put_from_multi_buffers is not None:
self._put_tensor_from_buffers(
pool_key,
tensor,
replicate_config=replicate_config,
)
return
if replicate_config is None:
put_fn = getattr(self.store, "put_tensor", None)
if put_fn is None:
raise HiddenStoreSaveError(
"Mooncake Hidden Store requires put_tensor or pub_tensor "
"support for single-object hidden tensors."
)
ret = put_fn(key, tensor)
else:
put_fn = getattr(self.store, "pub_tensor", None)
if put_fn is None:
raise HiddenStoreSaveError(
"Mooncake Hidden Store requires pub_tensor support when "
"a ReplicateConfig is configured."
)
ret = put_fn(key, tensor, replicate_config)
if ret != 0:
raise HiddenStoreSaveError(
f"failed to put hidden tensor for {pool_key.to_string()}: {ret}"
)
def _put_tensor_from_buffers(
self,
pool_key: HiddenPoolKey,
tensor: Any,
*,
replicate_config: Any | None,
) -> None:
if not tensor.is_contiguous():
raise HiddenStoreSaveError(
"hidden tensor must be contiguous before batch buffer put"
)
data_size = tensor.numel() * tensor.element_size()
metadata = _encode_mooncake_tensor_metadata(tensor)
metadata_buffer = (ctypes.c_ubyte * len(metadata)).from_buffer_copy(metadata)
metadata_ptr = ctypes.addressof(metadata_buffer)
payload_ptr = tensor.data_ptr()
registered_addrs: list[int] = []
try:
self.register_tensor(payload_ptr, data_size)
registered_addrs.append(payload_ptr)
self.register_tensor(metadata_ptr, len(metadata))
registered_addrs.append(metadata_ptr)
key = make_hidden_data_key(pool_key)
results = self.store.batch_put_from_multi_buffers(
[key],
[[metadata_ptr, payload_ptr]],
[[len(metadata), data_size]],
replicate_config,
)
failed = [result for result in results if result < 0]
if failed:
raise HiddenStoreSaveError(
"failed to put hidden tensor for "
f"{pool_key.to_string()}: {failed}"
)
finally:
for addr in reversed(registered_addrs):
self.unregister_tensor(addr)
def register_tensor(self, addr: int, size: int) -> None:
ret = self.store.register_buffer(addr, size)
if ret != 0:
raise HiddenStoreError(
f"failed to register hidden buffer addr={addr:#x} size={size}: {ret}"
)
def unregister_tensor(self, addr: int) -> None:
unregister_fn = getattr(self.store, "unregister_buffer", None)
if unregister_fn is None:
return
try:
ret = unregister_fn(addr)
except Exception:
logger.warning(
"failed to unregister hidden buffer addr=%#x",
addr,
exc_info=True,
)
return
if ret != 0:
logger.warning(
"unregister hidden buffer failed addr=%#x ret=%s",
addr,
ret,
)
def get_tensor_payload(
self,
pool_key: HiddenPoolKey,
addr: int,
size: int,
src_offset: int,
) -> int:
self.register_tensor(addr, size)
try:
key = make_hidden_data_key(pool_key)
results = self.store.get_into_ranges(
[addr],
[[key]],
[[[0]]],
[[[src_offset]]],
[[[size]]],
)
result = _single_range_result(results)
if result != size:
raise HiddenStoreLoadError(
"failed to get hidden tensor payload for "
f"{pool_key.to_string()}: {result}"
)
return result
finally:
self.unregister_tensor(addr)
def _read_range(
self,
pool_key: HiddenPoolKey,
*,
src_offset: int,
size: int,
) -> bytes | None:
buffer = (ctypes.c_ubyte * size)()
buffer_ptr = ctypes.addressof(buffer)
self.register_tensor(buffer_ptr, size)
key = make_hidden_data_key(pool_key)
try:
results = self.store.get_into_ranges(
[buffer_ptr],
[[key]],
[[[0]]],
[[[src_offset]]],
[[[size]]],
)
finally:
self.unregister_tensor(buffer_ptr)
if _single_range_result(results) != size:
return None
return bytes(buffer)
def _single_range_result(results: Any) -> int:
try:
return int(results[0][0][0])
except Exception:
return -1
def _decode_mooncake_tensor_metadata(
pool_key: HiddenPoolKey,
metadata: bytes,
) -> TensorMeta:
if len(metadata) < MOONCAKE_TENSOR_METADATA_NBYTES:
raise HiddenStoreLoadError(
f"hidden tensor metadata is too small: {len(metadata)}"
)
(
magic,
version,
header_size,
dtype,
ndim,
_layout_kind,
_reserved_flags,
data_offset,
data_bytes,
) = struct.unpack_from(_MOONCAKE_TENSOR_HEADER_FORMAT, metadata, 0)
if (
magic != _MOONCAKE_TENSOR_OBJECT_MAGIC
or version != _MOONCAKE_TENSOR_OBJECT_VERSION
or header_size != MOONCAKE_TENSOR_METADATA_NBYTES
):
raise HiddenStoreLoadError(
"invalid Mooncake tensor metadata header for " f"{pool_key.to_string()}"
)
if ndim < 0 or ndim > 8:
raise HiddenStoreLoadError(
f"invalid hidden tensor ndim for {pool_key.to_string()}: {ndim}"
)
if dtype not in _MOONCAKE_DTYPE_TO_TORCH_DTYPE:
raise HiddenStoreLoadError(
f"unsupported Mooncake tensor dtype for {pool_key.to_string()}: {dtype}"
)
local_shape = struct.unpack_from(
"<8q",
metadata,
_MOONCAKE_TENSOR_LOCAL_SHAPE_OFFSET,
)
shape = tuple(int(dim) for dim in local_shape[:ndim])
if any(dim < 0 for dim in shape):
raise HiddenStoreLoadError(
f"invalid hidden tensor shape for {pool_key.to_string()}: {shape}"
)
return TensorMeta(
pool_key=pool_key,
protocol_version=HIDDEN_PROTOCOL_VERSION,
layout=HIDDEN_TENSOR_LAYOUT,
shape=shape,
dtype=_MOONCAKE_DTYPE_TO_TORCH_DTYPE[dtype],
nbytes=int(data_bytes),
device_type="cpu",
data_offset=int(data_offset),
)
def _encode_mooncake_tensor_metadata(tensor: Any) -> bytes:
dtype = str(tensor.dtype)
_validate_supported_hidden_tensor_dtype(tensor)
shape = tuple(int(dim) for dim in tensor.shape)
if len(shape) > 8:
raise HiddenStoreSaveError(
f"hidden tensor has too many dimensions: {len(shape)}"
)
nbytes = tensor.numel() * tensor.element_size()
header = struct.pack(
_MOONCAKE_TENSOR_HEADER_FORMAT,
_MOONCAKE_TENSOR_OBJECT_MAGIC,
_MOONCAKE_TENSOR_OBJECT_VERSION,
MOONCAKE_TENSOR_METADATA_NBYTES,
_TORCH_DTYPE_TO_MOONCAKE_DTYPE[dtype],
len(shape),
0,
0,
MOONCAKE_TENSOR_METADATA_NBYTES,
nbytes,
)
dims = shape + (-1,) * (8 - len(shape))
tensor_shape = struct.pack("<8q", *dims)
axes = b"\0" * (32 * 4)
metadata = header + tensor_shape + tensor_shape + struct.pack("<II", 0, 0) + axes
if len(metadata) != MOONCAKE_TENSOR_METADATA_NBYTES:
raise HiddenStoreSaveError(
f"invalid Mooncake tensor metadata size: {len(metadata)}"
)
return metadata
def _validate_supported_hidden_tensor_dtype(tensor: Any) -> None:
dtype = str(tensor.dtype)
if dtype not in _SUPPORTED_HIDDEN_TORCH_DTYPES:
raise HiddenStoreSaveError(f"unsupported hidden tensor dtype: {dtype}")
def _make_hidden_replicate_config(
replicate_config: Any | None,
*,
with_soft_pin: bool,
) -> Any | None:
if replicate_config is None:
return None
config = _clone_replicate_config(replicate_config)
hidden_state_data_type = _get_hidden_state_object_data_type()
if hidden_state_data_type is not None and hasattr(config, "data_type"):
config.data_type = hidden_state_data_type
if hasattr(config, "with_soft_pin"):
config.with_soft_pin = bool(config.with_soft_pin) or with_soft_pin
return config
def _clone_replicate_config(replicate_config: Any) -> Any:
try:
return copy.copy(replicate_config)
except Exception:
config = type(replicate_config)()
for attr in (
"replica_num",
"nof_replica_num",
"with_soft_pin",
"with_hard_pin",
"preferred_segments",
"preferred_nof_segments",
"preferred_segment",
"prefer_alloc_in_same_node",
"data_type",
"group_ids",
):
if hasattr(replicate_config, attr) and hasattr(config, attr):
setattr(config, attr, getattr(replicate_config, attr))
return config
def _get_hidden_state_object_data_type() -> Any | None:
try:
from mooncake.store import ObjectDataType # type: ignore
except Exception:
return None
hidden_state_type = getattr(ObjectDataType, "HIDDEN_STATE", None)
if hidden_state_type is not None:
return hidden_state_type
return getattr(ObjectDataType, "TENSOR", None)

View File

@ -0,0 +1,643 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Worker-side hidden-state load/save logic for Mooncake Store."""
from __future__ import annotations
import os
import queue
import socket
import threading
import time
from concurrent.futures import Future, ThreadPoolExecutor
import torch
import zmq
import vllm.envs as envs
from vllm.config import VllmConfig
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.data import (
HiddenKeyMetadata,
HiddenPoolKey,
HiddenSaveRequest,
HiddenStoreOperationStats,
HiddenTensorDatabase,
MMMeta,
build_tensor_meta,
validate_loaded_tensor,
)
from vllm.distributed.ec_transfer.ec_connector.mooncake_store_hidden.store_client import (
HiddenStoreLoadError,
MooncakeHiddenStoreClient,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import (
get_mooncake_dp_engine_index,
)
from vllm.logger import init_logger
from vllm.utils.network_utils import make_zmq_socket
logger = init_logger(__name__)
LOOKUP_MSG = b"LOOKUP"
BATCH_LOOKUP_MSG = b"BATCH_LOOKUP"
RESP_BATCH = b"BATCH"
RESP_HIT = b"HIT"
RESP_MISS = b"MISS"
RESP_ERR = b"ERR"
THREAD_JOIN_TIMEOUT_SECONDS = 5.0
class HiddenStoreWorker:
"""Synchronous hidden tensor load/save path used by the EC connector."""
def __init__(
self,
store_client: MooncakeHiddenStoreClient,
tensor_database: HiddenTensorDatabase | None = None,
key_metadata: HiddenKeyMetadata | None = None,
):
self.store_client = store_client
self.tensor_database = tensor_database or HiddenTensorDatabase()
self.key_metadata = key_metadata
self.sending_thread: HiddenStoreSendingThread | None = None
self._operation_stats_lock = threading.Lock()
self._operation_stats = HiddenStoreOperationStats()
def make_pool_key(self, identifier: str) -> HiddenPoolKey:
assert self.key_metadata is not None
return HiddenPoolKey(
key_metadata=self.key_metadata,
identifier=identifier,
)
def start_sending_thread(self) -> None:
if self.sending_thread is not None:
return
self.sending_thread = HiddenStoreSendingThread(self)
self.sending_thread.start()
def enqueue_save(self, request: HiddenSaveRequest) -> None:
if self.sending_thread is None:
self.save_tensor(
request.pool_key,
request.tensor,
with_soft_pin=request.with_soft_pin,
)
return
self.sending_thread.add_request(request)
def get_finished_sending(self) -> set[str]:
if self.sending_thread is None:
return set()
return self.sending_thread.get_and_clear_finished_identifiers()
def get_failed_sending(self) -> dict[str, str]:
if self.sending_thread is None:
return {}
return self.sending_thread.get_and_clear_failure_reasons()
def get_operation_stats(self) -> HiddenStoreOperationStats | None:
with self._operation_stats_lock:
if self._operation_stats.is_empty():
return None
stats = self._operation_stats
self._operation_stats = HiddenStoreOperationStats()
return stats
def _record_operation(
self,
operation: str,
duration_seconds: float,
num_keys: int,
*,
num_bytes: int = 0,
status: str = "ok",
num_failed_keys: int = 0,
) -> None:
with self._operation_stats_lock:
self._operation_stats.record_operation(
operation=operation,
duration_seconds=duration_seconds,
num_keys=num_keys,
num_bytes=num_bytes,
status=status,
num_failed_keys=num_failed_keys,
)
def shutdown(self) -> None:
if self.sending_thread is not None:
self.sending_thread.close()
self.sending_thread = None
close_fn = getattr(self.store_client, "close", None)
if close_fn is not None:
close_fn()
def lookup(self, identifier: str) -> bool:
"""Return whether the hidden object exists in Mooncake Store."""
return self.lookup_batch([identifier]).get(identifier, False)
def lookup_batch(self, identifiers: list[str]) -> dict[str, bool]:
"""Return whether hidden objects exist in Mooncake Store."""
pool_keys = [self.make_pool_key(identifier) for identifier in identifiers]
started = time.perf_counter()
try:
exists = self.store_client.batch_exists(pool_keys)
except Exception:
self._record_operation(
"lookup_exists",
time.perf_counter() - started,
len(pool_keys),
status="error",
num_failed_keys=len(pool_keys),
)
raise
failed_keys = sum(1 for hit in exists if not hit)
self._record_operation(
"lookup_exists",
time.perf_counter() - started,
len(pool_keys),
status="miss" if failed_keys else "ok",
num_failed_keys=failed_keys,
)
results = dict(zip(identifiers, exists, strict=True))
for pool_key, hit in zip(pool_keys, exists, strict=True):
if hit:
logger.info(
"hidden_store_lookup_hit identifier=%s hidden_pool_key=%s",
pool_key.identifier,
pool_key.to_string(),
)
else:
logger.info(
"hidden_store_lookup_miss identifier=%s hidden_pool_key=%s "
"reason=missing_object",
pool_key.identifier,
pool_key.to_string(),
)
return results
def save_tensor(
self,
pool_key: HiddenPoolKey,
tensor: torch.Tensor,
with_soft_pin: bool = False,
) -> None:
exists_started = time.perf_counter()
try:
exists = self.store_client.exists(pool_key)
except Exception:
self._record_operation(
"save_exists",
time.perf_counter() - exists_started,
1,
status="error",
num_failed_keys=1,
)
raise
self._record_operation(
"save_exists",
time.perf_counter() - exists_started,
1,
status="ok" if exists else "miss",
)
if exists:
logger.info(
"hidden_store_save_skip identifier=%s hidden_pool_key=%s "
"reason=exists",
pool_key.identifier,
pool_key.to_string(),
)
return
started = time.perf_counter()
stored_tensor = tensor if tensor.is_contiguous() else tensor.contiguous()
used_staging = stored_tensor is not tensor
tensor_meta = build_tensor_meta(pool_key, stored_tensor)
try:
self.store_client.put_tensor(
pool_key,
stored_tensor,
with_soft_pin=with_soft_pin,
)
except Exception:
self._record_operation(
"save_put",
time.perf_counter() - started,
1,
num_bytes=tensor_meta.nbytes,
status="error",
num_failed_keys=1,
)
raise
self._record_operation(
"save_put",
time.perf_counter() - started,
1,
num_bytes=tensor_meta.nbytes,
status="ok",
)
logger.info(
"hidden_store_put identifier=%s hidden_pool_key=%s nbytes=%d "
"used_staging=%s hidden_store_put_ms=%.3f",
pool_key.identifier,
pool_key.to_string(),
tensor_meta.nbytes,
used_staging,
(time.perf_counter() - started) * 1000.0,
)
def load(
self,
items: list[MMMeta],
encoder_cache: dict[str, torch.Tensor],
*,
device: torch.device | str | None = None,
) -> None:
for item in items:
load_spec = item.load_spec
if load_spec is None or not load_spec.can_load:
continue
if item.identifier in encoder_cache:
logger.debug(
"hidden_store_load_skip identifier=%s "
"reason=local_encoder_cache",
item.identifier,
)
continue
started = time.perf_counter()
pool_key = self.make_pool_key(item.identifier)
tensor_meta = None
load_stage = "metadata"
try:
tensor_meta = self.store_client.get_tensor_meta(pool_key)
if tensor_meta is None:
raise HiddenStoreLoadError(
"failed to load hidden tensor metadata for "
f"{pool_key.to_string()}"
)
load_stage = "allocate"
target_device = device
if target_device is None:
target_device = "cuda" if torch.cuda.is_available() else None
target = torch.empty(
tensor_meta.shape,
dtype=_resolve_torch_dtype(tensor_meta.dtype),
device=target_device,
)
_data_key, addrs, sizes = self.tensor_database.prepare_value(
pool_key,
target,
)
load_stage = "payload"
self.store_client.get_tensor_payload(
pool_key,
addrs[0],
sizes[0],
tensor_meta.data_offset,
)
load_stage = "validate"
validate_loaded_tensor(target, tensor_meta)
except Exception as e:
self._record_operation(
"load_get",
time.perf_counter() - started,
1,
num_bytes=tensor_meta.nbytes if tensor_meta is not None else 0,
status="error",
num_failed_keys=1,
)
logger.exception(
"hidden_store_load_failed identifier=%s hidden_pool_key=%s "
"stage=%s shape=%s dtype=%s nbytes=%s error=%s",
item.identifier,
pool_key.to_string(),
load_stage,
tensor_meta.shape if tensor_meta is not None else None,
tensor_meta.dtype if tensor_meta is not None else None,
tensor_meta.nbytes if tensor_meta is not None else 0,
e,
)
raise
encoder_cache[item.identifier] = target
self._record_operation(
"load_get",
time.perf_counter() - started,
1,
num_bytes=tensor_meta.nbytes,
status="ok",
)
logger.info(
"hidden_store_get identifier=%s hidden_pool_key=%s nbytes=%d "
"hidden_store_get_ms=%.3f",
item.identifier,
pool_key.to_string(),
tensor_meta.nbytes,
(time.perf_counter() - started) * 1000.0,
)
def _resolve_torch_dtype(dtype: str) -> torch.dtype:
if dtype == "torch.float16":
return torch.float16
if dtype == "torch.bfloat16":
return torch.bfloat16
if dtype == "torch.float32":
return torch.float32
raise HiddenStoreLoadError(f"unsupported hidden tensor dtype: {dtype}")
class HiddenStoreSendingThread(threading.Thread):
"""Background thread for storing hidden tensors to the store."""
def __init__(self, store_worker: HiddenStoreWorker):
super().__init__(daemon=True, name="HiddenStoreSendingThread")
self.store_worker = store_worker
self.request_queue: queue.Queue[HiddenSaveRequest | None] = queue.Queue()
self.done_task_lock = threading.Lock()
self.finished_identifiers: set[str] = set()
self.failed_identifiers: set[str] = set()
self.failure_reasons: dict[str, str] = {}
self._closed = threading.Event()
def add_request(self, request: HiddenSaveRequest) -> None:
self.request_queue.put(request)
def get_and_clear_finished_identifiers(self) -> set[str]:
with self.done_task_lock:
finished = self.finished_identifiers.copy()
self.finished_identifiers.clear()
return finished
def get_and_clear_failed_identifiers(self) -> set[str]:
with self.done_task_lock:
failed = self.failed_identifiers.copy()
self.failed_identifiers.clear()
return failed
def get_and_clear_failure_reasons(self) -> dict[str, str]:
with self.done_task_lock:
failures = {
identifier: self.failure_reasons.get(identifier, "")
for identifier in self.failed_identifiers
}
for identifier in self.failed_identifiers:
self.failure_reasons.pop(identifier, None)
self.failed_identifiers.clear()
return failures
def set_finished_identifier(self, identifier: str) -> None:
with self.done_task_lock:
self.finished_identifiers.add(identifier)
def set_failed_identifier(self, identifier: str, error: Exception) -> None:
with self.done_task_lock:
self.failed_identifiers.add(identifier)
self.failure_reasons[identifier] = str(error)
def run(self) -> None:
while True:
request = self.request_queue.get()
try:
if request is None:
return
self.store_worker.save_tensor(
request.pool_key,
request.tensor,
with_soft_pin=request.with_soft_pin,
)
self.set_finished_identifier(request.identifier)
except Exception as e:
if request is not None:
self.set_failed_identifier(request.identifier, e)
logger.error("Error in %s: %s", self.name, e)
finally:
self.request_queue.task_done()
def close(self) -> None:
if self._closed.is_set():
return
self._closed.set()
self.request_queue.put(None)
if threading.current_thread() is not self:
self.join(timeout=THREAD_JOIN_TIMEOUT_SECONDS)
if self.is_alive():
logger.warning(
"%s did not exit within %.1f seconds",
self.name,
THREAD_JOIN_TIMEOUT_SECONDS,
)
class HiddenLookupServer:
"""Worker rank-0 admin channel for scheduler-side hidden lookups."""
def __init__(
self,
store_worker: HiddenStoreWorker,
vllm_config: VllmConfig,
):
self.ctx = zmq.Context() # type: ignore[attr-defined]
socket_path = get_zmq_rpc_path_hidden_lookup(vllm_config)
self._ipc_path = socket_path.removeprefix("ipc://")
if os.path.exists(self._ipc_path):
os.unlink(self._ipc_path)
self.socket = make_zmq_socket(
self.ctx,
socket_path,
zmq.REP, # type: ignore[attr-defined]
bind=True,
)
self.store_worker = store_worker
self.running = True
def process_request():
while self.running:
try:
all_frames = self.socket.recv_multipart(copy=False)
except zmq.error.ZMQError:
if not self.running:
return
logger.exception("HiddenLookupServer recv failed")
continue
msg_type = bytes(all_frames[0])
if msg_type == LOOKUP_MSG:
try:
identifier = bytes(all_frames[1]).decode("utf-8")
exists = self.store_worker.lookup(identifier)
if not exists:
self.socket.send_multipart([RESP_MISS])
else:
self.socket.send_multipart([RESP_HIT])
except Exception:
logger.exception("HiddenLookupServer lookup failed")
self.socket.send_multipart([RESP_ERR])
elif msg_type == BATCH_LOOKUP_MSG:
try:
identifiers = [
bytes(frame).decode("utf-8") for frame in all_frames[1:]
]
exists = self.store_worker.lookup_batch(identifiers)
frames = [
RESP_HIT if exists.get(identifier, False) else RESP_MISS
for identifier in identifiers
]
self.socket.send_multipart([RESP_BATCH, *frames])
except Exception:
logger.exception("HiddenLookupServer batch lookup failed")
self.socket.send_multipart([RESP_ERR])
else:
logger.warning(
"HiddenLookupServer received unknown msg_type: %r",
msg_type,
)
self.socket.send_multipart([RESP_ERR])
self.thread = threading.Thread(target=process_request, daemon=True)
self.thread.start()
def close(self):
self.running = False
self.socket.close(linger=0)
self.thread.join(timeout=THREAD_JOIN_TIMEOUT_SECONDS)
if self.thread.is_alive():
logger.warning(
"HiddenLookupServer thread did not exit within %.1f seconds",
THREAD_JOIN_TIMEOUT_SECONDS,
)
_close_zmq_context(self.ctx)
if os.path.exists(self._ipc_path):
os.unlink(self._ipc_path)
class HiddenLookupClient:
"""Scheduler-side client for worker rank-0 hidden lookup queries."""
def __init__(self, vllm_config: VllmConfig):
self.ctx = zmq.Context() # type: ignore[attr-defined]
socket_path = get_zmq_rpc_path_hidden_lookup(vllm_config)
self.socket = make_zmq_socket(
self.ctx,
socket_path,
zmq.REQ, # type: ignore[attr-defined]
bind=False,
)
self.executor = ThreadPoolExecutor(
max_workers=1,
thread_name_prefix="HiddenLookupClient",
)
self.futures: dict[str, Future[dict[str, bool]]] = {}
def lookup(self, identifier: str) -> bool:
result = self.lookup_batch([identifier], non_block=False)
assert result is not None
return result.get(identifier, False)
def _lookup_batch(self, identifiers: list[str]) -> dict[str, bool]:
self.socket.send_multipart(
[
BATCH_LOOKUP_MSG,
*(identifier.encode("utf-8") for identifier in identifiers),
]
)
resp = self.socket.recv_multipart()
msg_type = bytes(resp[0])
if msg_type == RESP_BATCH:
states = [bytes(frame) == RESP_HIT for frame in resp[1:]]
if len(states) != len(identifiers):
logger.warning(
"HiddenLookupClient received malformed batch response: "
"identifiers=%d states=%d",
len(identifiers),
len(states),
)
return {identifier: False for identifier in identifiers}
return dict(zip(identifiers, states, strict=True))
if msg_type == RESP_ERR:
return {identifier: False for identifier in identifiers}
logger.warning("HiddenLookupClient received unknown response: %r", msg_type)
return {identifier: False for identifier in identifiers}
def lookup_batch(
self,
identifiers: list[str],
non_block: bool = False,
) -> dict[str, bool] | None:
identifiers = list(dict.fromkeys(identifiers))
if not identifiers:
return {}
new_identifiers = [
identifier for identifier in identifiers if identifier not in self.futures
]
if new_identifiers:
future = self.executor.submit(self._lookup_batch, new_identifiers)
for identifier in new_identifiers:
self.futures[identifier] = future
if non_block and any(
not self.futures[identifier].done() for identifier in identifiers
):
return None
results: dict[str, bool] = {}
for identifier in identifiers:
future = self.futures[identifier]
try:
batch_results = future.result()
results[identifier] = batch_results.get(identifier, False)
except Exception as e:
logger.error("Async hidden lookup failed for %s: %s", identifier, e)
results[identifier] = False
finally:
self.futures.pop(identifier, None)
return results
def discard(self, identifier: str) -> None:
future = self.futures.pop(identifier, None)
if future is None:
return
if not any(existing is future for existing in self.futures.values()):
future.cancel()
def close(self):
self.executor.shutdown(wait=False, cancel_futures=True)
self.futures.clear()
self.socket.close(linger=0)
_close_zmq_context(self.ctx)
def get_zmq_rpc_path_hidden_lookup(vllm_config: VllmConfig) -> str:
"""Construct IPC path for Hidden Store lookup socket."""
assert vllm_config.ec_transfer_config is not None
dp_rank = get_mooncake_dp_engine_index(vllm_config.parallel_config)
base_url = envs.VLLM_RPC_BASE_PATH
hostname = socket.gethostname()
extra_config = vllm_config.ec_transfer_config.ec_connector_extra_config
rpc_port = extra_config.get(
"hidden_lookup_rpc_port",
extra_config.get("lookup_rpc_port", 0),
)
logger.debug("Hidden lookup Base URL: %s, RPC Port: %s", base_url, rpc_port)
return (
f"ipc://{base_url}/hidden_lookup_rpc_port_{rpc_port}_host_{hostname}"
f"_dp_rank{dp_rank}"
)
def _close_zmq_context(ctx) -> None:
try:
destroy = getattr(ctx, "destroy", None)
if destroy is not None:
destroy(linger=0)
return
term = getattr(ctx, "term", None)
if term is not None:
term()
except Exception:
logger.warning("failed to close hidden lookup ZMQ context", exc_info=True)