[ROCm][P/D] Fix MoRIIO WRITE mode for mixed KV layouts (#46290)

Signed-off-by: Tan Pin Siang <tanpinsiang@gmail.com>
Co-authored-by: vllmellm <vllm.ellm@embeddedllm.com>
Co-authored-by: Hongxia Yang <hongxia.yang@amd.com>
Co-authored-by: Jun Kang Chow <junkangchow@gmail.com>
Co-authored-by: Chun Fang <chun.fang@amd.com>
Co-authored-by: TianDi101 <ditian12@amd.com>
Co-authored-by: functionstackx <47992694+functionstackx@users.noreply.github.com>
This commit is contained in:
Tan Pin Siang
2026-06-23 12:12:51 +08:00
committed by GitHub
co-authored by vllmellm Hongxia Yang Jun Kang Chow Chun Fang TianDi101 functionstackx
parent a8481be7a9
commit 7e47fb72b5
6 changed files with 1104 additions and 195 deletions
@@ -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"
)
@@ -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()})
@@ -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),
)
@@ -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",
@@ -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)
@@ -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(