[CCF Archive] vLLM EPD hidden connector submission #4
|
|
@ -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 仓库和分支为准。
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Binary file not shown.
Loading…
Reference in New Issue