forked from Karylab-cklius/vllm
[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:
co-authored by
vllmellm
Hongxia Yang
Jun Kang Chow
Chun Fang
TianDi101
functionstackx
parent
a8481be7a9
commit
7e47fb72b5
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user