diff --git a/tests/v1/kv_connector/unit/test_moriio_connector.py b/tests/v1/kv_connector/unit/test_moriio_connector.py index a8da6cf36d1..ee296292eac 100644 --- a/tests/v1/kv_connector/unit/test_moriio_connector.py +++ b/tests/v1/kv_connector/unit/test_moriio_connector.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import importlib.util +import socket import uuid from unittest.mock import MagicMock, patch @@ -9,7 +10,6 @@ import pytest import torch import zmq -from tests.conftest import _find_free_port from vllm.config import ( CacheConfig, DeviceConfig, @@ -23,12 +23,14 @@ from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_common import ( MoRIIOAgentMetadata, MoRIIOConnectorMetadata, MoRIIOConstants, + MoRIIOMode, resolve_host_ip, zmq_ctx, ) from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_connector import ( KVConnectorRole, MoRIIOConnector, + MoRIIOConnectorScheduler, MoRIIOConnectorWorker, ) from vllm.platforms import current_platform @@ -46,6 +48,12 @@ from vllm.v1.kv_cache_interface import ( from .utils import create_request, create_scheduler +def _find_free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("", 0)) + return sock.getsockname()[1] + + def _make_test_kv_cache_config() -> KVCacheConfig: layer_names = ["layer0", "layer1", "layer2"] return KVCacheConfig( @@ -121,6 +129,16 @@ def _setup_kv_transfer_request( return request +def _write_consumer_scheduler_for_finished_request(tp_size: int = 2): + scheduler = MoRIIOConnectorScheduler.__new__(MoRIIOConnectorScheduler) + scheduler.is_producer = False + scheduler.mode = MoRIIOMode.WRITE + scheduler.tp_size = tp_size + scheduler._reqs_need_recv = {} + scheduler.unmap_request_id = MagicMock() + return scheduler + + class FakeMoRIIOWrapper: # A fake MoRIIOWrapper for testing purposes def __init__(self, *args, **kwargs): @@ -177,7 +195,7 @@ class FakeMoRIIOWrapper: def _handle_completion_message(self, msg: str): pass - def send_notify(self, req_ids, remote_ip, remote_port): + def send_notify(self, req_ids, remote_ip, remote_port, message_type=None): pass def pop_finished_req_ids(self): @@ -434,6 +452,61 @@ def test_read_mode_loads_remote_block_ids(): assert block_id == block.block_id, f"{block_id} != {block.block_id}" +@pytest.mark.parametrize( + ("transfer_id", "extra_params", "expected_notifications"), + [ + pytest.param( + "xfer-7", + {"remote_host": "127.0.0.1", "remote_notify_port": 7000}, + [ + ("xfer-7", "127.0.0.1", 7000), + ("xfer-7", "127.0.0.1", 7001), + ], + id="address-available", + ), + pytest.param("xfer-8", {}, [], id="address-unavailable-plain-id"), + ], +) +def test_write_mode_finished_before_alloc_releases_prefill_blocks( + transfer_id, extra_params, expected_notifications +): + scheduler = _write_consumer_scheduler_for_finished_request(tp_size=2) + notifications = [] + scheduler._send_transfer_release = lambda transfer_id, host, port: ( + notifications.append((transfer_id, host, port)) + ) + request = create_request(request_id=7, do_remote_prefill=True) + request.request_id = "plain-decode-id" + request.kv_transfer_params = { + "do_remote_prefill": True, + "do_remote_decode": False, + "transfer_id": transfer_id, + } | extra_params + + delay_free, new_params = scheduler.request_finished(request, block_ids=[]) + + assert not delay_free + assert new_params is None + assert request.kv_transfer_params["do_remote_prefill"] is False + assert scheduler._reqs_need_recv == {} + assert notifications == expected_notifications + + +def test_send_transfer_release_sends_structured_release_message(): + scheduler = _write_consumer_scheduler_for_finished_request() + path = make_zmq_path("tcp", "127.0.0.1", 7000) + sock = MagicMock() + scheduler.paths = {path: sock} + + scheduler._send_transfer_release("xfer-7", "127.0.0.1", 7000) + + payload = sock.send.call_args.args[0] + assert msgspec.msgpack.decode(payload) == { + "type": "release", + "transfer_id": "xfer-7", + } + + @pytest.mark.skipif( not aiter_available, reason="Requires aiter package for ROCm FlashAttention backend" ) diff --git a/tests/v1/kv_connector/unit/test_moriio_kv_layout.py b/tests/v1/kv_connector/unit/test_moriio_kv_layout.py index 5b3219db867..61146ce3c86 100644 --- a/tests/v1/kv_connector/unit/test_moriio_kv_layout.py +++ b/tests/v1/kv_connector/unit/test_moriio_kv_layout.py @@ -2,7 +2,11 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import importlib.util +import threading +from collections import OrderedDict, defaultdict +from queue import Queue from types import SimpleNamespace +from typing import Any import pytest import torch @@ -19,16 +23,33 @@ if not (current_platform.is_rocm() and mori_available): allow_module_level=True, ) +moriio_common = importlib.import_module( + "vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_common" +) +moriio_engine = importlib.import_module( + "vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_engine" +) moriio_layout = importlib.import_module( "vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_layout" ) +msgpack = importlib.import_module("msgpack") + +ROLE = moriio_common.ROLE +MoRIIOError = moriio_common.MoRIIOError +RemoteAllocInfo = moriio_common.RemoteAllocInfo +WriteTask = moriio_common.WriteTask +set_role = moriio_common.set_role +MoRIIOWrapper = moriio_engine.MoRIIOWrapper +MoRIIOWriter = moriio_engine.MoRIIOWriter -def _full_spec(block_size: int = 4) -> FullAttentionSpec: +def _full_spec( + block_size: int = 4, num_kv_heads: int = 2, head_size: int = 3 +) -> FullAttentionSpec: return FullAttentionSpec( block_size=block_size, - num_kv_heads=2, - head_size=3, + num_kv_heads=num_kv_heads, + head_size=head_size, dtype=torch.bfloat16, ) @@ -59,55 +80,216 @@ def _remote_meta(num_blocks: int = 16) -> SimpleNamespace: return SimpleNamespace(num_blocks=num_blocks) -def test_separated_kv_layout_uses_kv_axis_zero_and_block_axis_one(): - cache = torch.empty((2, 8, 4, 2, 3), dtype=torch.bfloat16) - worker = _worker({"layer": cache}, {"layer": _full_spec()}) +def _writer_with_fake_worker(fake_worker: Any) -> Any: + writer = MoRIIOWriter.__new__(MoRIIOWriter) + writer._worker_ref = lambda: fake_worker + writer._write_task_q = Queue() + writer._write_state_lock = threading.Lock() + writer._scheduled_writes = defaultdict(int) + writer._scheduled_layers = defaultdict(set) + writer._sealed_writes = {} + writer.ensure_worker_started = lambda: None + return writer + + +def _wrapper_for_messages() -> Any: + wrapper = MoRIIOWrapper.__new__(MoRIIOWrapper) + wrapper.lock = threading.Lock() + wrapper.done_remote_allocate_req_dict = {} + wrapper.done_req_ids = [] + wrapper.done_write_cache_req_ids = [] + wrapper._terminal_transfer_ids = OrderedDict() + return wrapper + + +def _write_task(layer_name: str, transfer_id: str = "xfer") -> Any: + return WriteTask( + request_id="req", + transfer_id=transfer_id, + dst_engine_id="remote-engine", + local_block_ids=[1, 3], + remote_block_ids_hint=None, + layer_name=layer_name, + event=None, + remote_notify_port=7000, + remote_ip="127.0.0.1", + ) + + +@pytest.mark.parametrize( + ("shape", "spec", "remote_num_blocks", "expected_geometry", "expected_offsets"), + [ + pytest.param( + (2, 8, 4, 2, 3), + _full_spec(), + 16, + { + "block_stride": 24, + "local_kv_stride": 192, + "remote_kv_stride": 384, + "split_kv_regions": True, + }, + ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48]), + id="separated", + ), + pytest.param( + (8, 2, 4, 2, 3), + _full_spec(), + 16, + { + "block_stride": 48, + "local_kv_stride": 24, + "remote_kv_stride": 24, + "split_kv_regions": False, + }, + ([96, 288], [384, 480], [96, 96]), + id="interleaved", + ), + pytest.param( + (2, 8, 2, 4, 3), + _full_spec(), + 16, + { + "block_size": 4, + "block_stride": 24, + "local_kv_stride": 192, + "remote_kv_stride": 384, + "split_kv_regions": True, + }, + ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48]), + id="shuffled-separated", + ), + pytest.param( + (8, 2, 2, 4, 3), + _full_spec(), + 16, + { + "block_size": 4, + "block_stride": 48, + "local_kv_stride": 24, + "remote_kv_stride": 24, + "split_kv_regions": False, + }, + ([96, 288], [384, 480], [96, 96]), + id="shuffled-interleaved", + ), + pytest.param( + (2, 16, 2, 2, 3), + _full_spec(), + 16, + { + "num_blocks": 8, + "block_size": 4, + "block_stride": 24, + "local_kv_stride": 192, + "remote_kv_stride": 384, + "split_kv_regions": True, + }, + ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48]), + id="separated-kernel-blocks", + ), + pytest.param( + (16, 2, 2, 2, 3), + _full_spec(), + 16, + { + "num_blocks": 8, + "block_size": 4, + "block_stride": 48, + "local_kv_stride": None, + "remote_kv_stride": None, + "transfers_per_block": 1, + }, + ([96, 288], [384, 480], [96, 96]), + id="interleaved-kernel-blocks", + ), + pytest.param( + (2, 32, 8, 2, 3), + _full_spec(block_size=16, num_kv_heads=8), + 8, + { + "num_blocks": 4, + "block_size": 16, + "block_len": 768, + "block_stride": 384, + "local_kv_stride": 1536, + "remote_kv_stride": 3072, + "split_kv_regions": True, + }, + ( + [768, 2304, 3840, 5376], + [3072, 3840, 9216, 9984], + [768, 768, 768, 768], + ), + id="separated-kernel-axis-from-spec", + ), + pytest.param( + (32, 2, 8, 2, 3), + _full_spec(block_size=16, num_kv_heads=8), + 8, + { + "num_blocks": 4, + "block_size": 16, + "block_len": 1536, + "block_stride": 768, + "local_kv_stride": None, + "remote_kv_stride": None, + "transfers_per_block": 1, + }, + ([1536, 4608], [6144, 7680], [1536, 1536]), + id="interleaved-kernel-axis-from-spec", + ), + pytest.param( + (8, 4, 3), + _mla_spec(), + 16, + { + "block_stride": 12, + "local_kv_stride": None, + "remote_kv_stride": None, + "transfers_per_block": 1, + }, + ([24, 72], [96, 120], [24, 24]), + id="mla-key-only", + ), + ], +) +def test_supported_layouts_compute_expected_geometry_and_offsets( + shape, spec, remote_num_blocks, expected_geometry, expected_offsets +): + cache = torch.empty(shape, dtype=torch.bfloat16) + worker = _worker({"layer": cache}, {"layer": spec}) geometry = moriio_layout.get_layer_transfer_geometry( - "layer", cache, worker.layer_to_spec, remote_num_blocks=16 + "layer", cache, worker.layer_to_spec, remote_num_blocks=remote_num_blocks ) - assert geometry.block_stride == 24 - assert geometry.local_kv_stride == 192 - assert geometry.remote_kv_stride == 384 - assert geometry.split_kv_regions + for field, expected in expected_geometry.items(): + assert getattr(geometry, field) == expected - assert moriio_layout.compute_block_transfer_offsets( - "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks - ) == ([48, 144, 432, 528], [192, 240, 960, 1008], [48, 48, 48, 48]) - - -def test_interleaved_kv_layout_uses_block_axis_zero_and_kv_axis_one(): - cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16) - worker = _worker({"layer": cache}, {"layer": _full_spec()}) - - geometry = moriio_layout.get_layer_transfer_geometry( - "layer", cache, worker.layer_to_spec, remote_num_blocks=16 + assert ( + moriio_layout.compute_block_transfer_offsets( + "layer", + cache, + worker.layer_to_spec, + [1, 3], + [4, 5], + remote_num_blocks, + ) + == expected_offsets ) - assert geometry.block_stride == 48 - assert geometry.local_kv_stride == 24 - assert geometry.remote_kv_stride == 24 - assert not geometry.split_kv_regions - - assert moriio_layout.compute_block_transfer_offsets( - "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks - ) == ([96, 288], [384, 480], [96, 96]) -def test_mla_key_only_layout_transfers_one_slab_per_block(): - cache = torch.empty((8, 4, 3), dtype=torch.bfloat16) - worker = _worker({"layer": cache}, {"layer": _mla_spec()}) - - geometry = moriio_layout.get_layer_transfer_geometry( - "layer", cache, worker.layer_to_spec, remote_num_blocks=16 +def test_kernel_block_layout_without_spec_dimensions_rejects_ambiguous_axes(): + cache = torch.empty((2, 32, 8, 2, 3), dtype=torch.bfloat16) + worker = _worker( + {"layer": cache}, + {"layer": SimpleNamespace(block_size=16)}, ) - assert geometry.block_stride == 12 - assert geometry.local_kv_stride is None - assert geometry.remote_kv_stride is None - assert geometry.transfers_per_block == 1 - assert moriio_layout.compute_block_transfer_offsets( - "layer", cache, worker.layer_to_spec, [1, 3], [4, 5], _remote_meta().num_blocks - ) == ([24, 72], [96, 120], [24, 24]) + with pytest.raises(ValueError, match="Ambiguous MoRIIO kernel-block"): + moriio_layout.get_layer_transfer_geometry( + "layer", cache, worker.layer_to_spec, remote_num_blocks=8 + ) def test_mixed_layers_compute_distinct_offsets_per_layer(): @@ -155,6 +337,319 @@ def test_mixed_layers_compute_distinct_offsets_per_layer(): assert interleaved != indexer +def test_write_transfer_plan_caches_offsets_per_geometry(): + kv_caches = { + "dense0": torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16), + "dense1": torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16), + "indexer": torch.empty((8, 4, 3), dtype=torch.bfloat16), + } + calls: list[str] = [] + + class FakeWorker: + kv_caches: dict[str, torch.Tensor] + layer_name_to_local_kv_cache_metadata: dict[str, list[Any]] + + def _compute_block_transfer_offsets( + self, layer_name, local_block_ids, remote_block_ids, remote_moriio_meta + ): + calls.append(layer_name) + call_id = len(calls) + return ([call_id], [call_id + 10], [call_id + 20]) + + fake_worker = FakeWorker() + fake_worker.kv_caches = kv_caches + fake_worker.layer_name_to_local_kv_cache_metadata = {name: [] for name in kv_caches} + writer = MoRIIOWriter.__new__(MoRIIOWriter) + writer._worker_ref = lambda: fake_worker + request_info = RemoteAllocInfo(block_ids=[4, 5]) + remote_meta = _remote_meta() + + dense0_plan = writer._prepare_transfer_plan( + SimpleNamespace( + layer_name="dense0", + local_block_ids=[1, 3], + request_id="req", + transfer_id="xfer", + ), + request_info, + remote_meta, + ) + dense1_plan = writer._prepare_transfer_plan( + SimpleNamespace( + layer_name="dense1", + local_block_ids=[1, 3], + request_id="req", + transfer_id="xfer", + ), + request_info, + remote_meta, + ) + indexer_plan = writer._prepare_transfer_plan( + SimpleNamespace( + layer_name="indexer", + local_block_ids=[1, 3], + request_id="req", + transfer_id="xfer", + ), + request_info, + remote_meta, + ) + + assert calls == ["dense0", "indexer"] + assert dense0_plan.transfer_local_offsets == [1] + assert dense1_plan.transfer_local_offsets == [1] + assert indexer_plan.transfer_local_offsets == [2] + assert len(request_info.transfer_offsets) == 2 + + +def test_write_scheduler_deduplicates_layers_and_seals_expected_count(): + request_info = RemoteAllocInfo(block_ids=[4, 5]) + wrapper = _wrapper_for_messages() + wrapper.done_remote_allocate_req_dict["xfer"] = request_info + writer = _writer_with_fake_worker(SimpleNamespace(moriio_wrapper=wrapper)) + + assert writer.schedule_write(_write_task("dense0")) + assert not writer.schedule_write(_write_task("dense0")) + assert writer.schedule_write(_write_task("indexer")) + + assert writer._write_task_q.qsize() == 2 + writer.seal_pending_transfers() + + assert request_info.writes_expected == 2 + assert writer._sealed_writes["xfer"] == 2 + + +def test_write_completion_notifies_once_after_all_sealed_writes_finish(): + class FakeWrapper: + def __init__(self): + self.done_remote_allocate_req_dict = {} + self.done_req_ids = [] + self.lock = threading.Lock() + self.notifications = [] + self.wait_count = 0 + self.waited_statuses = [] + self._terminal_transfer_ids = OrderedDict() + + def waiting_for_transfer_complete(self, transfer_statuses=None): + self.wait_count += 1 + self.waited_statuses.append(list(transfer_statuses or [])) + + def _is_transfer_terminal_locked(self, transfer_id): + return transfer_id in self._terminal_transfer_ids + + def _mark_transfer_terminal_locked(self, transfer_id): + self._terminal_transfer_ids[transfer_id] = None + + def send_notify(self, transfer_id, remote_ip, remote_port, message_type=None): + self.notifications.append( + (transfer_id, remote_ip, remote_port, message_type) + ) + + wrapper = FakeWrapper() + request_info = RemoteAllocInfo(block_ids=[4, 5], writes_expected=2) + request_info.transfer_statuses.extend(["status-a", "status-b"]) + request_info.completion_request_id = "req" + request_info.completion_remote_notify_port = 7000 + request_info.completion_remote_ip = "127.0.0.1" + wrapper.done_remote_allocate_req_dict["xfer"] = request_info + writer = _writer_with_fake_worker( + SimpleNamespace(moriio_wrapper=wrapper, tp_rank=2) + ) + writer._scheduled_writes["xfer"] = 2 + writer._scheduled_layers["xfer"] = {"dense0", "indexer"} + writer._sealed_writes["xfer"] = 2 + + writer._mark_write_done("xfer", request_info) + assert wrapper.notifications == [] + writer._mark_write_done("xfer", request_info) + writer._finalize_if_complete("xfer", request_info) + + assert wrapper.notifications == [("xfer", "127.0.0.1", 7002, "write_done")] + assert wrapper.done_req_ids == ["xfer"] + assert wrapper.done_remote_allocate_req_dict == {} + assert wrapper.wait_count == 1 + assert wrapper.waited_statuses == [["status-a", "status-b"]] + assert request_info.transfer_statuses == [] + assert wrapper._is_transfer_terminal_locked("xfer") + + +def test_moriio_wrapper_waits_scoped_statuses_without_global_drain(): + class FakeStatus: + def __init__(self): + self.checked = 0 + + def Succeeded(self): + self.checked += 1 + return True + + def Failed(self): + return False + + wrapper = MoRIIOWrapper.__new__(MoRIIOWrapper) + wrapper.lock = threading.Lock() + wrapper._transfer_timeout = 1 + global_status = FakeStatus() + scoped_status = FakeStatus() + wrapper.transfer_status = [global_status] + + wrapper.waiting_for_transfer_complete([scoped_status]) + + assert scoped_status.checked == 1 + assert global_status.checked == 0 + assert wrapper.transfer_status == [global_status] + + +def test_write_failure_marks_terminal_and_clears_scheduled_state(): + wrapper = _wrapper_for_messages() + wrapper.done_remote_allocate_req_dict["xfer"] = RemoteAllocInfo(block_ids=[4, 5]) + writer = _writer_with_fake_worker(SimpleNamespace(moriio_wrapper=wrapper)) + writer._scheduled_writes["xfer"] = 2 + writer._scheduled_layers["xfer"] = {"dense0", "indexer"} + writer._sealed_writes["xfer"] = 2 + + writer._mark_request_done("xfer") + + assert wrapper.done_req_ids == ["xfer"] + assert wrapper.done_remote_allocate_req_dict == {} + assert wrapper._is_transfer_terminal_locked("xfer") + assert "xfer" not in writer._scheduled_writes + assert "xfer" not in writer._scheduled_layers + assert "xfer" not in writer._sealed_writes + + +def test_schedule_write_rejects_terminal_transfer_without_recreating_state(): + wrapper = _wrapper_for_messages() + wrapper.done_remote_allocate_req_dict["xfer"] = RemoteAllocInfo(block_ids=[4, 5]) + writer = _writer_with_fake_worker(SimpleNamespace(moriio_wrapper=wrapper)) + writer._scheduled_writes["xfer"] = 1 + writer._scheduled_layers["xfer"] = {"dense0"} + writer._sealed_writes["xfer"] = 1 + + writer._mark_request_done("xfer") + + assert not writer.schedule_write(_write_task("indexer")) + assert writer._write_task_q.empty() + assert "xfer" not in writer._scheduled_writes + assert "xfer" not in writer._scheduled_layers + assert "xfer" not in writer._sealed_writes + + +def test_late_remote_blocks_message_is_ignored_after_transfer_done(): + set_role(ROLE.PRODUCER) + wrapper = _wrapper_for_messages() + with wrapper.lock: + wrapper._mark_transfer_terminal_locked("xfer") + + wrapper._handle_message( + msgpack.dumps( + { + "type": "remote_blocks", + "req_id": "req", + "transfer_id": "xfer", + "block_notify_list": [4, 5], + "decode_rank": 3, + } + ) + ) + + assert "xfer" not in wrapper.done_remote_allocate_req_dict + + +@pytest.mark.parametrize( + ("role", "payload", "expected"), + [ + pytest.param( + ROLE.PRODUCER, + msgpack.dumps( + { + "type": "remote_blocks", + "req_id": "req", + "transfer_id": "xfer", + "block_notify_list": [4, 5], + "decode_rank": 3, + } + ), + "remote_blocks", + id="remote-blocks", + ), + pytest.param( + ROLE.CONSUMER, + msgpack.dumps({"type": "write_done", "transfer_id": "xfer"}), + "write_done", + id="write-done", + ), + pytest.param( + ROLE.PRODUCER, + msgpack.dumps({"type": "release", "transfer_id": "xfer"}), + "release", + id="release", + ), + pytest.param(None, b"xfer", "plain", id="plain-string"), + ], +) +def test_moriio_wrapper_routes_valid_messages(role, payload, expected): + wrapper = _wrapper_for_messages() + completions: list[str] = [] + if role is not None: + set_role(role) + if expected == "plain": + wrapper._handle_completion_message = completions.append + + wrapper._handle_message(payload) + + if expected == "remote_blocks": + request_info = wrapper.done_remote_allocate_req_dict["xfer"] + assert request_info.block_ids == [4, 5] + assert request_info.decode_dp_rank == 3 + elif expected == "write_done": + assert wrapper.done_write_cache_req_ids == ["xfer"] + elif expected == "release": + assert wrapper.done_req_ids == ["xfer"] + assert wrapper._is_transfer_terminal_locked("xfer") + else: + assert completions == ["xfer"] + + +@pytest.mark.parametrize( + ("role", "payload", "match"), + [ + pytest.param( + None, + msgpack.dumps({"type": "unknown", "transfer_id": "xfer"}), + "Unhandled structured message type", + id="unknown-structured-type", + ), + pytest.param( + ROLE.PRODUCER, + msgpack.dumps( + { + "type": "remote_blocks", + "req_id": "req", + "transfer_id": "xfer", + "block_notify_list": [], + } + ), + "block_notify_list cannot be empty", + id="empty-remote-blocks", + ), + pytest.param( + None, + b"", + "Unhandled message format", + id="empty-completion", + ), + ], +) +def test_moriio_wrapper_rejects_invalid_messages(role, payload, match): + wrapper = _wrapper_for_messages() + if role is not None: + set_role(role) + wrapper._handle_completion_message = lambda msg: None + + with pytest.raises(MoRIIOError, match=match): + wrapper._handle_message(payload) + + def test_block_id_length_mismatch_raises_value_error(): cache = torch.empty((8, 2, 4, 2, 3), dtype=torch.bfloat16) worker = _worker({"layer": cache}, {"layer": _full_spec()}) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_common.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_common.py index 73b3d2e1484..07fae409429 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_common.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_common.py @@ -78,8 +78,17 @@ class RemoteAllocInfo: block_ids: list[int] writes_done: int = 0 + writes_expected: int | None = None decode_dp_rank: int = 0 - transfer_offset: tuple[list[int], list[int], list[int]] | None = None + completion_request_id: str | None = None + completion_remote_notify_port: int | None = None + completion_remote_ip: str | None = None + completion_notified: bool = False + transfer_statuses: list[Any] = field(default_factory=list) + transfer_offsets: dict[ + tuple[tuple[int, ...], tuple[int, ...], torch.dtype], + tuple[list[int], list[int], list[int]], + ] = field(default_factory=dict) class ROLE(Enum): @@ -434,12 +443,21 @@ class MoRIIOConnectorMetadata(KVConnectorMetadata): ): transfer_id = kv_transfer_params["transfer_id"] - # Parse host/ports from the request_id. The router embeds both zmq_addresses - # in the request_id - peer_zmq = get_peer_zmq_from_request_id(request_id, is_producer=write_mode) - remote_host, remote_handshake_port, remote_notify_port = ( - parse_moriio_zmq_address(peer_zmq) - ) + remote_host = kv_transfer_params.get("remote_host") + remote_handshake_port = kv_transfer_params.get("remote_handshake_port") + remote_notify_port = kv_transfer_params.get("remote_notify_port") + if ( + remote_host is None + or remote_handshake_port is None + or remote_notify_port is None + ): + # Parse host/ports from the request_id. The router embeds both + # zmq_addresses in PD request IDs, but WRITE decode requests may carry + # a plain request ID and get the remote address via kv_transfer_params. + peer_zmq = get_peer_zmq_from_request_id(request_id, is_producer=write_mode) + remote_host, remote_handshake_port, remote_notify_port = ( + parse_moriio_zmq_address(peer_zmq) + ) _req = ReqMeta( transfer_id=transfer_id, @@ -447,9 +465,9 @@ class MoRIIOConnectorMetadata(KVConnectorMetadata): remote_block_ids=kv_transfer_params["remote_block_ids"], remote_engine_id=kv_transfer_params["remote_engine_id"], remote_host=remote_host, - remote_port=remote_handshake_port, - remote_handshake_port=remote_handshake_port, - remote_notify_port=remote_notify_port, + remote_port=int(remote_handshake_port), + remote_handshake_port=int(remote_handshake_port), + remote_notify_port=int(remote_notify_port), tp_size=kv_transfer_params.get("tp_size", 1), remote_dp_size=kv_transfer_params.get("remote_dp_size", 1), ) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py index a41bb5789f0..e119ea7b7d4 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_connector.py @@ -234,7 +234,13 @@ class MoRIIOConnector(KVConnectorBase_V1): return None def wait_for_save(self): - pass + if self.mode != MoRIIOMode.WRITE or get_role() != ROLE.PRODUCER: + return + assert self.connector_worker is not None + assert isinstance(self._connector_metadata, MoRIIOConnectorMetadata), ( + "Connector metadata not initialized yet" + ) + self.connector_worker.wait_for_save(self._connector_metadata) def shutdown(self): if self.connector_worker is not None: @@ -390,6 +396,49 @@ class MoRIIOConnectorScheduler: serialized_data = msgpack.dumps(data) self.paths[path].send(serialized_data) + def _send_transfer_release(self, transfer_id: TransferId, host: str, port: int): + path = make_zmq_path("tcp", host, port) + if path not in self.paths: + ctx = zmq.Context.instance() + sock = make_zmq_socket( + ctx=ctx, path=path, socket_type=zmq.DEALER, bind=False + ) + self.paths[path] = sock + + self.paths[path].send( + msgpack.dumps({"type": "release", "transfer_id": transfer_id}) + ) + + def _release_write_prefill_blocks(self, request_id: ReqId, params: dict[str, Any]): + transfer_id = params.get("transfer_id") + if transfer_id is None: + logger.warning( + "Cannot release WRITE prefill blocks for request %s: " + "missing transfer_id", + request_id, + ) + return + + remote_dp_rank = params.get("remote_dp_rank", 0) + remote_host = params.get("remote_host") + remote_notify_port = params.get("remote_notify_port") + if remote_host is None or remote_notify_port is None: + try: + peer_zmq = get_peer_zmq_from_request_id(request_id, is_producer=False) + remote_host, _, remote_notify_port = parse_moriio_zmq_address(peer_zmq) + except ValueError: + logger.warning( + "Cannot release WRITE prefill blocks for request %s: " + "missing remote notify address", + request_id, + ) + return + + remote_notify_port = int(remote_notify_port) + for tp_index in range(self.tp_size): + target_port = remote_notify_port + get_port_offset(remote_dp_rank, tp_index) + self._send_transfer_release(transfer_id, remote_host, target_port) + def update_state_after_alloc( self, request: "Request", @@ -443,11 +492,20 @@ class MoRIIOConnectorScheduler: ) remote_dp_rank = request.kv_transfer_params.get("remote_dp_rank", 0) - - peer_zmq = get_peer_zmq_from_request_id( - request.request_id, is_producer=False + remote_host = request.kv_transfer_params.get("remote_host") + remote_notify_port = request.kv_transfer_params.get( + "remote_notify_port" ) - remote_host, _, remote_notify_port = parse_moriio_zmq_address(peer_zmq) + if remote_host is None or remote_notify_port is None: + peer_zmq = get_peer_zmq_from_request_id( + request.request_id, is_producer=False + ) + remote_host, _, remote_notify_port = parse_moriio_zmq_address( + peer_zmq + ) + remote_notify_port = int(remote_notify_port) + + block_ids = blocks.get_block_ids()[0] for tp_index in range(self.tp_size): target_port = remote_notify_port + get_port_offset( @@ -457,7 +515,7 @@ class MoRIIOConnectorScheduler: self.send_notify_block( req_id=request.request_id, transfer_id=request.kv_transfer_params["transfer_id"], - block_notify_list=blocks.get_block_ids()[0], + block_notify_list=block_ids, host=remote_host, port=target_port, ) @@ -473,60 +531,31 @@ class MoRIIOConnectorScheduler: meta = MoRIIOConnectorMetadata() meta.transfer_id_to_request_id = self.transfer_id_to_request_id - if self.mode == MoRIIOMode.WRITE: - # when async_load_kv finished, - # new reqs will be added to scheduler_output.scheduled_new_reqs + if self.mode == MoRIIOMode.WRITE and get_role() == ROLE.PRODUCER: + # This is the logic for checking against chunked prefill. + # When the last chunk is identified, + # It places the request metadata into the saving queue. - if get_role() == ROLE.CONSUMER: - for new_req in scheduler_output.scheduled_new_reqs: - red_id = new_req.req_id - local_block_ids = list(new_req.block_ids)[0] - assert new_req.sampling_params is not None, ( - f"sampling_params is None for req {new_req.req_id}" - ) - assert hasattr(new_req.sampling_params, "extra_args"), ( - f"sampling_params missing extra_args for req {new_req.req_id}" - ) - kv_transfer_params = ( - new_req.sampling_params.extra_args.get("kv_transfer_params", {}) - if new_req.sampling_params.extra_args - else {} - ) - meta.add_new_req( - red_id, - local_block_ids, - kv_transfer_params, - ) - if get_role() == ROLE.PRODUCER: - # This is the logic for checking against chunked prefill. - # When the last chunk is identified, - # It places the request metadata into the saving queue. + for i, req_id in enumerate(scheduler_output.scheduled_cached_reqs.req_ids): + new_block_ids = scheduler_output.scheduled_cached_reqs.new_block_ids[i] - for i, req_id in enumerate( - scheduler_output.scheduled_cached_reqs.req_ids - ): - new_block_ids = ( - scheduler_output.scheduled_cached_reqs.new_block_ids[i] - ) - - if new_block_ids is not None: - block_ids = new_block_ids[0] - # TODO : hybrid attn, etc - req, existing_blocks = self._reqs_need_pending_save[req_id] - updated_blocks = list(existing_blocks) + (block_ids) - self._reqs_need_pending_save[req_id] = (req, updated_blocks) - if ( - len(self._reqs_need_pending_save[req_id][1]) - * self.block_size - >= req.num_prompt_tokens - ): - meta.add_new_req( - request_id=req_id, - local_block_ids=self._reqs_need_pending_save[req_id][1], - kv_transfer_params=req.kv_transfer_params or {}, - write_mode=True, - ) - del self._reqs_need_pending_save[req_id] + if new_block_ids is not None: + block_ids = new_block_ids[0] + # TODO : hybrid attn, etc + req, existing_blocks = self._reqs_need_pending_save[req_id] + updated_blocks = list(existing_blocks) + (block_ids) + self._reqs_need_pending_save[req_id] = (req, updated_blocks) + if ( + len(self._reqs_need_pending_save[req_id][1]) * self.block_size + >= req.num_prompt_tokens + ): + meta.add_new_req( + request_id=req_id, + local_block_ids=self._reqs_need_pending_save[req_id][1], + kv_transfer_params=req.kv_transfer_params or {}, + write_mode=True, + ) + del self._reqs_need_pending_save[req_id] # Loop through scheduled reqs and convert to ReqMeta. for req_id, (req, block_ids) in self._reqs_need_recv.items(): @@ -601,9 +630,15 @@ class MoRIIOConnectorScheduler: # update_state_after_alloc must not have been called (the request # must have been aborted before it was scheduled). # To avoid stranding the prefill blocks in the prefill instance, - # we must add empty block_ids to _reqs_need_recv so that our - # worker side will notify and free blocks in the prefill instance. - self._reqs_need_recv[request.request_id] = (request, []) + # READ mode adds empty block_ids to _reqs_need_recv so the worker + # side notifies the prefill instance. WRITE mode should notify the + # producer directly: there is no decode allocation for the producer + # to write into, and a plain request_id may not contain router- + # embedded MoRIIO ZMQ addresses. + if self.mode == MoRIIOMode.WRITE: + self._release_write_prefill_blocks(request.request_id, params) + else: + self._reqs_need_recv[request.request_id] = (request, []) params["do_remote_prefill"] = False return False, None @@ -635,6 +670,10 @@ class MoRIIOConnectorScheduler: do_remote_decode=False, remote_block_ids=computed_block_ids, remote_engine_id=self.engine_id, + remote_host=self.host_ip, + remote_handshake_port=self.handshake_port, + remote_notify_port=self.side_notify_port, + remote_dp_size=self.vllm_config.parallel_config.data_parallel_size, tp_size=self.vllm_config.parallel_config.tensor_parallel_size, transfer_id=params["transfer_id"], ) @@ -1480,7 +1519,7 @@ class MoRIIOConnectorWorker: metadata: MoRIIOConnectorMetadata, layer_name: str, kv_layer: torch.Tensor, - attn_metadata: "AttentionMetadata", + attn_metadata: "AttentionMetadata | None", **kwargs, ): if not self.is_producer: @@ -1604,6 +1643,12 @@ class MoRIIOConnectorWorker: self._reqs_to_send.update(metadata.reqs_to_send) + def wait_for_save(self, metadata: MoRIIOConnectorMetadata): + if self.mode == MoRIIOMode.WRITE and self.is_producer: + for layer_name, kv_layer in self.kv_caches.items(): + self.save_kv_layer(metadata, layer_name, kv_layer, None) + self._writer.seal_pending_transfers() + def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): logger.debug( "Remote agent %s available, calling _read_blocks for req %s", diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_engine.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_engine.py index 3ca5f37ca90..3a90151f7e6 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_engine.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_engine.py @@ -2,6 +2,8 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import threading import time +from collections import OrderedDict, defaultdict +from queue import Empty, Queue from typing import TYPE_CHECKING, Any from weakref import ref as weakref_ref @@ -18,8 +20,6 @@ from vllm.utils.network_utils import ( if TYPE_CHECKING: from mori.io import BackendType -from queue import Empty, Queue - from vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_common import ( ROLE, HandshakeError, @@ -61,10 +61,25 @@ except ImportError: """Write task execution logic for MoRIIO connector.""" +_MAX_TERMINAL_TRANSFER_IDS = 4096 + + +WriteGeometryKey = tuple[tuple[int, ...], tuple[int, ...], torch.dtype] + + +def _get_write_geometry_key(kv_cache: torch.Tensor) -> WriteGeometryKey: + return (tuple(kv_cache.shape), tuple(kv_cache.stride()), kv_cache.dtype) + + class MoRIIOWriter: """Handles write operations for KV cache transfers. - Implements distributed KV cache transfer using the MoRIIO library - for RDMA-based communication between prefill and decode instances.""" + + WRITE mode state machine: + D sends destination block allocation, P schedules one write per layer + after the layer CUDA event, P seals the scheduled write count after + forward, then P notifies D and releases P blocks after all scheduled + writes complete. + """ def __init__(self, worker: "MoRIIOConnectorWorker"): """Initialize the writer. @@ -76,7 +91,11 @@ class MoRIIOWriter: self._write_task_q: Queue[WriteTask] = Queue() self._write_worker_started = False self._write_worker_lock = threading.Lock() + self._write_state_lock = threading.Lock() self._deferred_tasks: list[WriteTask] = [] + self._scheduled_writes: dict[TransferId, int] = defaultdict(int) + self._scheduled_layers: dict[TransferId, set[str]] = defaultdict(set) + self._sealed_writes: dict[TransferId, int] = {} self._defer_timeout = worker.moriio_config.defer_timeout @property @@ -106,14 +125,55 @@ class MoRIIOWriter: thread.start() logger.info("Started MoRIIO write worker thread") - def schedule_write(self, task: WriteTask) -> None: + def schedule_write(self, task: WriteTask) -> bool: """Schedule a write task. Args: task: The write task to schedule """ self.ensure_worker_started() + if self._is_transfer_terminal(task.transfer_id): + return False + + with self._write_state_lock: + if self._is_transfer_terminal(task.transfer_id): + return False + if task.layer_name in self._scheduled_layers[task.transfer_id]: + return False + self._scheduled_layers[task.transfer_id].add(task.layer_name) + self._scheduled_writes[task.transfer_id] += 1 self._write_task_q.put(task) + return True + + def is_scheduled(self, transfer_id: TransferId, layer_name: str) -> bool: + with self._write_state_lock: + return layer_name in self._scheduled_layers.get(transfer_id, set()) + + def seal_pending_transfers(self) -> None: + """Seal expected WRITE counts after the model forward has run. + + `save_kv_layer` is only invoked for attention layers whose backend uses + the standard KV connector hook. Hybrid models can register more KV + cache tensors than the number of hooks that fire in a forward, so WRITE + completion must be based on the tasks actually queued for the transfer. + """ + pending: list[tuple[TransferId, RemoteAllocInfo]] = [] + with self._write_state_lock: + for transfer_id, write_count in self._scheduled_writes.items(): + if transfer_id in self._sealed_writes: + continue + self._sealed_writes[transfer_id] = write_count + request_info = ( + self.worker.moriio_wrapper.done_remote_allocate_req_dict.get( + transfer_id + ) + ) + if request_info is not None: + request_info.writes_expected = write_count + pending.append((transfer_id, request_info)) + + for transfer_id, request_info in pending: + self._finalize_if_complete(transfer_id, request_info) def _write_worker_loop(self) -> None: """Main loop for the write worker thread.""" @@ -128,6 +188,9 @@ class MoRIIOWriter: except Empty: continue + if self._is_transfer_terminal(task.transfer_id): + continue + # Check if remote blocks are ready if not self._is_remote_ready(task): # task.retry_count += 1 @@ -158,6 +221,8 @@ class MoRIIOWriter: still_deferred: list[WriteTask] = [] for task in self._deferred_tasks: + if self._is_transfer_terminal(task.transfer_id): + continue if now - task.enqueue_time > defer_timeout: logger.error( "Deferred write task for request %s expired after %.1fs " @@ -181,12 +246,25 @@ class MoRIIOWriter: self._deferred_tasks = still_deferred + def _clear_transfer_state(self, transfer_id: TransferId) -> None: + with self._write_state_lock: + self._scheduled_writes.pop(transfer_id, None) + self._scheduled_layers.pop(transfer_id, None) + self._sealed_writes.pop(transfer_id, None) + + def _is_transfer_terminal(self, transfer_id: TransferId) -> bool: + wrapper = self.worker.moriio_wrapper + with wrapper.lock: + return wrapper._is_transfer_terminal_locked(transfer_id) + def _mark_request_done(self, transfer_id: str) -> None: """Mark a request done so its blocks are freed, even on transfer failure.""" wrapper = self.worker.moriio_wrapper with wrapper.lock: wrapper.done_req_ids.append(transfer_id) - wrapper.done_remote_allocate_req_dict.pop(transfer_id, None) + wrapper.done_remote_allocate_req_dict.pop(transfer_id, None) + wrapper._mark_transfer_terminal_locked(transfer_id) + self._clear_transfer_state(transfer_id) def _is_remote_ready(self, task: WriteTask) -> bool: """Check if remote blocks are allocated for this task. @@ -229,6 +307,12 @@ class MoRIIOWriter: """ # Get remote allocation info request_info = self._get_remote_alloc_info(task.transfer_id) + with self._write_state_lock: + request_info.completion_request_id = task.request_id + request_info.completion_remote_notify_port = task.remote_notify_port + request_info.completion_remote_ip = task.remote_ip + if task.transfer_id in self._sealed_writes: + request_info.writes_expected = self._sealed_writes[task.transfer_id] if request_info.block_ids is None: logger.debug( @@ -259,10 +343,12 @@ class MoRIIOWriter: plan = self._prepare_transfer_plan(task, request_info, remote_moriio_meta) # Execute transfer - self._do_layer_write(plan, sessions) + transfer_statuses = self._do_layer_write(plan, sessions) + with self._write_state_lock: + request_info.transfer_statuses.extend(transfer_statuses) # Finalize if all layers complete - self._finalize_if_complete(task, request_info) + self._mark_write_done(task.transfer_id, request_info) def _prepare_transfer_plan( self, @@ -279,21 +365,23 @@ class MoRIIOWriter: Returns: The transfer plan """ - # Compute offsets if not cached - if request_info.transfer_offset is None: + layer_cache = self.worker.kv_caches[task.layer_name] + geometry_key = _get_write_geometry_key(layer_cache) + offsets = request_info.transfer_offsets.get(geometry_key) + if offsets is None: offsets = self.worker._compute_block_transfer_offsets( task.layer_name, task.local_block_ids, request_info.block_ids, remote_moriio_meta, ) - request_info.transfer_offset = offsets + request_info.transfer_offsets[geometry_key] = offsets # Get session index layer_names = list(self.worker.layer_name_to_local_kv_cache_metadata.keys()) sess_idx = layer_names.index(task.layer_name) - local_off, remote_off, sizes = request_info.transfer_offset + local_off, remote_off, sizes = offsets return LayerTransferPlan( request_id=task.request_id, @@ -306,7 +394,7 @@ class MoRIIOWriter: use_batch=True, ) - def _do_layer_write(self, plan: LayerTransferPlan, sessions: list) -> None: + def _do_layer_write(self, plan: LayerTransferPlan, sessions: list) -> list[Any]: """Perform the actual layer write. Args: @@ -314,59 +402,82 @@ class MoRIIOWriter: sessions: List of transfer sessions """ if plan.use_batch: - self.worker.moriio_wrapper.write_remote_data( - plan.transfer_sizes, - plan.transfer_local_offsets, - plan.transfer_remote_offsets, - sessions[plan.sess_idx], - ) - else: - for i in range(len(plan.transfer_local_offsets)): + return [ + self.worker.moriio_wrapper.write_remote_data( + plan.transfer_sizes, + plan.transfer_local_offsets, + plan.transfer_remote_offsets, + sessions[plan.sess_idx], + ) + ] + + transfer_statuses: list[Any] = [] + for i in range(len(plan.transfer_local_offsets)): + transfer_statuses.append( self.worker.moriio_wrapper.write_remote_data_single( plan.transfer_sizes[i], plan.transfer_local_offsets[i], plan.transfer_remote_offsets[i], plan.sess_idx, ) + ) + return transfer_statuses + + def _mark_write_done( + self, transfer_id: TransferId, request_info: RemoteAllocInfo + ) -> None: + """Record one completed WRITE task and finalize if sealed.""" + with self._write_state_lock: + request_info.writes_done += 1 + self._finalize_if_complete(transfer_id, request_info) def _finalize_if_complete( - self, task: WriteTask, request_info: RemoteAllocInfo + self, transfer_id: TransferId, request_info: RemoteAllocInfo ) -> None: - """Finalize transfer if all layers are complete. + """Finalize transfer if all scheduled writes are complete.""" + with self._write_state_lock: + expected = request_info.writes_expected + if expected is None or request_info.writes_done < expected: + return + if request_info.completion_notified: + return + request_id = request_info.completion_request_id + remote_notify_port = request_info.completion_remote_notify_port + remote_ip = request_info.completion_remote_ip + if request_id is None or remote_notify_port is None or remote_ip is None: + return + transfer_statuses = list(request_info.transfer_statuses) + request_info.transfer_statuses.clear() + request_info.completion_notified = True - Args: - task: The write task - request_info: Remote allocation information - """ - request_info.writes_done += 1 + # Wait for this request's transfers to complete. + self.worker.moriio_wrapper.waiting_for_transfer_complete(transfer_statuses) - if request_info.writes_done >= self.worker.num_layers: - # Wait for transfer to complete - self.worker.moriio_wrapper.waiting_for_transfer_complete() + remote_port = remote_notify_port + get_port_offset( + request_info.decode_dp_rank, self.worker.tp_rank + ) + # Consider using RDMA immediate data in decode side + # to eliminate the need for this notification. + # Consider including the first gen token from prefill in the notification - remote_port = task.remote_notify_port + get_port_offset( - request_info.decode_dp_rank, self.worker.tp_rank - ) - # Consider using RDMA immediate data in decode side - # to eliminate the need for this notification. - # Consider including the first gen token from prefill in the notification - - # Send completion notification - self.worker.moriio_wrapper.send_notify( - task.transfer_id, task.remote_ip, remote_port - ) - # mark request as done, then we can free the blocks - with self.worker.moriio_wrapper.lock: - self.worker.moriio_wrapper.done_req_ids.append(task.transfer_id) - del self.worker.moriio_wrapper.done_remote_allocate_req_dict[ - task.transfer_id - ] - logger.debug( - "Completed transfer for (request, transfer) %s, %s, notified port %d", - task.request_id, - task.transfer_id, - remote_port, + # Send completion notification + self.worker.moriio_wrapper.send_notify( + transfer_id, remote_ip, remote_port, message_type="write_done" + ) + # mark request as done, then we can free the blocks + with self.worker.moriio_wrapper.lock: + self.worker.moriio_wrapper.done_req_ids.append(transfer_id) + self.worker.moriio_wrapper.done_remote_allocate_req_dict.pop( + transfer_id, None ) + self.worker.moriio_wrapper._mark_transfer_terminal_locked(transfer_id) + self._clear_transfer_state(transfer_id) + logger.debug( + "Completed transfer for (request, transfer) %s, %s, notified port %d", + request_id, + transfer_id, + remote_port, + ) class MoRIIOWrapper: @@ -400,6 +511,7 @@ class MoRIIOWrapper: self.done_req_ids: list[str] = [] self.done_remote_allocate_req_dict: dict[TransferId, RemoteAllocInfo] = {} self.done_write_cache_req_ids: list[str] = [] + self._terminal_transfer_ids: OrderedDict[TransferId, None] = OrderedDict() self._transfer_timeout = transfer_timeout self.notify_thread: threading.Thread | None = None self.sessions: list[IOEngine.Session] = [] @@ -506,8 +618,7 @@ class MoRIIOWrapper: transfer_status = session.batch_write( local_offset, remote_offset, transfer_size_byte, write_uid ) - with self.lock: - self.transfer_status.append(transfer_status) + return transfer_status def write_remote_data_single( self, transfer_size_byte, local_offset=0, remote_offset=0, sess_idx=0 @@ -520,17 +631,19 @@ class MoRIIOWrapper: transfer_size_byte, self.moriio_engine.allocate_transfer_uid(), ) - with self.lock: - self.transfer_status.append(transfer_status) + return transfer_status - def waiting_for_transfer_complete(self): - if not self.transfer_status: + def waiting_for_transfer_complete(self, transfer_statuses: list[Any] | None = None): + if transfer_statuses is None: + with self.lock: + transfers_to_wait = self.transfer_status[:] + self.transfer_status.clear() + else: + transfers_to_wait = list(transfer_statuses) + + if not transfers_to_wait: return - with self.lock: - transfers_to_wait = self.transfer_status[:] - self.transfer_status.clear() - timeout = self._transfer_timeout deadline = time.monotonic() + timeout remaining = list(transfers_to_wait) @@ -598,49 +711,105 @@ class MoRIIOWrapper: # [read] mode: receives block release messages from decode side # Decode Role: # [write] mode: receives KV cache write completion notifications + msg_str = repr(msg) handled = False try: data = msgpack.loads(msg) - if isinstance(data, dict) and "req_id" in data: + if isinstance(data, dict): self._handle_structured_message(data) - return - except (msgpack.exceptions.ExtraData, msgpack.exceptions.UnpackException): + except ( + msgpack.exceptions.ExtraData, + msgpack.exceptions.UnpackException, + ValueError, + ): logger.debug("Failed to decode msgpack message, will try as string") pass try: msg_str = msg.decode("UTF-8") - if msg_str.startswith(MoRIIOConstants.TRANSFER_PREFIX): + if msg_str: self._handle_completion_message(msg_str) handled = True except UnicodeDecodeError: - logger.warning("Received non-UTF8 message: %s", msg_str) + logger.warning("Received non-UTF8 message: %r", msg) if not handled: raise MoRIIOError(f"Unhandled message format: {msg_str}") def _handle_structured_message(self, data: dict): + message_type = data.get("type") + if message_type is None and "req_id" in data: + message_type = "remote_blocks" + + if message_type == "remote_blocks": + self._handle_remote_blocks_message(data) + elif message_type == "write_done": + self._handle_write_done_message(data) + elif message_type == "release": + self._handle_release_message(data) + else: + raise MoRIIOError(f"Unhandled structured message type: {message_type}") + + def _handle_remote_blocks_message(self, data: dict): assert get_role() == ROLE.PRODUCER, "Only prefill can get block messages" transfer_id = data["transfer_id"] block_notify_list = data.get("block_notify_list", []) decode_dp_rank = data.get("decode_rank", 0) - assert len(block_notify_list) > 0, ( - "block_notify_list cannot be empty in remote allocate message" - ) + if not block_notify_list: + raise MoRIIOError( + "block_notify_list cannot be empty in remote allocate message" + ) with self.lock: + if self._is_transfer_terminal_locked(transfer_id): + logger.debug( + "Ignoring remote allocation for terminal transfer %s", + transfer_id, + ) + return self.done_remote_allocate_req_dict[transfer_id] = RemoteAllocInfo( block_ids=block_notify_list, decode_dp_rank=decode_dp_rank ) + def _handle_write_done_message(self, data: dict): + assert get_role() != ROLE.PRODUCER, ( + "Only decode can get WRITE completion messages" + ) + transfer_id = data["transfer_id"] + with self.lock: + self.done_write_cache_req_ids.append(transfer_id) + + def _handle_release_message(self, data: dict): + assert get_role() == ROLE.PRODUCER, ( + "Only prefill can get transfer release messages" + ) + transfer_id = data["transfer_id"] + with self.lock: + self.done_req_ids.append(transfer_id) + self.done_remote_allocate_req_dict.pop(transfer_id, None) + self._mark_transfer_terminal_locked(transfer_id) + def _handle_completion_message(self, msg: str): with self.lock: if get_role() == ROLE.PRODUCER: self.done_req_ids.append(msg) + self.done_remote_allocate_req_dict.pop(msg, None) + self._mark_transfer_terminal_locked(msg) else: self.done_write_cache_req_ids.append(msg) - def send_notify(self, req_ids, remote_ip, remote_port): + def _is_transfer_terminal_locked(self, transfer_id: TransferId) -> bool: + return transfer_id in self._terminal_transfer_ids + + def _mark_transfer_terminal_locked(self, transfer_id: TransferId) -> None: + self._terminal_transfer_ids[transfer_id] = None + self._terminal_transfer_ids.move_to_end(transfer_id) + while len(self._terminal_transfer_ids) > _MAX_TERMINAL_TRANSFER_IDS: + self._terminal_transfer_ids.popitem(last=False) + + def send_notify( + self, req_ids, remote_ip, remote_port, message_type: str | None = None + ): if not remote_ip or not remote_port: logger.warning("Missing remote_ip or remote_port for notification") return @@ -664,7 +833,12 @@ class MoRIIOWrapper: "Invalid req_id type: %s, expected str", type(req_id) ) continue - sock.send(req_id.encode("utf-8")) + if message_type is None: + sock.send(req_id.encode("utf-8")) + else: + sock.send( + msgpack.dumps({"type": message_type, "transfer_id": req_id}) + ) except Exception as e: logger.error("Failed to send notification to %s: %s", path, e) self.paths.pop(path, None) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py index 8a6aced9daa..39dbec3886c 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/moriio/moriio_layout.py @@ -56,6 +56,42 @@ def is_mla_cache_layer( return isinstance(spec, (MLAAttentionSpec, SlidingWindowMLASpec)) +def _spec_dim_matches(value: int, expected: int | None) -> bool: + return expected is None or value == expected + + +def _kernel_layout_matches( + spec: KVCacheSpec, kernel_block_size: int, num_kv_heads: int, head_dim: int +) -> bool: + if kernel_block_size <= 0 or spec.block_size % kernel_block_size != 0: + return False + return _spec_dim_matches( + num_kv_heads, getattr(spec, "num_kv_heads", None) + ) and _spec_dim_matches(head_dim, getattr(spec, "head_size", None)) + + +def _select_kernel_block_layout( + layer_name: str, shape: torch.Size, spec: KVCacheSpec +) -> tuple[int, int, int]: + axis2_matches = _kernel_layout_matches(spec, shape[2], shape[3], shape[4]) + axis3_matches = _kernel_layout_matches(spec, shape[3], shape[2], shape[4]) + + if axis2_matches and axis3_matches and shape[2] != shape[3]: + raise ValueError( + f"Ambiguous MoRIIO kernel-block K/V cache shape for layer " + f"{layer_name}: {tuple(shape)}" + ) + if axis2_matches: + return shape[2], shape[3], shape[4] + if axis3_matches: + return shape[3], shape[2], shape[4] + + raise ValueError( + f"Unsupported MoRIIO K/V cache shape for layer {layer_name}: " + f"{tuple(shape)} does not contain block size {spec.block_size}" + ) + + def get_layer_transfer_geometry( layer_name: str, kv_cache: torch.Tensor, @@ -65,6 +101,7 @@ def get_layer_transfer_geometry( shape = kv_cache.shape stride = kv_cache.stride() element_size = kv_cache.element_size() + spec = layer_to_spec[layer_name] is_mla_cache = is_mla_cache_layer(layer_to_spec, layer_name) if is_mla_cache and len(shape) == 3: @@ -85,25 +122,92 @@ def get_layer_transfer_geometry( ) if not is_mla_cache and len(shape) == 5 and shape[0] == 2: - _, num_blocks, block_size, num_kv_heads, head_dim = shape + _, num_blocks = shape[:2] + kernel_blocks_per_block = 1 + if shape[2] == spec.block_size: + block_size, num_kv_heads, head_dim = shape[2:] + elif shape[3] == spec.block_size: + num_kv_heads, block_size, head_dim = shape[2:] + else: + kernel_num_blocks = num_blocks + kernel_block_size, num_kv_heads, head_dim = _select_kernel_block_layout( + layer_name, shape, spec + ) + kernel_blocks_per_block = spec.block_size // kernel_block_size + if kernel_num_blocks % kernel_blocks_per_block != 0: + raise ValueError( + f"Unsupported MoRIIO K/V cache shape for layer {layer_name}: " + f"{tuple(shape)} has {kernel_num_blocks} kernel blocks, " + f"not divisible by {kernel_blocks_per_block}" + ) + num_blocks = kernel_num_blocks // kernel_blocks_per_block + block_size = spec.block_size slot_size_bytes = num_kv_heads * head_dim * element_size block_len = block_size * slot_size_bytes - remote_kv_stride = stride[1] * (remote_num_blocks or num_blocks) return LayerTransferGeometry( num_blocks=num_blocks, block_size=block_size, block_len=block_len, slot_size_bytes=slot_size_bytes, - block_stride=stride[1], + block_stride=stride[1] * kernel_blocks_per_block, local_kv_stride=stride[0], - remote_kv_stride=remote_kv_stride, + remote_kv_stride=( + stride[1] * kernel_blocks_per_block * (remote_num_blocks or num_blocks) + ), transfers_per_block=2, regions_per_block=1, split_kv_regions=True, ) if not is_mla_cache and len(shape) == 5 and shape[1] == 2: - num_blocks, _, block_size, num_kv_heads, head_dim = shape + num_blocks = shape[0] + if shape[2] == spec.block_size: + block_size, num_kv_heads, head_dim = shape[2:] + slot_size_bytes = num_kv_heads * head_dim * element_size + block_len = block_size * slot_size_bytes + return LayerTransferGeometry( + num_blocks=num_blocks, + block_size=block_size, + block_len=block_len, + slot_size_bytes=slot_size_bytes, + block_stride=stride[0], + local_kv_stride=stride[1], + remote_kv_stride=stride[1], + transfers_per_block=2, + regions_per_block=2, + split_kv_regions=False, + ) + elif shape[3] == spec.block_size: + num_kv_heads, block_size, head_dim = shape[2:] + else: + kernel_num_blocks = num_blocks + kernel_block_size, _, _ = _select_kernel_block_layout( + layer_name, shape, spec + ) + kernel_blocks_per_block = spec.block_size // kernel_block_size + if kernel_num_blocks % kernel_blocks_per_block != 0: + raise ValueError( + f"Unsupported MoRIIO K/V cache shape for layer {layer_name}: " + f"{tuple(shape)} has {kernel_num_blocks} kernel blocks, " + f"not divisible by {kernel_blocks_per_block}" + ) + num_blocks = kernel_num_blocks // kernel_blocks_per_block + block_size = spec.block_size + block_stride = stride[0] * kernel_blocks_per_block + block_len = block_stride * element_size + slot_size_bytes = block_len // block_size + return LayerTransferGeometry( + num_blocks=num_blocks, + block_size=block_size, + block_len=block_len, + slot_size_bytes=slot_size_bytes, + block_stride=block_stride, + local_kv_stride=None, + remote_kv_stride=None, + transfers_per_block=1, + regions_per_block=1, + split_kv_regions=False, + ) slot_size_bytes = num_kv_heads * head_dim * element_size block_len = block_size * slot_size_bytes return LayerTransferGeometry(