diff --git a/docs/features/mooncake_store_connector_usage.md b/docs/features/mooncake_store_connector_usage.md new file mode 100644 index 00000000000..3520cd6e872 --- /dev/null +++ b/docs/features/mooncake_store_connector_usage.md @@ -0,0 +1,161 @@ +# MooncakeStoreConnector Usage Guide + +MooncakeStoreConnector is a KV cache connector that uses [MooncakeDistributedStore](https://github.com/kvcache-ai/Mooncake) as a shared KV cache pool. Unlike `MooncakeConnector` which does direct point-to-point KV transfer between prefiller and decoder, MooncakeStoreConnector enables KV cache offloading to an external distributed store, supporting: + +- **CPU offloading**: Extend effective KV cache capacity by offloading to CPU memory via Mooncake's transfer engine. +- **Prefix caching across instances**: Hash-based deduplication allows multiple vLLM instances to share cached KV blocks through the store. +- **Single-node and multi-node deployment**: Works both as a standalone KV cache extension and in disaggregated prefill-decode setups. + +## Prerequisites + +### Install Mooncake + +Install mooncake through pip: + +```bash +uv pip install mooncake-transfer-engine +``` + +Refer to the [Mooncake official repository](https://github.com/kvcache-ai/Mooncake) for more installation instructions and building from source. + +### Start the Mooncake Master Server + +The Mooncake master manages metadata and coordinates the distributed store. Start it before launching vLLM: + +```bash +mooncake_master --port 50051 +``` + +Default ports: + +- RPC: 50051 + +Multiple vLLM instances can share the same master server. + +### Configure Mooncake + +Create a JSON configuration file (e.g., `mooncake_config.json`): + +```json +{ + "metadata_server": "P2PHANDSHAKE", + "master_server_address": "127.0.0.1:50051", + "global_segment_size": "80GB", + "local_buffer_size": "4GB", + "protocol": "rdma", + "device_name": "" +} +``` + +- `protocol`: Use `"rdma"` for best performance. `"tcp"` works as a fallback. +- `global_segment_size`: CPU memory contributed to the distributed pool (per GPU). +- `local_buffer_size`: Private buffer for this node's own operations (per GPU). + +Set the config path via environment variable: + +```bash +export MOONCAKE_CONFIG_PATH=/path/to/mooncake_config.json +``` + +## Usage + +### Single-Node KV Cache Offloading + +Use MooncakeStoreConnector to offload KV cache to CPU memory, extending the effective cache size: + +```bash +MOONCAKE_CONFIG_PATH=mooncake_config.json \ +vllm serve meta-llama/Llama-3.1-8B-Instruct \ + --kv-transfer-config '{"kv_connector":"MooncakeStoreConnector","kv_role":"kv_both"}' +``` + +### Disaggregated Prefill-Decode (XpYd) + +In disaggregated prefill-decode mode, use `MultiConnector` to combine `MooncakeConnector` (point-to-point KV transfer) with `MooncakeStoreConnector` (shared KV cache pool). This enables both direct P2P transfer between prefiller and decoder, and cross-instance prefix cache sharing via the distributed store. +**Prefiller Node:** + +```bash +MOONCAKE_CONFIG_PATH=mooncake_config.json \ +VLLM_MOONCAKE_BOOTSTRAP_PORT=50052 \ +vllm serve meta-llama/Llama-3.1-8B-Instruct \ + --port 8100 \ + --kv-transfer-config '{ + "kv_connector": "MultiConnector", + "kv_role": "kv_producer", + "kv_connector_extra_config": { + "connectors": [ + { + "kv_connector": "MooncakeConnector", + "kv_role": "kv_producer" + }, + { + "kv_connector": "MooncakeStoreConnector", + "kv_role": "kv_producer" + } + ] + } + }' +``` + +**Decoder Node:** + +```bash +MOONCAKE_CONFIG_PATH=mooncake_config.json \ +VLLM_MOONCAKE_BOOTSTRAP_PORT=50053 \ +vllm serve meta-llama/Llama-3.1-8B-Instruct \ + --port 8200 \ + --kv-transfer-config '{ + "kv_connector": "MultiConnector", + "kv_role": "kv_consumer", + "kv_connector_extra_config": { + "connectors": [ + { + "kv_connector": "MooncakeConnector", + "kv_role": "kv_consumer" + }, + { + "kv_connector": "MooncakeStoreConnector", + "kv_role": "kv_consumer" + } + ] + } + }' +``` + +**Proxy:** + +A disaggregation proxy is required to route requests between prefiller and decoder nodes. The proxy assigns `do_remote_prefill=True` / `do_remote_decode=True` to coordinate P2P transfer via `MooncakeConnector`. Refer to the [MooncakeConnector usage guide](mooncake_connector_usage.md) for proxy setup details. + +## Environment Variables + +| Variable | Description | Default | +| --- | --- | --- | +| `MOONCAKE_CONFIG_PATH` | Path to Mooncake JSON config file | (required) | +| `VLLM_MOONCAKE_BOOTSTRAP_PORT` | Bootstrap port for MooncakeConnector P2P transfer (disagg mode only) | 8998 | + +## KV Transfer Config + +### KV Role Options + +- **kv_producer**: For prefiller instances that store KV caches to the pool. +- **kv_consumer**: For decoder instances that load KV caches from the pool. +- **kv_both**: The instance both stores and loads KV caches. Use this for single-node CPU offloading. + +### kv_connector_extra_config + +- `load_async` (bool): Enable asynchronous loading for better compute-I/O overlap. Default: `true`. +- `enable_cross_layers_blocks` (bool): Enable cross-layer block packing for reduced store operations. Default: `false`. +- `discard_partial_chunks` (bool): Discard partial block chunks during store. Default: `true`. +- `lookup_rpc_port` (int): Custom port for the ZMQ lookup RPC socket. Default: `0`. + +## Notes + +### Cross-DP Prefix Cache Hits + +When running with data parallelism, set a fixed `PYTHONHASHSEED` so that block hashes are consistent across DP ranks: + +```bash +PYTHONHASHSEED=0 vllm serve ... +``` + +Without this, identical prompts may produce different block hashes on different DP ranks, preventing cross-instance prefix cache hits. diff --git a/tests/v1/kv_connector/unit/test_mooncake_store_connector.py b/tests/v1/kv_connector/unit/test_mooncake_store_connector.py new file mode 100644 index 00000000000..f1715c1d989 --- /dev/null +++ b/tests/v1/kv_connector/unit/test_mooncake_store_connector.py @@ -0,0 +1,258 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from unittest.mock import MagicMock, patch + +from vllm.config import set_current_vllm_config +from vllm.distributed.kv_events import BlockStored +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorRole, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import ( + connector, + worker, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501 + MooncakeStoreConnectorMetadata, +) +from vllm.v1.outputs import KVConnectorOutput + +from .utils import create_vllm_config + + +def _make_vllm_config(): + return create_vllm_config( + kv_connector="MooncakeStoreConnector", + kv_role="kv_both", + ) + + +def _make_block_stored() -> BlockStored: + return BlockStored( + block_hashes=[b"hash"], + parent_block_hash=None, + token_ids=[1, 2, 3], + block_size=16, + lora_id=None, + medium="cpu", + lora_name=None, + ) + + +def test_scheduler_role_initializes_store_scheduler_only(): + vllm_config = _make_vllm_config() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreScheduler" + ) as mock_scheduler, + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER) + + mock_scheduler.assert_called_once_with(vllm_config) + mock_worker.assert_not_called() + assert conn.connector_scheduler is mock_scheduler.return_value + assert conn.connector_worker is None + + +def test_worker_role_initializes_store_worker_on_rank0(): + vllm_config = _make_vllm_config() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreScheduler" + ) as mock_scheduler, + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + mock_scheduler.assert_not_called() + mock_worker.assert_called_once_with(vllm_config) + assert conn.connector_scheduler is None + assert conn.connector_worker is mock_worker.return_value + + +def test_worker_role_initializes_on_nonzero_rank(): + vllm_config = _make_vllm_config() + vllm_config.parallel_config.rank = 1 + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker, + ): + connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + mock_worker.assert_called_once_with(vllm_config) + + +def test_lookup_rpc_path_uses_data_parallel_index_in_dense_dp(): + vllm_config = _make_vllm_config() + vllm_config.parallel_config.data_parallel_rank = 0 + vllm_config.parallel_config.data_parallel_index = 3 + + path = worker.get_zmq_rpc_path_lookup(vllm_config) + + assert path.endswith("_dp_rank3") + + +def test_lookup_rpc_path_uses_local_rank_when_local_engines_only(): + vllm_config = _make_vllm_config() + vllm_config.parallel_config.data_parallel_index = 7 + vllm_config.parallel_config.data_parallel_rank_local = 1 + vllm_config.parallel_config.data_parallel_hybrid_lb = True + + path = worker.get_zmq_rpc_path_lookup(vllm_config) + + assert path.endswith("_dp_rank1") + + +def test_worker_methods_delegate_to_store_worker(): + vllm_config = _make_vllm_config() + kv_caches = {"layer0": MagicMock()} + metadata = MooncakeStoreConnectorMetadata(set(), set()) + finished_req_ids = {"req-1"} + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker_cls, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + worker_inst = mock_worker_cls.return_value + worker_inst.get_finished.return_value = ({"req-1"}, {"req-2"}) + conn.bind_connector_metadata(metadata) + + conn.register_kv_caches(kv_caches) + result = conn.get_finished(finished_req_ids) + + worker_inst.register_kv_caches.assert_called_once_with(kv_caches) + worker_inst.get_finished.assert_called_once_with(finished_req_ids, metadata) + assert result == ({"req-1"}, {"req-2"}) + + +def test_get_kv_connector_kv_cache_events_returns_none_when_empty(): + vllm_config = _make_vllm_config() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker_cls, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + mock_worker_cls.return_value.get_kv_events.return_value = [] + assert conn.get_kv_connector_kv_cache_events() is None + + +def test_get_kv_connector_kv_cache_events_wraps_worker_events(): + vllm_config = _make_vllm_config() + event = _make_block_stored() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker_cls, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + mock_worker_cls.return_value.get_kv_events.return_value = [event] + kv_events = conn.get_kv_connector_kv_cache_events() + + assert isinstance(kv_events, connector.MooncakeStoreKVEvents) + assert kv_events.get_number_of_workers() == 1 + assert kv_events.get_all_events() == [event] + + +def test_prefer_cross_layer_blocks_from_config(): + # Default: disabled + vllm_config = _make_vllm_config() + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreScheduler" + ), + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER) + assert conn.prefer_cross_layer_blocks is False + + # Enabled via config + vllm_config_enabled = create_vllm_config( + kv_connector="MooncakeStoreConnector", + kv_role="kv_both", + kv_connector_extra_config={"enable_cross_layers_blocks": "true"}, + ) + with ( + set_current_vllm_config(vllm_config_enabled), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreScheduler" + ), + ): + conn_enabled = connector.MooncakeStoreConnector( + vllm_config_enabled, KVConnectorRole.SCHEDULER + ) + assert conn_enabled.prefer_cross_layer_blocks is True + + +def test_register_cross_layers_kv_cache_delegates_to_worker(): + vllm_config = _make_vllm_config() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreWorker" + ) as mock_worker_cls, + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.WORKER) + + fake_tensor = MagicMock() + fake_backend = MagicMock() + conn.register_cross_layers_kv_cache(fake_tensor, fake_backend) + + worker_inst = mock_worker_cls.return_value + worker_inst.register_cross_layers_kv_caches.assert_called_once_with(fake_tensor) + + +def test_update_connector_output_and_take_events(): + vllm_config = _make_vllm_config() + event = _make_block_stored() + + with ( + set_current_vllm_config(vllm_config), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "connector.MooncakeStoreScheduler" + ), + ): + conn = connector.MooncakeStoreConnector(vllm_config, KVConnectorRole.SCHEDULER) + + kv_events = connector.MooncakeStoreKVEvents(num_workers=1) + kv_events.add_events([event]) + conn.update_connector_output(KVConnectorOutput(kv_cache_events=kv_events)) + + assert conn._kv_cache_events is kv_events + assert list(conn.take_events()) == [event] + assert conn._kv_cache_events is None diff --git a/tests/v1/kv_connector/unit/test_mooncake_store_worker.py b/tests/v1/kv_connector/unit/test_mooncake_store_worker.py new file mode 100644 index 00000000000..b808e148045 --- /dev/null +++ b/tests/v1/kv_connector/unit/test_mooncake_store_worker.py @@ -0,0 +1,300 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import threading +from unittest.mock import MagicMock, patch + +import torch + +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import ( + worker, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501 + ChunkedTokenDatabase, + KeyMetadata, + ReqMeta, +) + + +def _make_store_sending_thread( + store: MagicMock, +) -> worker.KVCacheStoreSendingThread: + token_database = ChunkedTokenDatabase( + KeyMetadata("test-model", 0, 0, 0, 0), block_size=16 + ) + token_database.set_kv_caches_base_addr([0x1000]) + token_database.set_block_len([256]) + thread = worker.KVCacheStoreSendingThread( + store=store, + token_database=token_database, + block_size=16, + tp_rank=0, + put_step=1, + kv_role="kv_producer", + ready_event=threading.Event(), + ) + thread.request_queue.task_done = MagicMock() + return thread + + +def _make_store_req(req_id: str, block_hashes: list[bytes]) -> ReqMeta: + return ReqMeta( + req_id=req_id, + token_len_chunk=32, + block_ids=[0, 1], + block_hashes=block_hashes, + can_save=True, + original_block_size=16, + ) + + +def test_store_sending_thread_skips_request_during_cpu_pressure(): + store = MagicMock() + store.batch_is_exist.side_effect = lambda keys: [0] * len(keys) + store.batch_put_from_multi_buffers.side_effect = [ + [-200, -200], + [256, 256], + [256, 256], + ] + thread = _make_store_sending_thread(store) + + thread.add_stored_request("req-a") + thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"])) + + assert thread._store_pressure_active is True + assert "req-a" in thread._skip_store_requests + assert store.batch_put_from_multi_buffers.call_count == 1 + + thread.add_stored_request("req-a") + thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"])) + + assert store.batch_put_from_multi_buffers.call_count == 1 + + thread.add_stored_request("req-b") + thread._handle_request(_make_store_req("req-b", [b"b0", b"b1"])) + + assert thread._store_pressure_active is False + assert "req-a" not in thread._skip_store_requests + assert store.batch_put_from_multi_buffers.call_count == 2 + + thread.add_stored_request("req-a") + thread._handle_request(_make_store_req("req-a", [b"a4", b"a5"])) + + assert store.batch_put_from_multi_buffers.call_count == 3 + + +def test_store_sending_thread_only_skips_on_no_available_handle(): + store = MagicMock() + store.batch_is_exist.side_effect = lambda keys: [0] * len(keys) + store.batch_put_from_multi_buffers.side_effect = [ + [-500, -500], + [256, 256], + ] + thread = _make_store_sending_thread(store) + + thread.add_stored_request("req-a") + thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"])) + + assert thread._store_pressure_active is False + assert "req-a" not in thread._skip_store_requests + assert store.batch_put_from_multi_buffers.call_count == 1 + + thread.add_stored_request("req-a") + thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"])) + + assert store.batch_put_from_multi_buffers.call_count == 2 + + +# --------------------------------------------------------------------------- +# Helpers for register_kv_caches tests +# --------------------------------------------------------------------------- + + +def _auto_set_ready_event(*args, **kwargs): + """Side effect for mocked thread constructors that auto-sets ready_event.""" + for arg in args: + if isinstance(arg, threading.Event): + arg.set() + for val in kwargs.values(): + if isinstance(val, threading.Event): + val.set() + return MagicMock() + + +def _make_bare_worker( + *, + num_gpu_blocks: int = 10, + block_size: int = 16, + kv_role: str = "kv_both", +) -> worker.MooncakeStoreWorker: + """Construct a MooncakeStoreWorker via __new__, bypassing __init__. + + Sets only the attributes that register_kv_caches() reads so we can + test the stride-based layout detection without a real + MooncakeDistributedStore. + """ + w = object.__new__(worker.MooncakeStoreWorker) + w.cache_config = MagicMock() + w.cache_config.num_gpu_blocks = num_gpu_blocks + w.store = MagicMock() + w.store.register_buffer.return_value = 0 + w.use_mla = False + w.token_database = ChunkedTokenDatabase( + KeyMetadata("test-model", 0, 0, 0, 0), block_size=block_size + ) + w.kv_role = kv_role + w.block_size = block_size + w.tp_rank = 0 + w.put_step = 1 + w.enable_kv_events = False + w.kv_send_thread = None + w.kv_recv_thread = None + return w + + +# --------------------------------------------------------------------------- +# register_kv_caches tests +# --------------------------------------------------------------------------- + + +def test_register_kv_caches_blocks_first_single_segment(): + """Blocks-first layout (FlashInfer/MLA): one segment per layer.""" + num_blocks = 10 + page_size_elements = 64 # elements per block + w = _make_bare_worker(num_gpu_blocks=num_blocks) + + # Shape: (num_blocks, page_size_elements) — blocks outermost, no outer_dims + tensor = torch.zeros(num_blocks, page_size_elements, dtype=torch.float16) + + with ( + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreSendingThread", + side_effect=_auto_set_ready_event, + ), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreRecvingThread", + side_effect=_auto_set_ready_event, + ), + ): + w.register_kv_caches({"layer0": tensor}) + + assert len(w.kv_caches_base_addr) == 1 + assert w.kv_caches_base_addr[0] == tensor.untyped_storage().data_ptr() + + expected_block_len = tensor.untyped_storage().nbytes() // num_blocks + assert len(w.block_len) == 1 + assert w.block_len[0] == expected_block_len + + w.store.register_buffer.assert_called_once_with( + tensor.untyped_storage().data_ptr(), + tensor.untyped_storage().nbytes(), + ) + + +def test_register_kv_caches_kv_first_two_segments(): + """K/V-first layout (FlashAttn): two segments (K, V) per layer.""" + num_blocks = 10 + block_size_tokens = 16 + num_kv_heads = 4 + head_size = 8 + + w = _make_bare_worker(num_gpu_blocks=num_blocks) + + # Shape: (2, num_blocks, block_size, num_kv_heads, head_size) — K/V outermost + tensor = torch.zeros( + 2, + num_blocks, + block_size_tokens, + num_kv_heads, + head_size, + dtype=torch.float16, + ) + + with ( + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreSendingThread", + side_effect=_auto_set_ready_event, + ), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreRecvingThread", + side_effect=_auto_set_ready_event, + ), + ): + w.register_kv_caches({"layer0": tensor}) + + # K/V-first: dim 0 has stride > page_size, so 2 segments + assert len(w.kv_caches_base_addr) == 2 + assert len(w.block_len) == 2 + + el = tensor.element_size() + seg_stride = tensor.stride(0) * el # stride of the K/V dim in bytes + base = tensor.untyped_storage().data_ptr() + assert w.kv_caches_base_addr[0] == base + assert w.kv_caches_base_addr[1] == base + seg_stride + assert w.block_len[0] == seg_stride // num_blocks + assert w.block_len[1] == seg_stride // num_blocks + + +def test_register_kv_caches_cross_layer_single_segment(): + """Cross-layer tensor: single segment with block_len = page_size * num_layers.""" + num_blocks = 10 + num_layers = 4 + per_layer_page_elements = 64 # elements per layer per block + + w = _make_bare_worker(num_gpu_blocks=num_blocks) + + # Cross-layer blocks-first tensor: all layers packed into a single + # contiguous block. Shape (num_blocks, num_layers * per_layer_page) + # mimics the physical layout after stride reordering. + total_page_elements = num_layers * per_layer_page_elements + tensor = torch.zeros(num_blocks, total_page_elements, dtype=torch.float16) + + with ( + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreSendingThread", + side_effect=_auto_set_ready_event, + ), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreRecvingThread", + side_effect=_auto_set_ready_event, + ), + ): + # Use the cross-layer wrapper key, same as register_cross_layers_kv_caches + w.register_kv_caches({"__cross_layer__": tensor}) + + assert len(w.kv_caches_base_addr) == 1 + assert w.kv_caches_base_addr[0] == tensor.untyped_storage().data_ptr() + + expected_block_len = tensor.untyped_storage().nbytes() // num_blocks + # block_len should be per_layer_page_size * num_layers + assert ( + expected_block_len + == num_layers * per_layer_page_elements * tensor.element_size() + ) + assert len(w.block_len) == 1 + assert w.block_len[0] == expected_block_len + + # Also verify via register_cross_layers_kv_caches wrapper + w2 = _make_bare_worker(num_gpu_blocks=num_blocks) + with ( + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreSendingThread", + side_effect=_auto_set_ready_event, + ), + patch( + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store." + "worker.KVCacheStoreRecvingThread", + side_effect=_auto_set_ready_event, + ), + ): + w2.register_cross_layers_kv_caches(tensor) + + assert w2.kv_caches_base_addr == w.kv_caches_base_addr + assert w2.block_len == w.block_len diff --git a/vllm/distributed/kv_transfer/kv_connector/factory.py b/vllm/distributed/kv_transfer/kv_connector/factory.py index b1ecf3c56ad..30df07c0918 100644 --- a/vllm/distributed/kv_transfer/kv_connector/factory.py +++ b/vllm/distributed/kv_transfer/kv_connector/factory.py @@ -197,6 +197,11 @@ KVConnectorFactory.register_connector( "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector", "MooncakeConnector", ) +KVConnectorFactory.register_connector( + "MooncakeStoreConnector", + "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.connector", + "MooncakeStoreConnector", +) KVConnectorFactory.register_connector( "FlexKVConnectorV1", "vllm.distributed.kv_transfer.kv_connector.v1.flexkv_connector", diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_utils.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_utils.py index 2d158387f06..a3c9200e826 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_utils.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_utils.py @@ -8,6 +8,7 @@ import uvicorn from fastapi import FastAPI, HTTPException from pydantic import BaseModel +from vllm.config import ParallelConfig from vllm.distributed.kv_transfer.kv_connector.utils import EngineId from vllm.logger import init_logger @@ -16,6 +17,15 @@ WorkerAddr = str logger = init_logger(__name__) +def get_mooncake_dp_engine_index(parallel_config: ParallelConfig) -> int: + """Return the per-engine DP index used for Mooncake side channels.""" + if parallel_config.local_engines_only: + assert parallel_config.data_parallel_rank_local is not None + return parallel_config.data_parallel_rank_local + + return parallel_config.data_parallel_index + + class RegisterWorkerPayload(BaseModel): engine_id: EngineId dp_rank: int diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/__init__.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/__init__.py new file mode 100644 index 00000000000..208f01a7cb5 --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/connector.py new file mode 100644 index 00000000000..184501a96a5 --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/connector.py @@ -0,0 +1,229 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# Adapted from vllm-project/vllm-ascend +# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/). +"""MooncakeStoreConnector - KV cache connector using MooncakeDistributedStore. + +Unlike MooncakeConnector which does direct P2P transfer, this connector +uses MooncakeDistributedStore as a shared KV cache pool. Both producer +and consumer instances read/write KV to/from the store independently, +enabling prefix caching via hash-based deduplication. +""" + +from collections.abc import Iterable +from typing import Any + +import torch + +from vllm.config import VllmConfig +from vllm.distributed.kv_events import ( + KVCacheEvent, + KVConnectorKVEvents, + KVEventAggregator, +) +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorBase_V1, + KVConnectorMetadata, + KVConnectorRole, +) +from vllm.forward_context import ForwardContext +from vllm.logger import init_logger +from vllm.v1.attention.backend import AttentionMetadata +from vllm.v1.core.kv_cache_manager import KVCacheBlocks +from vllm.v1.core.sched.output import SchedulerOutput +from vllm.v1.kv_cache_interface import KVCacheConfig +from vllm.v1.outputs import KVConnectorOutput +from vllm.v1.request import Request + +from .data import MooncakeStoreConnectorMetadata +from .scheduler import MooncakeStoreScheduler +from .worker import MooncakeStoreWorker + +logger = init_logger(__name__) + + +class MooncakeStoreKVEvents(KVConnectorKVEvents): + """KV event aggregation for MooncakeStoreConnector.""" + + def __init__(self, num_workers: int) -> None: + self._aggregator = KVEventAggregator(num_workers) + + def add_events(self, events: list[KVCacheEvent]) -> None: + self._aggregator.add_events(events) + + def aggregate(self) -> "MooncakeStoreKVEvents": + common_events = self._aggregator.get_common_events() + self._aggregator.clear_events() + self._aggregator.add_events(common_events) + self._aggregator.reset_workers() + return self + + def increment_workers(self, count: int = 1) -> None: + self._aggregator.increment_workers(count) + + def get_all_events(self) -> list[KVCacheEvent]: + return self._aggregator.get_all_events() + + def get_number_of_workers(self) -> int: + return self._aggregator.get_number_of_workers() + + def clear_events(self) -> None: + self._aggregator.clear_events() + self._aggregator.reset_workers() + + def __repr__(self) -> str: + return f"" + + +class MooncakeStoreConnector(KVConnectorBase_V1): + """KV connector using MooncakeDistributedStore as shared KV pool.""" + + @property + def prefer_cross_layer_blocks(self) -> bool: + extra_config = self._kv_transfer_config.kv_connector_extra_config + return ( + str(extra_config.get("enable_cross_layers_blocks", "False")).lower() + == "true" + ) + + def __init__( + self, + vllm_config: VllmConfig, + role: KVConnectorRole, + kv_cache_config: KVCacheConfig | None = None, + ): + super().__init__( + vllm_config=vllm_config, + role=role, + kv_cache_config=kv_cache_config, # type: ignore[arg-type] + ) + assert vllm_config.kv_transfer_config is not None + self.kv_role = vllm_config.kv_transfer_config.kv_role + self._kv_cache_events: MooncakeStoreKVEvents | None = None + + self.connector_scheduler: MooncakeStoreScheduler | None = None + self.connector_worker: MooncakeStoreWorker | None = None + + if role == KVConnectorRole.SCHEDULER: + self.connector_scheduler = MooncakeStoreScheduler(vllm_config) + else: + self.connector_worker = MooncakeStoreWorker(vllm_config) + + # ============================================================ + # Scheduler-side methods + # ============================================================ + + def get_num_new_matched_tokens( + self, + request: Request, + num_computed_tokens: int, + ) -> tuple[int, bool]: + assert self.connector_scheduler is not None + return self.connector_scheduler.get_num_new_matched_tokens( + request, num_computed_tokens + ) + + def update_state_after_alloc( + self, + request: Request, + blocks: KVCacheBlocks, + num_external_tokens: int, + ): + assert self.connector_scheduler is not None + return self.connector_scheduler.update_state_after_alloc( + request, blocks, num_external_tokens + ) + + def build_connector_meta( + self, + scheduler_output: SchedulerOutput, + ) -> KVConnectorMetadata: + assert self.connector_scheduler is not None + return self.connector_scheduler.build_connector_meta(scheduler_output) + + def request_finished( + self, + request: Request, + block_ids: list[int], + ) -> tuple[bool, dict[str, Any] | None]: + assert self.connector_scheduler is not None + return self.connector_scheduler.request_finished(request, block_ids) + + def update_connector_output(self, connector_output: KVConnectorOutput): + kv_cache_events = connector_output.kv_cache_events + if not kv_cache_events or not isinstance( + kv_cache_events, MooncakeStoreKVEvents + ): + return + + if self._kv_cache_events is None: + self._kv_cache_events = kv_cache_events + else: + self._kv_cache_events.add_events(kv_cache_events.get_all_events()) + self._kv_cache_events.increment_workers( + kv_cache_events.get_number_of_workers() + ) + + def take_events(self) -> Iterable[KVCacheEvent]: + if self._kv_cache_events is not None: + self._kv_cache_events.aggregate() + yield from self._kv_cache_events.get_all_events() + self._kv_cache_events.clear_events() + self._kv_cache_events = None + + # ============================================================ + # Worker-side methods + # ============================================================ + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): + assert self.connector_worker is not None + self.connector_worker.register_kv_caches(kv_caches) + + def register_cross_layers_kv_cache( + self, kv_cache: torch.Tensor, attn_backend: type + ): + assert self.connector_worker is not None + self.connector_worker.register_cross_layers_kv_caches(kv_cache) + + def start_load_kv(self, forward_context: ForwardContext, **kwargs: Any) -> None: + # No-op: loads are issued in get_finished() for compute overlap. + pass + + def wait_for_layer_load(self, layer_name: str) -> None: + # No layerwise support - no-op + return + + def save_kv_layer( + self, + layer_name: str, + kv_layer: torch.Tensor, + attn_metadata: AttentionMetadata, + **kwargs: Any, + ) -> None: + # No layerwise support - no-op + return + + def wait_for_save(self): + # No-op: stores are issued in get_finished() for compute overlap. + pass + + def get_finished( + self, finished_req_ids: set[str] + ) -> tuple[set[str] | None, set[str] | None]: + assert self.connector_worker is not None + metadata = self._get_connector_metadata() + assert isinstance(metadata, MooncakeStoreConnectorMetadata) + return self.connector_worker.get_finished(finished_req_ids, metadata) + + def get_kv_connector_kv_cache_events( + self, + ) -> MooncakeStoreKVEvents | None: + assert self.connector_worker is not None + events = self.connector_worker.get_kv_events() + if not events: + return None + + kv_events = MooncakeStoreKVEvents(num_workers=1) + kv_events.add_events(events) + return kv_events diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/data.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/data.py new file mode 100644 index 00000000000..acdeedceeaa --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/data.py @@ -0,0 +1,276 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# Adapted from vllm-project/vllm-ascend +# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/). +"""Data classes for MooncakeStoreConnector.""" + +from collections.abc import Iterable +from dataclasses import dataclass + +import torch + +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorMetadata, +) +from vllm.logger import init_logger +from vllm.utils.math_utils import cdiv +from vllm.v1.core.kv_cache_utils import BlockHash + +logger = init_logger(__name__) + + +@dataclass +class KeyMetadata: + """Metadata for constructing pool keys.""" + + model_name: str + tp_rank: int + pcp_rank: int + dcp_rank: int + pp_rank: int + + +@dataclass(order=True) +class PoolKey: + """Key for addressing KV cache blocks in the distributed store.""" + + key_metadata: KeyMetadata + chunk_hash: str + + def __hash__(self): + return hash( + ( + self.key_metadata.model_name, + self.key_metadata.tp_rank, + self.key_metadata.pcp_rank, + self.key_metadata.dcp_rank, + self.key_metadata.pp_rank, + self.chunk_hash, + ) + ) + + def to_string(self) -> str: + return ( + f"{self.key_metadata.model_name}" + f"@tp_rank:{self.key_metadata.tp_rank}" + f"@pcp{self.key_metadata.pcp_rank}" + f"@dcp{self.key_metadata.dcp_rank}" + f"@pp_rank:{self.key_metadata.pp_rank}" + f"@{self.chunk_hash}" + ) + + +class ChunkedTokenDatabase: + """Maps token positions to store keys and GPU memory addresses.""" + + def __init__(self, metadata: KeyMetadata, block_size: int): + self.metadata = metadata + self.block_size = block_size + self.kv_caches_base_addr: list[int] = [] + self.block_len: list[int] = [] + + def _make_key_by_hash(self, chunk_hash: str) -> PoolKey: + return PoolKey(self.metadata, chunk_hash) + + def set_kv_caches_base_addr(self, kv_caches_base_addr: list[int]): + self.kv_caches_base_addr = kv_caches_base_addr + + def set_block_len(self, block_len: list[int]): + for length in block_len: + if length % self.block_size != 0: + raise ValueError(f"block_len {length} % {self.block_size} != 0") + self.block_len = block_len + + def prepare_value( + self, start: int, end: int, block_ids: list[int] + ) -> tuple[list[int], list[int], int]: + """Compute memory addresses and sizes for a token range. + + Returns: + (addr_list, size_list, block_id) + """ + addr_list = [] + size_list = [] + block_id = block_ids[start // self.block_size] + length = len(self.block_len) + for index, base_addr in enumerate(self.kv_caches_base_addr): + addr = base_addr + block_id * self.block_len[index % length] + size = self.block_len[index % length] // self.block_size * (end - start) + addr_list.append(addr) + size_list.append(size) + return addr_list, size_list, block_id + + def process_tokens( + self, + token_len: int, + block_hashes: list[BlockHash] | list[str], + mask_num: int = 0, + ) -> Iterable[tuple[int, int, PoolKey]]: + """Process tokens and yield (start_idx, end_idx, pool_key) tuples. + + Args: + token_len: Total number of tokens. + block_hashes: Block hashes for each block. + mask_num: Number of tokens to skip from the beginning. + """ + if not block_hashes: + return + if not isinstance(block_hashes[0], str): + block_hashes = [ + h.hex() # type: ignore[union-attr] + for h in block_hashes + ] + for chunk_id, hash_val in enumerate(block_hashes): + start_idx = chunk_id * self.block_size + if start_idx >= token_len: + break + end_idx = min(start_idx + self.block_size, token_len) + if start_idx < mask_num: + continue + else: + yield ( + start_idx, + end_idx, + self._make_key_by_hash( + hash_val # type: ignore[arg-type] + ), + ) + + +@dataclass +class LoadSpec: + """Specification for loading KV cache from external store.""" + + vllm_cached_tokens: int + kvpool_cached_tokens: int + can_load: bool + token_len: int = 0 + + +@dataclass +class RequestTracker: + """Tracks per-request state across scheduler ticks.""" + + req_id: str + token_len: int + allocated_block_ids: list[int] + num_saved_tokens: int = 0 + token_ids: list[int] | None = None + # Snapshot of the prefill range length at tracker creation time. + # For a fresh request this is len(prompt). For a resumed-from-preemption + # request it includes previously-generated tokens, which are re-prefilled. + prefill_end_tokens: int = 0 + + def update( + self, + new_block_ids: tuple[list[int], ...] | list[int], + ) -> None: + if len(new_block_ids) == 0: + new_block_ids = [] + elif isinstance(new_block_ids, tuple): + new_block_ids = new_block_ids[0] + elif isinstance(new_block_ids, list): + pass + else: + raise ValueError(f"Unsupported new_block_ids type {type(new_block_ids)}") + self.allocated_block_ids.extend(new_block_ids) + + +@dataclass +class ReqMeta: + """Per-request metadata for store put/get operations.""" + + req_id: str + token_len_chunk: int + block_ids: list[int] + block_hashes: list[BlockHash] + + can_save: bool | None = None + load_spec: LoadSpec | None = None + is_last_chunk: bool | None = None + current_event: torch.cuda.Event | None = None + + token_ids: list[int] | None = None + original_block_size: int | None = None + + @staticmethod + def from_request_tracker( + tracker: RequestTracker, + block_size: int, + load_spec: LoadSpec | None = None, + skip_save: bool | None = False, + block_hashes: list[BlockHash] | None = None, + is_last_chunk: bool | None = None, + discard_partial_chunks: bool = True, + original_block_size: int | None = None, + ) -> "ReqMeta | None": + """Create ReqMeta from a RequestTracker.""" + if block_hashes is None: + block_hashes = [] + input_token_len = tracker.token_len + + chunk_boundary = ( + cdiv(tracker.num_saved_tokens + 1, block_size) * block_size + if discard_partial_chunks + else 0 + ) + num_tokens_to_save = ( + (input_token_len // block_size * block_size) + if discard_partial_chunks + else input_token_len + ) + + skip_save = skip_save or num_tokens_to_save < chunk_boundary + if skip_save and load_spec is None: + return None + + if not skip_save: + tracker.num_saved_tokens = num_tokens_to_save + + token_ids = None + if tracker.token_ids: + token_ids = tracker.token_ids + + if load_spec is not None and load_spec.can_load: + logger.debug( + "Scheduled to load %d tokens for request %s", + load_spec.kvpool_cached_tokens, + tracker.req_id, + ) + else: + load_spec = None + + logger.debug( + "request:%s, meta save spec:%s, meta load spec:%s", + tracker.req_id, + not skip_save, + load_spec, + ) + return ReqMeta( + req_id=tracker.req_id, + token_len_chunk=num_tokens_to_save, + block_ids=tracker.allocated_block_ids, + can_save=not skip_save, + load_spec=load_spec, + block_hashes=block_hashes, + is_last_chunk=is_last_chunk, + token_ids=token_ids, + original_block_size=original_block_size, + ) + + +class MooncakeStoreConnectorMetadata(KVConnectorMetadata): + """Metadata passed from scheduler to worker.""" + + def __init__( + self, + unfinished_request_ids: set[str], + preempted_req_ids: set[str], + ): + self.requests: list[ReqMeta] = [] + self.unfinished_request_ids = unfinished_request_ids + self.preempted_req_ids = preempted_req_ids + + def add_request(self, req_meta: ReqMeta) -> None: + self.requests.append(req_meta) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py new file mode 100644 index 00000000000..5ce3278fee8 --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py @@ -0,0 +1,380 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# Adapted from vllm-project/vllm-ascend +# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/). +"""Scheduler-side logic for MooncakeStoreConnector.""" + +from typing import Any + +from vllm.config import VllmConfig +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorMetadata, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501 + LoadSpec, + MooncakeStoreConnectorMetadata, + ReqMeta, + RequestTracker, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.worker import ( # noqa: E501 + LookupKeyClient, +) +from vllm.logger import init_logger +from vllm.v1.core.kv_cache_manager import KVCacheBlocks +from vllm.v1.core.sched.output import NewRequestData, SchedulerOutput +from vllm.v1.request import Request + +logger = init_logger(__name__) + + +def _new_req_prefill_tokens(request: NewRequestData) -> list[int]: + """Tokens this prefill will compute KV for. + + Under the v2 model runner, resumed-from-preemption requests appear in + ``scheduled_new_reqs`` with ``prefill_token_ids`` set to the request's full + token list (prompt + previously-generated). For all other cases this falls + back to the original prompt. + """ + if request.prefill_token_ids is not None: + return request.prefill_token_ids + assert request.prompt_token_ids is not None + return request.prompt_token_ids + + +class MooncakeStoreScheduler: + """Scheduler-side component for MooncakeStoreConnector.""" + + def __init__(self, vllm_config: VllmConfig): + assert vllm_config.kv_transfer_config is not None + self.kv_role = vllm_config.kv_transfer_config.kv_role + self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get( + "load_async", True + ) + self.client = LookupKeyClient(vllm_config) + + self.pcp_size = vllm_config.parallel_config.prefill_context_parallel_size + self.dcp_size = vllm_config.parallel_config.decode_context_parallel_size + self.original_block_size = vllm_config.cache_config.block_size + self._block_size = vllm_config.cache_config.block_size + if self.pcp_size > 1: + self._block_size *= self.pcp_size + if self.dcp_size > 1: + self._block_size *= self.dcp_size + + self._discard_partial_chunks = ( + vllm_config.kv_transfer_config.get_from_extra_config( + "discard_partial_chunks", True + ) + ) + + # Per-request state + self.load_specs: dict[str, LoadSpec] = {} # to be loaded + self._request_trackers: dict[str, RequestTracker] = {} # scheduled new requests + self._preempted_req_ids: set[str] = set() # preempted requests + self._unfinished_requests: dict[str, tuple[Request, list[int]]] = {} + self._unfinished_request_ids: set[str] = set() + + def get_num_new_matched_tokens( + self, + request: Request, + num_computed_tokens: int, + ) -> tuple[int, bool]: + """Check for external KV cache hit.""" + # Look up against the full prefill range, not just the prompt. + if self._discard_partial_chunks: + token_len = request.num_tokens // self._block_size * self._block_size + else: + token_len = request.num_tokens + + if token_len < self._block_size: + return 0, False + + num_external_hit_tokens = self.client.lookup(token_len, request.block_hashes) + + if num_external_hit_tokens == request.num_tokens: + num_external_hit_tokens -= 1 + + if num_external_hit_tokens < num_computed_tokens: + need_to_allocate = 0 + else: + need_to_allocate = num_external_hit_tokens - num_computed_tokens + + logger.debug( + "Reqid: %s, Total tokens %d, kvpool hit tokens: %d, need to load: %d", + request.request_id, + request.num_tokens, + num_external_hit_tokens, + need_to_allocate, + ) + + if need_to_allocate <= 0: + return 0, False + + self.load_specs[request.request_id] = LoadSpec( + vllm_cached_tokens=num_computed_tokens, + kvpool_cached_tokens=num_external_hit_tokens, + can_load=False, + ) + + return need_to_allocate, self.load_async + + def update_state_after_alloc( + self, + request: Request, + blocks: KVCacheBlocks, + num_external_tokens: int, + ): + """Update state after block allocation.""" + local_block_ids: list[int] = [] + if num_external_tokens > 0: + local_block_ids = blocks.get_block_ids()[0] + + self._unfinished_requests[request.request_id] = (request, local_block_ids) + self._unfinished_request_ids.add(request.request_id) + + if request.request_id not in self.load_specs: + return + + if num_external_tokens == 0: + self.load_specs[request.request_id].can_load = False + return + + assert ( + num_external_tokens > 0 + and num_external_tokens + == self.load_specs[request.request_id].kvpool_cached_tokens + - self.load_specs[request.request_id].vllm_cached_tokens + ), ( + f"Mismatch in number of tokens: {num_external_tokens} vs " + f"{self.load_specs[request.request_id].kvpool_cached_tokens} - " + f"{self.load_specs[request.request_id].vllm_cached_tokens}" + f" for request {request.request_id}" + ) + + self.load_specs[request.request_id].can_load = True + + def build_connector_meta( + self, scheduler_output: SchedulerOutput + ) -> KVConnectorMetadata: + """Build connector metadata for this scheduler step.""" + force_skip_save = self.kv_role == "kv_consumer" + + for finished_req_id in scheduler_output.finished_req_ids: + self.load_specs.pop(finished_req_id, None) + self._request_trackers.pop(finished_req_id, None) + self._unfinished_requests.pop(finished_req_id, None) + self._unfinished_request_ids.discard(finished_req_id) + self._preempted_req_ids.discard(finished_req_id) + + preempted_ids = scheduler_output.preempted_req_ids or set() + self._preempted_req_ids.update(preempted_ids) + for req_id in preempted_ids: + self._request_trackers.pop(req_id, None) + self._unfinished_requests.pop(req_id, None) + + meta = MooncakeStoreConnectorMetadata( + self._unfinished_request_ids, + preempted_ids, + ) + + # Handle new requests + for request in scheduler_output.scheduled_new_reqs: + load_spec = self.load_specs.pop(request.req_id, None) + num_tokens_to_compute = ( + request.num_computed_tokens + + scheduler_output.num_scheduled_tokens[request.req_id] + ) + assert request.req_id in self._unfinished_requests + request_tuple = self._unfinished_requests.get(request.req_id) + request_real = request_tuple[0] # type: ignore[index] + + if not isinstance(request.block_ids[0], list): + unfolded_block_ids = request.block_ids.copy() + else: + # TODO: support HMA + unfolded_block_ids = request.block_ids[0].copy() + + prefill_tokens = _new_req_prefill_tokens(request) + request_tracker = RequestTracker( + req_id=request.req_id, + token_len=num_tokens_to_compute, + allocated_block_ids=unfolded_block_ids, + num_saved_tokens=0, + token_ids=prefill_tokens[:num_tokens_to_compute], + prefill_end_tokens=len(prefill_tokens), + ) + self._request_trackers[request.req_id] = request_tracker + + last_chunk_tokens_num = ( + (len(prefill_tokens) // self._block_size * self._block_size) + if self._discard_partial_chunks + else len(prefill_tokens) + ) + + req_meta = ReqMeta.from_request_tracker( + request_tracker, + self._block_size, + load_spec=load_spec, + skip_save=force_skip_save, + block_hashes=request_real.block_hashes, + is_last_chunk=(request_tracker.token_len >= last_chunk_tokens_num), + discard_partial_chunks=self._discard_partial_chunks, + original_block_size=self.original_block_size, + ) + if req_meta is not None: + meta.add_request(req_meta) + + # Handle cached (running, or MRV1 resumed-from-preemption) requests + cached_reqs = scheduler_output.scheduled_cached_reqs + if not force_skip_save: + for i, req_id in enumerate(cached_reqs.req_ids): + new_block_ids = cached_reqs.new_block_ids[i] + if not new_block_ids: + continue + + req_meta = None + if req_id in self._preempted_req_ids: + # Resumed after preemption + if isinstance(new_block_ids, tuple): + block_ids_list = new_block_ids[0].copy() + else: + block_ids_list = new_block_ids.copy() + self._preempted_req_ids.discard(req_id) + load_spec = self.load_specs.pop(req_id, None) + request_tuple = self._unfinished_requests.get(req_id) + request_real = request_tuple[0] # type: ignore[index] + num_tokens_to_compute = ( + request_real.num_computed_tokens + + scheduler_output.num_scheduled_tokens[req_id] + ) + # On resume, the request re-prefills prompt + previously + # generated tokens (all_token_ids). + prefill_tokens = list(request_real.all_token_ids) + request_tracker = RequestTracker( + req_id=req_id, + token_len=num_tokens_to_compute, + allocated_block_ids=block_ids_list, + num_saved_tokens=0, + token_ids=prefill_tokens[:num_tokens_to_compute].copy(), + prefill_end_tokens=len(prefill_tokens), + ) + self._request_trackers[req_id] = request_tracker + + last_chunk_tokens_num = ( + (len(prefill_tokens) // self._block_size * self._block_size) + if self._discard_partial_chunks + else len(prefill_tokens) + ) + req_meta = ReqMeta.from_request_tracker( + request_tracker, + self._block_size, + load_spec=load_spec, + skip_save=force_skip_save, + block_hashes=request_real.block_hashes, + is_last_chunk=( + request_tracker.token_len >= last_chunk_tokens_num + ), + discard_partial_chunks=self._discard_partial_chunks, + original_block_size=self.original_block_size, + ) + else: + # Decode/chunked request + request_tracker = self._request_trackers[req_id] + num_new_tokens = scheduler_output.num_scheduled_tokens[req_id] + req_tuple = self._unfinished_requests.get(req_id) + if req_tuple: + unfinished_req = req_tuple[0] + num_current_tokens = request_tracker.token_len + new_token_ids = unfinished_req.all_token_ids[ + num_current_tokens : num_current_tokens + num_new_tokens + ] + request_tracker.token_len += len(new_token_ids) + else: + raise ValueError( + f"Request {req_id} is not in _unfinished_requests" + ) + num_computed_token = cached_reqs.num_computed_tokens[i] + # Use the tracker's snapshot of the prefill range so resumed + # requests keep saving past the original prompt boundary. + prefill_end = request_tracker.prefill_end_tokens + if num_computed_token >= prefill_end: + continue + request_tracker.update(new_block_ids) + + last_chunk_tokens_num = ( + (prefill_end // self._block_size * self._block_size) + if self._discard_partial_chunks + else prefill_end + ) + req_meta = ReqMeta.from_request_tracker( + request_tracker, + self._block_size, + load_spec=None, + skip_save=force_skip_save, + block_hashes=unfinished_req.block_hashes, + is_last_chunk=( + request_tracker.token_len >= last_chunk_tokens_num + ), + discard_partial_chunks=self._discard_partial_chunks, + original_block_size=self.original_block_size, + ) + + if req_meta is not None: + meta.add_request(req_meta) + + # Handle requests with pending load specs not yet scheduled + request_ids = [req.req_id for req in scheduler_output.scheduled_new_reqs] + for request_id, ( + unfinished_req, + block_ids, + ) in self._unfinished_requests.items(): + if request_id not in request_ids and request_id not in cached_reqs.req_ids: + load_spec = self.load_specs.pop(request_id, None) + if not load_spec: + continue + num_tokens_to_compute = load_spec.kvpool_cached_tokens + if (num_tokens_to_compute % self._block_size != 0) and ( + num_tokens_to_compute == unfinished_req.num_tokens - 1 + ): + num_tokens_to_compute = num_tokens_to_compute + 1 + request_tracker = RequestTracker( + req_id=request_id, + token_len=num_tokens_to_compute, + allocated_block_ids=block_ids, + num_saved_tokens=0, + ) + self._request_trackers[request_id] = request_tracker + req_meta = ReqMeta.from_request_tracker( + request_tracker, + self._block_size, + load_spec=load_spec, + skip_save=None, + block_hashes=unfinished_req.block_hashes, + discard_partial_chunks=self._discard_partial_chunks, + ) + if req_meta is not None: + meta.add_request(req_meta) + + return meta + + def request_finished( + self, + request: Request, + block_ids: list[int], + ) -> tuple[bool, dict[str, Any] | None]: + """Determine whether to delay freeing blocks for async save.""" + if self.kv_role == "kv_consumer": + return False, None + tracker = self._request_trackers.get(request.request_id) + assert tracker is not None + if tracker.num_saved_tokens <= 0: + return False, None + delay_free_blocks = len(block_ids) > 0 + if delay_free_blocks: + logger.debug( + "Delaying free of %d blocks for request %s", + len(block_ids), + request.request_id, + ) + return delay_free_blocks, None diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py new file mode 100644 index 00000000000..487542c5917 --- /dev/null +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py @@ -0,0 +1,979 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# The transfer-thread scaffolding (KVTransferThread, KVCacheStoreSendingThread, +# KVCacheStoreRecvingThread) is adapted from vllm-project/vllm-ascend +# (vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/). +"""Worker-side logic for MooncakeStoreConnector. + +Includes the store worker, transfer threads, lookup server, +and MooncakeDistributedStore integration. +""" + +import json +import os +import queue +import threading +from collections import defaultdict +from dataclasses import dataclass +from typing import Any + +import regex as re +import torch +import zmq + +import vllm.envs as envs +from vllm.config import VllmConfig +from vllm.distributed import ( + get_dcp_group, + get_pcp_group, + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, +) +from vllm.distributed.kv_events import BlockStored +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import ( + get_mooncake_dp_engine_index, +) +from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import ( # noqa: E501 + ChunkedTokenDatabase, + KeyMetadata, + MooncakeStoreConnectorMetadata, + ReqMeta, +) +from vllm.logger import init_logger +from vllm.utils.network_utils import get_ip, make_zmq_socket +from vllm.v1.core.kv_cache_utils import BlockHash, maybe_convert_block_hash +from vllm.v1.serial_utils import MsgpackDecoder, MsgpackEncoder + +logger = init_logger(__name__) + +DEFAULT_GLOBAL_SEGMENT_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB +DEFAULT_LOCAL_BUFFER_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB +MOONCAKE_NO_AVAILABLE_HANDLE = -200 + + +@dataclass +class MooncakeStoreConfig: + """Configuration for MooncakeDistributedStore.""" + + metadata_server: str + global_segment_size: int + local_buffer_size: int + protocol: str + device_name: str + master_server_address: str + + @staticmethod + def from_file(file_path: str) -> "MooncakeStoreConfig": + with open(file_path) as file: + config = json.load(file) + return MooncakeStoreConfig( + metadata_server=config.get("metadata_server", ""), + global_segment_size=_parse_size( + config.get("global_segment_size", DEFAULT_GLOBAL_SEGMENT_SIZE) + ), + local_buffer_size=_parse_size( + config.get("local_buffer_size", DEFAULT_LOCAL_BUFFER_SIZE) + ), + protocol=config.get("protocol", "rdma"), + device_name=config.get("device_name", ""), + master_server_address=config.get("master_server_address", ""), + ) + + @staticmethod + def load_from_env() -> "MooncakeStoreConfig": + config_path = os.getenv("MOONCAKE_CONFIG_PATH") + if not config_path: + raise ValueError( + "The environment variable 'MOONCAKE_CONFIG_PATH' is not set." + ) + return MooncakeStoreConfig.from_file(config_path) + + +def _parse_size(value: Any) -> int: + """Parse storage size strings with units: GB, MB, KB, B.""" + if isinstance(value, int): + return value + if not isinstance(value, str): + try: + return int(value) + except (TypeError, ValueError) as e: + raise TypeError(f"Unsupported type for size: {type(value)}") from e + + cleaned = value.strip().lower() + if not cleaned: + raise ValueError("Size cannot be empty.") + + unit_multipliers = { + "gb": 1024**3, + "mb": 1024**2, + "kb": 1024, + "b": 1, + } + match = re.match(r"^\s*([\d.]+)\s*(gb|mb|kb|b)?\s*$", cleaned) + if not match: + raise ValueError(f"Invalid format: '{value}'") + + number_str = match.group(1) + unit = match.group(2) or "b" + multiplier = unit_multipliers[unit] + + try: + numeric_value = float(number_str) + except ValueError as exc: + raise ValueError(f"Invalid numeric value '{number_str}' in: '{value}'") from exc + return int(numeric_value * multiplier) + + +# ============================================================ +# Transfer Threads +# ============================================================ + + +class KVTransferThread(threading.Thread): + """Base class for async KV cache transfer threads.""" + + def __init__( + self, + store: Any, + token_database: ChunkedTokenDatabase, + block_size: int, + tp_rank: int, + ready_event: threading.Event, + name: str, + ): + super().__init__(daemon=True, name=name) + self.store = store + self.ready_event = ready_event + self.block_size = block_size + self.tp_rank = tp_rank + self.token_database = token_database + self.done_task_lock = threading.Lock() + self.request_queue: queue.Queue[Any] = queue.Queue() + self.finished_requests: set[str] = set() + self.kv_event_lock = threading.Lock() + self.kv_events: list[BlockStored] = [] + + def add_request(self, request: ReqMeta) -> None: + self.request_queue.put(request) + + def get_and_clear_finished_requests(self) -> set[str]: + with self.done_task_lock: + finished = self.finished_requests.copy() + self.finished_requests.clear() + return finished + + def set_finished_request(self, req_id: str): + with self.done_task_lock: + self.finished_requests.add(req_id) + + def run(self): + self.ready_event.set() + while True: + try: + request_data = self.request_queue.get() + if request_data is None: + logger.warning("Received a None request!") + self.request_queue.task_done() + continue + self._handle_request(request_data) + except Exception as e: + logger.error("Error in %s: %s", self.name, e) + + def _handle_request(self, req_meta: Any): + pass + + def update_kv_event(self, events: list[BlockStored]): + with self.kv_event_lock: + self.kv_events.extend(events) + + def get_kv_events(self) -> list[BlockStored]: + with self.kv_event_lock: + events = self.kv_events.copy() + self.kv_events.clear() + return events + + +class KVCacheStoreSendingThread(KVTransferThread): + """Background thread for storing KV cache blocks to the store.""" + + def __init__( + self, + store: Any, + token_database: ChunkedTokenDatabase, + block_size: int, + tp_rank: int, + put_step: int, + kv_role: str, + ready_event: threading.Event, + enable_kv_event: bool = False, + ): + super().__init__( + store, + token_database, + block_size, + tp_rank, + ready_event, + name="KVCacheStoreSendingThread", + ) + self.put_step = put_step + self.kv_role = kv_role + self.stored_requests: defaultdict[str, int] = defaultdict(int) + self.enable_kv_event = enable_kv_event + + # Pause store requests when CPU offloading is under pressure. + self._store_pressure_active = False + self._skip_store_requests: set[str] = set() + + def add_stored_request(self, req_id: str): + with self.done_task_lock: + self.stored_requests[req_id] += 1 + + def dec_stored_request(self, req_id: str): + with self.done_task_lock: + if req_id in self.stored_requests: + self.stored_requests[req_id] -= 1 + + def delete_finished_stored_request(self, req_id: str): + with self.done_task_lock: + if req_id in self.stored_requests: + del self.stored_requests[req_id] + self._skip_store_requests.discard(req_id) + + def _should_skip_request(self, req_id: str) -> bool: + with self.done_task_lock: + return self._store_pressure_active and req_id in self._skip_store_requests + + def _mark_request_skipped_for_pressure(self, req_id: str) -> bool: + with self.done_task_lock: + already_skipped = req_id in self._skip_store_requests + self._store_pressure_active = True + self._skip_store_requests.add(req_id) + return already_skipped + + def _clear_store_pressure(self) -> bool: + with self.done_task_lock: + if not self._store_pressure_active and not self._skip_store_requests: + return False + self._store_pressure_active = False + self._skip_store_requests.clear() + return True + + def _handle_request(self, req_meta: ReqMeta): + token_len = req_meta.token_len_chunk + block_ids = req_meta.block_ids + req_id = req_meta.req_id + current_event = req_meta.current_event + + if req_id not in self.stored_requests: + self.request_queue.task_done() + return + if self._should_skip_request(req_id): + logger.debug( + "Skipping Mooncake store for request %s while CPU offloading " + "is under pressure", + req_id, + ) + self.dec_stored_request(req_id) + self.request_queue.task_done() + return + + starts = [] + ends = [] + keys = [] + block_hashes: list[BlockHash] = [] + for index, (start, end, key) in enumerate( + self.token_database.process_tokens(token_len, req_meta.block_hashes) + ): + starts.append(start) + ends.append(end) + keys.append(key.to_string()) + block_hashes.append(req_meta.block_hashes[index]) + + # Apply put_step striding for TP + starts = starts[self.tp_rank % self.put_step :: self.put_step] + ends = ends[self.tp_rank % self.put_step :: self.put_step] + keys = keys[self.tp_rank % self.put_step :: self.put_step] + block_hashes = block_hashes[self.tp_rank % self.put_step :: self.put_step] + + if not keys: + self.dec_stored_request(req_id) + return + + # Check which blocks already exist (dedup) + exists_states = self.store.batch_is_exist(keys) + missing_indices = [i for i, exists in enumerate(exists_states) if exists != 1] + + if not missing_indices: + self.dec_stored_request(req_id) + return + + starts = [starts[i] for i in missing_indices] + ends = [ends[i] for i in missing_indices] + keys = [keys[i] for i in missing_indices] + block_hashes = [block_hashes[i] for i in missing_indices] + + logger.debug( + "Storing KV cache for %d out of %d blocks " + "(missing_count=%d) for request %s", + len(keys), + token_len // self.block_size, + len(missing_indices), + req_id, + ) + + addrs = [] + sizes = [] + stored_events: list[BlockStored] = [] + prev_key = None + new_block_hashes = [maybe_convert_block_hash(bh) for bh in block_hashes] + + for index, start in enumerate(starts): + addr, size, _ = self.token_database.prepare_value( + start, ends[index], block_ids + ) + addrs.append(addr) + sizes.append(size) + + if self.enable_kv_event: + token_ids = ( + req_meta.token_ids[start : ends[index]] + if req_meta.token_ids is not None + else None + ) + stored_event = BlockStored( + block_hashes=[new_block_hashes[index]], + parent_block_hash=prev_key, + token_ids=token_ids, + block_size=req_meta.original_block_size, + lora_id=None, + medium="cpu", + lora_name=None, + ) + stored_events.append(stored_event) + prev_key = new_block_hashes[index] + + if current_event is not None: + current_event.synchronize() + + try: + res = self.store.batch_put_from_multi_buffers(keys, addrs, sizes) + failed = [i for i, v in enumerate(res) if v < 0] + if failed: + # Compute total bytes attempted for this batch + total_bytes = sum(sum(s) if isinstance(s, list) else s for s in sizes) + failed_codes = set(res[i] for i in failed) + logger.warning( + "batch_put failed: %d/%d keys failed " + "(codes=%s, batch_bytes=%d, num_keys=%d), " + "first_key=%s", + len(failed), + len(keys), + failed_codes, + total_bytes, + len(keys), + keys[0] if keys else "N/A", + ) + if ( + MOONCAKE_NO_AVAILABLE_HANDLE in failed_codes + and not self._mark_request_skipped_for_pressure(req_id) + ): + logger.warning( + "Detected Mooncake CPU offloading pressure " + "(NO_AVAILABLE_HANDLE); skipping future store " + "batches for request %s until a later store " + "batch succeeds", + req_id, + ) + elif self._clear_store_pressure(): + logger.info( + "Mooncake CPU offloading pressure cleared after a " + "successful store batch" + ) + except Exception as e: + logger.error("Failed to put key %s, error: %s", keys, e) + + if self.enable_kv_event and stored_events: + self.update_kv_event(stored_events) + + self.dec_stored_request(req_id) + self.request_queue.task_done() + + +class KVCacheStoreRecvingThread(KVTransferThread): + """Background thread for loading KV cache blocks from the store.""" + + def __init__( + self, + store: Any, + token_database: ChunkedTokenDatabase, + block_size: int, + tp_rank: int, + ready_event: threading.Event, + ): + super().__init__( + store, + token_database, + block_size, + tp_rank, + ready_event, + name="KVCacheStoreRecvingThread", + ) + + def _handle_request(self, req_meta: ReqMeta): + token_len = req_meta.load_spec.token_len # type: ignore[union-attr] + req_id = req_meta.req_id + mask_num = ( + req_meta.load_spec.vllm_cached_tokens # type: ignore[union-attr] + // self.block_size + * self.block_size + ) + + addr_list = [] + size_list = [] + key_list = [] + for start, end, key in self.token_database.process_tokens( + token_len, req_meta.block_hashes, mask_num + ): + addr, size, _ = self.token_database.prepare_value( + start, end, req_meta.block_ids + ) + key_list.append(key.to_string()) + addr_list.append(addr) + size_list.append(size) + + # Rotate lists by tp_rank for load balancing + key_list_c = ( + key_list[self.tp_rank % len(key_list) :] + + key_list[: self.tp_rank % len(key_list)] + ) + addr_list_c = ( + addr_list[self.tp_rank % len(addr_list) :] + + addr_list[: self.tp_rank % len(addr_list)] + ) + size_list_c = ( + size_list[self.tp_rank % len(size_list) :] + + size_list[: self.tp_rank % len(size_list)] + ) + + try: + res = self.store.batch_get_into_multi_buffers( + key_list_c, addr_list_c, size_list_c + ) + failed = [ + (key, value) + for key, value in zip(key_list_c, res, strict=True) + if value < 0 + ] + if failed: + logger.warning( + "Failed to get %d Mooncake keys (batch_keys=%d, first_failures=%s)", + len(failed), + len(key_list_c), + failed[:3], + ) + except Exception as e: + logger.warning( + "Failed to get Mooncake batch %s, error: %s", + key_list_c[:3], + e, + ) + + self.set_finished_request(req_id) + self.request_queue.task_done() + + +# ============================================================ +# Store Worker +# ============================================================ + + +class MooncakeStoreWorker: + """Worker-side component for MooncakeStoreConnector.""" + + def __init__(self, vllm_config: VllmConfig): + try: + from mooncake.store import MooncakeDistributedStore # type: ignore + except ImportError as e: + raise ImportError( + "Please install mooncake by following the instructions at " + "https://github.com/kvcache-ai/Mooncake/blob/main/doc/" + "en/build.md to run vLLM with MooncakeStoreConnector." + ) from e + + model_config = vllm_config.model_config + parallel_config = vllm_config.parallel_config + + self.dp_rank = get_mooncake_dp_engine_index(parallel_config) + self.tp_rank = get_tensor_model_parallel_rank() + self.tp_size = get_tensor_model_parallel_world_size() + self.pp_size = parallel_config.pipeline_parallel_size + self.pp_rank = (parallel_config.rank // self.tp_size) % self.pp_size + + self.pcp_size = get_pcp_group().world_size + self.pcp_rank = get_pcp_group().rank_in_group if self.pcp_size > 1 else 0 + self.dcp_size = get_dcp_group().world_size + self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_size > 1 else 0 + + assert vllm_config.kv_transfer_config is not None + self.kv_role = vllm_config.kv_transfer_config.kv_role + self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get( + "load_async", True + ) + self.cache_config = vllm_config.cache_config + self.original_block_size = self.cache_config.block_size + self.block_size = self.cache_config.block_size + if self.pcp_size > 1: + self.block_size *= self.pcp_size + if self.dcp_size > 1: + self.block_size *= self.dcp_size + self.num_layers = model_config.get_num_layers(parallel_config) + + self.use_mla = False + if ( + hasattr(model_config, "use_mla") + and isinstance(model_config.use_mla, bool) + and model_config.use_mla + ): + self.use_mla = True + + if self.use_mla: + self.num_kv_head = 1 + else: + self.num_kv_head = model_config.get_total_num_kv_heads() + + if self.num_kv_head < self.tp_size: + self.put_step = self.tp_size // self.num_kv_head + self.head_or_tp_rank = self.tp_rank // self.put_step + else: + self.head_or_tp_rank = self.tp_rank + self.put_step = 1 + + self.metadata = KeyMetadata( + model_name=model_config.model.rstrip("/").split("/")[-1], + tp_rank=self.head_or_tp_rank, + pcp_rank=self.pcp_rank, + dcp_rank=self.dcp_rank, + pp_rank=self.pp_rank, + ) + + self.token_database = ChunkedTokenDatabase(self.metadata, self.block_size) + + # Initialize MooncakeDistributedStore with its own TransferEngine + store_config = MooncakeStoreConfig.load_from_env() + self.store = MooncakeDistributedStore() + + local_seg = get_ip() + config_dict = { + "local_hostname": local_seg, + "metadata_server": store_config.metadata_server, + "global_segment_size": str(store_config.global_segment_size), + "local_buffer_size": str(store_config.local_buffer_size), + "protocol": store_config.protocol, + "rdma_devices": store_config.device_name, + "master_server_addr": store_config.master_server_address, + } + ret = self.store.setup(config_dict) + if ret != 0: + msg = "Initialize MooncakeDistributedStore failed." + logger.error(msg) + raise RuntimeError(msg) + + kv_event_config = vllm_config.kv_events_config + self.enable_kv_events = False + if kv_event_config and kv_event_config.enable_kv_cache_events: + self.enable_kv_events = True + + self.kv_send_thread: KVCacheStoreSendingThread | None = None + self.kv_recv_thread: KVCacheStoreRecvingThread | None = None + self.finished_store_req: set[str] = set() + + # Start lookup server on rank 0 for scheduler-side prefix queries + self.lookup_server: LookupKeyServer | None = None + if vllm_config.parallel_config.rank == 0: + self.lookup_server = LookupKeyServer(self, vllm_config) + + def register_cross_layers_kv_caches(self, kv_cache: torch.Tensor) -> None: + """Register a cross-layers KV cache tensor. + + Wraps the unified tensor in a single-entry dict so that the + existing stride-based logic in register_kv_caches() produces + the correct single-segment result (block_len = page_size * num_layers). + """ + self.register_kv_caches({"__cross_layer__": kv_cache}) + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): + """Register KV cache tensors and start transfer threads.""" + # TODO(yifan): we haven't supported HMA yet. + first_kv_cache = next(iter(kv_caches.values())) + + # num_blocks from cache_config is authoritative (set after + # profiling, before KV cache allocation). + assert self.cache_config.num_gpu_blocks is not None + self.num_blocks = self.cache_config.num_gpu_blocks + + # Detect the KV cache memory layout using the stride-based + # approach from simple_kv_offload/worker.py. + # + # The physical layout varies across attention backends: + # FlashAttn/ROCm : (2, num_blocks, ...) → K/V outermost + # FlashInfer/MLA : (num_blocks, ...) → blocks outermost + # + # We derive page_size_bytes = storage.nbytes() // num_blocks, + # then classify dims: any dim whose byte-stride exceeds + # page_size_bytes must be an outer segment dim (e.g. the K/V + # dim of size 2). For those backends we register each segment + # (K, V) as a separate base-address so that the per-block + # offset arithmetic in prepare_value() stays correct. + storage = first_kv_cache.untyped_storage() + el = first_kv_cache.element_size() + page_size_bytes = storage.nbytes() // self.num_blocks + outer_dims = [ + d + for d in range(first_kv_cache.ndim) + if first_kv_cache.stride(d) * el > page_size_bytes + ] + + # Register buffers with the store (deduplicate shared storages) + # and record per-segment base addresses for every layer. + seen_ptrs: set[int] = set() + self.kv_caches_base_addr: list[int] = [] + self.block_len: list[int] = [] + + for cache in kv_caches.values(): + cache_storage = cache.untyped_storage() + base_addr = cache_storage.data_ptr() + region_len = cache_storage.nbytes() + + if base_addr not in seen_ptrs: + seen_ptrs.add(base_addr) + ret = self.store.register_buffer(base_addr, region_len) + if ret != 0: + logger.error( + "register_buffer failed for addr %#x len %d: %d", + base_addr, + region_len, + ret, + ) + + if not outer_dims: + # Blocks-first layout (FlashInfer / MLA): one segment. + self.kv_caches_base_addr.append(base_addr) + self.block_len.append(page_size_bytes) + else: + # K/V-first layout (FlashAttn / ROCm): split segments. + seg_stride = cache.stride(outer_dims[0]) * el + for idx in range(cache.shape[outer_dims[0]]): + self.kv_caches_base_addr.append(base_addr + idx * seg_stride) + self.block_len.append(seg_stride // self.num_blocks) + + logger.info( + "Registering KV_Caches. use_mla: %s, shape %s, " + "num_blocks: %d, block_len: %s, " + "per_key_bytes: %d, " + "num_segments: %d", + self.use_mla, + first_kv_cache.shape, + self.num_blocks, + list(set(self.block_len)), + sum(self.block_len), + len(self.kv_caches_base_addr), + ) + + self.token_database.set_kv_caches_base_addr(self.kv_caches_base_addr) + self.token_database.set_block_len(self.block_len) + + # Start transfer threads + if self.kv_role in ["kv_producer", "kv_both"]: + ready_event_sending = threading.Event() + self.kv_send_thread = KVCacheStoreSendingThread( + self.store, + self.token_database, + self.block_size, + self.tp_rank, + self.put_step, + self.kv_role, + ready_event_sending, + self.enable_kv_events, + ) + self.kv_send_thread.start() + + ready_event_recving = threading.Event() + self.kv_recv_thread = KVCacheStoreRecvingThread( + self.store, + self.token_database, + self.block_size, + self.tp_rank, + ready_event_recving, + ) + self.kv_recv_thread.start() + ready_event_recving.wait() + + def start_load_kv( + self, + metadata: MooncakeStoreConnectorMetadata, + ): + """No-op: loads are issued in get_finished() for overlap.""" + pass + + def wait_for_save( + self, + metadata: MooncakeStoreConnectorMetadata, + ): + """No-op: stores are issued in get_finished() for overlap.""" + pass + + def get_finished( + self, + finished_req_ids: set[str], + meta: MooncakeStoreConnectorMetadata, + ) -> tuple[set[str], set[str]]: + """Issue all I/O and get completed send/recv request IDs. + + All load and store I/O requests are issued here (after model + compute is launched on the compute stream) for better + compute-I/O overlap. + """ + # Issue async loads + for request in meta.requests: + load_spec = request.load_spec + if load_spec is None or not load_spec.can_load: + continue + + token_len = request.token_len_chunk + if (load_spec.kvpool_cached_tokens % self.block_size != 0) and ( + load_spec.kvpool_cached_tokens == token_len - 1 + ): + token_len = load_spec.kvpool_cached_tokens + 1 + else: + token_len = load_spec.kvpool_cached_tokens + load_spec.token_len = token_len + + assert self.kv_recv_thread is not None + self.kv_recv_thread.add_request(request) + + assert self.load_async, "load_async must be True for better performance." + # Issue stores with CUDA event synchronization + if self.kv_role in ["kv_producer", "kv_both"]: + current_event = None + for request in meta.requests: + if request.can_save: + current_event = torch.cuda.Event() + current_event.record() + break + + for request in meta.requests: + if not request.can_save: + continue + request.current_event = current_event + assert self.kv_send_thread is not None + self.kv_send_thread.add_stored_request(request.req_id) + self.kv_send_thread.add_request(request) + + # Check completion of previously queued transfers + done_sending = ( + self._get_and_clear_finished_sending(finished_req_ids, meta) + if self.kv_role in ["kv_producer", "kv_both"] + else set() + ) + + done_recving = ( + self.kv_recv_thread.get_and_clear_finished_requests() + if self.load_async and self.kv_recv_thread is not None + else set() + ) + + logger.debug( + "Completed send: %d, recv: %d, tp_rank: %d", + len(done_sending), + len(done_recving), + self.tp_rank, + ) + return done_sending, done_recving + + def _get_and_clear_finished_sending( + self, + finished_req_ids: set[str], + meta: MooncakeStoreConnectorMetadata, + ) -> set[str]: + assert self.kv_send_thread is not None + finished_sending: set[str] = set() + + for req_id in meta.preempted_req_ids: + self.kv_send_thread.delete_finished_stored_request(req_id) + + for req_id in self.kv_send_thread.stored_requests.copy(): + if ( + self.kv_send_thread.stored_requests[req_id] == 0 + and req_id in self.finished_store_req + ): + self.finished_store_req.remove(req_id) + finished_sending.add(req_id) + self.kv_send_thread.delete_finished_stored_request(req_id) + + for req_id in finished_req_ids: + req_remain_jobs = self.kv_send_thread.stored_requests.get(req_id) + if req_remain_jobs == 0: + finished_sending.add(req_id) + self.kv_send_thread.delete_finished_stored_request(req_id) + elif req_remain_jobs is not None: + self.finished_store_req.add(req_id) + + return finished_sending + + def lookup( + self, + token_len: int, + block_hashes: list[BlockHash], + ) -> int: + """Check how many prefix tokens exist in the store. + + Checks across all TP ranks and PP ranks. + """ + end = 0 + keys: list[str] = [] + try: + starts: list[int] = [] + for start, end, key in self.token_database.process_tokens( + token_len, block_hashes + ): + keys.append(key.to_string()) + starts.append(start) + + # Expand keys for all TP ranks + multi_tp_keys = keys[:] + for i in range(1, min(self.tp_size, self.num_kv_head)): + for item in keys: + new_str = item.replace("@tp_rank:0", f"@tp_rank:{i}", 1) + multi_tp_keys.append(new_str) + + # Expand keys for all PP ranks + pp_base_keys = multi_tp_keys.copy() + for i in range(1, self.pp_size): + for item in pp_base_keys: + new_str = item.replace("@pp_rank:0", f"@pp_rank:{i}", 1) + multi_tp_keys.append(new_str) + + res = self.store.batch_is_exist(multi_tp_keys) + + num_block = len(keys) + multi_tp_values = [ + res[i * num_block : (i + 1) * num_block] + for i in range(min(self.tp_size, self.num_kv_head) * self.pp_size) + ] + index = self._find_min_first_non_one_index(multi_tp_values) + if index != -1: + return starts[index] + except Exception as e: + logger.error("Remote connection failed in lookup: %s", e) + return 0 + return end + + @staticmethod + def _find_min_first_non_one_index( + arr: list[list[int]], + ) -> int: + try: + return min(idx for row in arr for idx, val in enumerate(row) if val != 1) + except ValueError: + return -1 + + def get_kv_events(self) -> list[BlockStored]: + if self.enable_kv_events and self.kv_send_thread is not None: + return self.kv_send_thread.get_kv_events() + return [] + + +# ============================================================ +# Lookup Key Server +# ============================================================ + + +class LookupKeyServer: + """ZMQ server on worker rank 0 for handling prefix lookup queries.""" + + def __init__( + self, + store_worker: MooncakeStoreWorker, + vllm_config: VllmConfig, + ): + self.decoder = MsgpackDecoder() + self.ctx = zmq.Context() # type: ignore[attr-defined] + socket_path = get_zmq_rpc_path_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: + all_frames = self.socket.recv_multipart(copy=False) + token_len = int.from_bytes(all_frames[0], byteorder="big") + hash_frames = all_frames[1:] + hashes_str = self.decoder.decode(hash_frames) + result = self.store_worker.lookup(token_len, hashes_str) + response = result.to_bytes(4, "big") + self.socket.send(response) + + self.thread = threading.Thread(target=process_request, daemon=True) + self.thread.start() + + def close(self): + self.socket.close(linger=0) + if os.path.exists(self._ipc_path): + os.unlink(self._ipc_path) + + +# ============================================================ +# Lookup Key Client +# ============================================================ + + +class LookupKeyClient: + """ZMQ client for querying prefix cache hits from worker.""" + + def __init__(self, vllm_config: VllmConfig): + self.encoder = MsgpackEncoder() + self.ctx = zmq.Context() # type: ignore[attr-defined] + socket_path = get_zmq_rpc_path_lookup(vllm_config) + self.socket = make_zmq_socket( + self.ctx, + socket_path, + zmq.REQ, # type: ignore[attr-defined] + bind=False, + ) + + def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int: + hash_strs = [h.hex() for h in block_hashes] + hash_frames = self.encoder.encode(hash_strs) + token_len_bytes = token_len.to_bytes(4, byteorder="big") + all_frames = [token_len_bytes] + list(hash_frames) + self.socket.send_multipart(all_frames, copy=False) + resp = self.socket.recv() + result = int.from_bytes(resp, "big") + return result + + def close(self): + self.socket.close(linger=0) + + +def get_zmq_rpc_path_lookup(vllm_config: VllmConfig) -> str: + """Construct IPC path for ZMQ lookup socket.""" + dp_rank = get_mooncake_dp_engine_index(vllm_config.parallel_config) + base_url = envs.VLLM_RPC_BASE_PATH + rpc_port = 0 + assert vllm_config.kv_transfer_config is not None + extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config + if "lookup_rpc_port" in extra_config: + rpc_port = extra_config["lookup_rpc_port"] + uid = os.getuid() + logger.debug("Base URL: %s, RPC Port: %s, UID: %s", base_url, rpc_port, uid) + return f"ipc://{base_url}/lookup_rpc_port_{rpc_port}_uid{uid}_dp_rank{dp_rank}"