Files
enginex-ascend-910-vllm/tests/ut/kv_offload/test_mooncake_connector.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

3173 lines
133 KiB
Python

import os
import queue
import socket
import sys
import threading
import time
import types
import unittest
from collections import OrderedDict, defaultdict, deque
from typing import Any, cast
from unittest.mock import MagicMock, patch
import msgspec
import torch
import zmq
from vllm.utils.network_utils import make_zmq_path
from vllm.v1.kv_cache_interface import FullAttentionSpec, MLAAttentionSpec, UniformTypeKVCacheSpecs
from vllm.v1.request import RequestStatus
fake_engine = types.ModuleType("mooncake.engine")
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
sys.modules["mooncake.engine"] = fake_engine
# Clean up stale mock modules installed by other test files
# (e.g., ascend_store/_mock_deps.py) that replace real kv_transfer
# subpackages with MagicMock/fake modules, breaking our imports.
# Save and restore so other test files (ascend_store) still see their mocks.
_kv_xfer = "vllm_ascend.distributed.kv_transfer"
_vllm_kv_xfer = "vllm.distributed.kv_transfer"
_saved_modules: dict[str, types.ModuleType] = {}
_to_remove = []
for k in list(sys.modules):
if k.startswith(_kv_xfer):
suffix = k[len(_kv_xfer) :]
if suffix == "" or suffix.startswith(".utils") or suffix.startswith(".kv_p2p"):
_to_remove.append(k)
elif k.startswith(_vllm_kv_xfer):
_to_remove.append(k)
for _m in _to_remove:
_saved_modules[_m] = sys.modules.pop(_m)
_mock_ascend_config = MagicMock(enable_kv_nz=False)
_mock_pp_group = MagicMock(rank_in_group=0, world_size=1)
_mock_tp_group = MagicMock(rank_in_group=0, world_size=4)
_mock_pcp_group = MagicMock(rank_in_group=0, world_size=1)
_mock_dcp_group = MagicMock(rank_in_group=0, world_size=1)
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_pp_group", return_value=_mock_pp_group).start()
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_tp_group", return_value=_mock_tp_group).start()
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_tensor_model_parallel_world_size", return_value=4
).start()
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_tensor_model_parallel_rank", return_value=0
).start()
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_pcp_group", return_value=_mock_pcp_group
).start()
patch("vllm.distributed.parallel_state._DCP", _mock_dcp_group).start()
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector import ( # noqa: E402
MAX_REQUESTS_PER_PEER_HANDLER,
GroupPull,
KVCacheRecvingThread,
KVCacheSendingThread,
KVCacheTaskTracker,
KVConnectorRole,
MooncakeAgentMetadata,
MooncakeConnector,
MooncakeConnectorMetadata,
MooncakeConnectorScheduler,
MooncakeConnectorWorker,
ReqMeta,
ensure_zmq_recv,
ensure_zmq_send,
group_concurrent_contiguous,
split_if_not_byte_contiguous,
string_to_int64_hash,
zmq_ctx,
)
for _k, _v in _saved_modules.items():
sys.modules[_k] = _v
GET_META_MSG = b"get_meta_msg"
DONE_RECVING_MSG = b"done_recving_msg"
def make_agent_metadata(**overrides: Any) -> MooncakeAgentMetadata:
metadata: dict[str, Any] = {
"engine_id": "engine1",
"te_rpc_port": 9090,
"kv_group2layeridx": {0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0])},
"block_size": 16,
"kv_caches_base_addr": [[12345678]],
"block_size_scale": [[1]],
"num_blocks": 2,
"block_lens": [[1024]],
"block_strides": [[1024]],
}
metadata.update(overrides)
return MooncakeAgentMetadata(**metadata)
class TestKVCacheTaskTrackerInit(unittest.TestCase):
def test_init_basic_properties(self):
tracker = KVCacheTaskTracker()
self.assertIsInstance(tracker.done_task_lock, type(threading.Lock()))
self.assertIsInstance(tracker.finished_requests, set)
self.assertIsInstance(tracker.delayed_free_requests, OrderedDict)
class TestGetAndClearFinishedSingleRequests(unittest.TestCase):
def setUp(self):
self.tracker = KVCacheTaskTracker()
self.tracker.finished_requests = set()
self.tracker.done_task_lock = threading.Lock()
def test_empty_requests(self):
result = self.tracker.get_and_clear_finished_requests()
self.assertEqual(result, set())
self.assertEqual(len(self.tracker.finished_requests), 0)
def test_single_request(self):
self.tracker.finished_requests = {"req_123"}
result = self.tracker.get_and_clear_finished_requests()
self.assertEqual(result, {"req_123"})
self.assertEqual(len(self.tracker.finished_requests), 0)
def test_multiple_requests(self):
self.tracker.finished_requests = {"req_1", "req_2", "req_3"}
result = self.tracker.get_and_clear_finished_requests()
self.assertSetEqual(result, {"req_1", "req_2", "req_3"})
self.assertEqual(len(self.tracker.finished_requests), 0)
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
def test_concurrent_access(self, mock_logger):
from concurrent.futures import ThreadPoolExecutor
self.tracker.finished_requests = {"req_1", "req_2"}
with ThreadPoolExecutor(max_workers=3) as executor:
futures = [executor.submit(self.tracker.get_and_clear_finished_requests) for _ in range(3)]
results = [f.result() for f in futures]
self.assertEqual(sum(1 for r in results if r), 1)
self.assertEqual(len(self.tracker.finished_requests), 0)
class TestKVCacheSendingThreadInit(unittest.TestCase):
def setUp(self):
kv_caches: dict[str, Any] = {}
self.common_args: dict[str, Any] = {
"tp_rank": 1,
"prefill_tp_size": 4,
"local_engine_id": "engine_1",
"side_channel_host": "localhost",
"side_channel_port": 5555,
"metadata": MagicMock(),
"vllm_config": MockVllmConfig(),
"ready_event": threading.Event(),
"kv_caches": kv_caches,
"pcp_rank": 0,
}
self.threads = []
def tearDown(self):
for thread in self.threads:
if hasattr(thread, "task_tracker") and hasattr(thread.task_tracker, "socket"):
thread.task_tracker.socket.close()
if hasattr(thread, "is_alive") and thread.is_alive():
thread.join(timeout=0.1)
def test_thread_daemon_property(self):
thread = KVCacheSendingThread(**self.common_args)
self.threads.append(thread)
self.assertTrue(thread.daemon)
def test_thread_name_format(self):
thread = KVCacheSendingThread(**self.common_args)
self.threads.append(thread)
self.assertEqual(thread.name, "KVCacheSendingThread")
def test_ready_event_reference(self):
custom_event = threading.Event()
args = self.common_args.copy()
args["ready_event"] = custom_event
thread = KVCacheSendingThread(**args)
self.threads.append(thread)
self.assertIs(thread.ready_event, custom_event)
class TestGetAndClearFinishedRequests(unittest.TestCase):
def setUp(self):
kv_caches: dict[str, Any] = {}
self.common_args: dict[str, Any] = {
"tp_rank": 1,
"prefill_tp_size": 4,
"local_engine_id": "engine_1",
"side_channel_host": "localhost",
"vllm_config": MockVllmConfig(),
"side_channel_port": 5555,
"metadata": {"test": "metadata"},
"ready_event": threading.Event(),
"kv_caches": kv_caches,
"pcp_rank": 0,
}
self.thread = KVCacheSendingThread(**self.common_args)
@patch.object(KVCacheTaskTracker, "get_and_clear_finished_requests")
def test_get_and_clear_finished_requests(self, mock_get_clear):
expected_requests = {"req1", "req2"}
mock_get_clear.return_value = expected_requests
result = self.thread.get_and_clear_finished_requests()
mock_get_clear.assert_called_once()
self.assertEqual(result, expected_requests)
class TestKVCacheSendingThread(unittest.TestCase):
def test_run_handles_get_meta_and_done_recv_msgs(self):
ready_event = threading.Event()
metadata = make_agent_metadata(
engine_id="engine1",
kv_caches_base_addr=[[12345678]],
num_blocks=2,
)
vllm_config = MockVllmConfig()
host = "127.0.0.1"
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("", 0))
base_port = s.getsockname()[1]
thread = KVCacheSendingThread(
tp_rank=0,
prefill_tp_size=1,
local_engine_id="engine1",
side_channel_host=host,
side_channel_port=base_port,
metadata=metadata,
vllm_config=vllm_config,
ready_event=ready_event,
kv_caches={},
pcp_rank=0,
)
thread.start()
actual_port = base_port + (
thread.pp_rank * thread.tp_size + thread.tp_rank + thread.pcp_rank * thread.prefill_tp_size
)
self.assertTrue(ready_event.wait(timeout=3), "Server thread startup timeout")
context = zmq.Context() # type: ignore
sock = context.socket(zmq.DEALER) # type: ignore
sock.connect(f"tcp://{host}:{actual_port}")
encoder = msgspec.msgpack.Encoder()
decoder = msgspec.msgpack.Decoder(type=MooncakeAgentMetadata)
sock.send_multipart([b"", encoder.encode((GET_META_MSG,))])
frames = sock.recv_multipart()
self.assertEqual(frames[0], b"")
meta = decoder.decode(frames[1])
self.assertEqual(meta.engine_id, "engine1")
self.assertEqual(meta.kv_caches_base_addr, [[12345678]])
self.assertEqual(meta.num_blocks, 2)
req_id = "request_42"
thread.task_tracker.add_req_to_process(req_id)
sock.send_multipart([b"", encoder.encode((DONE_RECVING_MSG, req_id, 0))])
frames = sock.recv_multipart()
self.assertEqual(frames[0], b"")
self.assertEqual(frames[1], b"ACK")
self.assertIn(req_id, thread.task_tracker.finished_requests)
sock.close()
context.term()
def test_reformat_kv_cache_hybrid_linear_uses_cache_block_size(self):
block_size = 4
num_blocks = 2
tp_num_need_pulls = 2
feature_size = 3
transferred = torch.arange(
num_blocks * tp_num_need_pulls * block_size * feature_size,
dtype=torch.float32,
).reshape(num_blocks, tp_num_need_pulls, block_size, feature_size)
cache = transferred.reshape(num_blocks, block_size, tp_num_need_pulls * feature_size).clone()
expected = transferred.transpose(1, 2).contiguous().reshape_as(cache)
thread = KVCacheRecvingThread.__new__(KVCacheRecvingThread)
thread.kv_caches = {"layer.0": (cache, cache.clone())}
group_kv_caches = {"layer.0": (cache.clone(), cache.clone())}
thread.reformat_kv_cache_hybrid_linear_torch(
[[0, 1]],
tp_num_need_pulls,
group_kv_caches,
)
reformatted_k_cache, reformatted_v_cache = group_kv_caches["layer.0"]
torch.testing.assert_close(reformatted_k_cache, expected)
torch.testing.assert_close(reformatted_v_cache, expected)
class TestMooncakeTransferGroups(unittest.TestCase):
def test_attention_group_uses_explicit_total_heads_for_unequal_pd_tp(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.num_key_value_heads = 16
mla_group = {
"kv_cache_spec_type": "AscendMLAAttentionSpec",
"kv_cache_spec": {"num_kv_heads": 1, "total_num_kv_heads": 1},
}
full_attention_decode_group = {
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_spec": {"num_kv_heads": 4, "total_num_kv_heads": 8},
}
replicated_prefill_group = {
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_spec": {"num_kv_heads": 1, "total_num_kv_heads": 8},
}
self.assertEqual(worker._get_attention_group_num_key_value_heads(mla_group), 1)
self.assertEqual(
worker._get_attention_group_num_key_value_heads(full_attention_decode_group),
8,
)
self.assertEqual(
worker._get_attention_group_num_key_value_heads(replicated_prefill_group),
8,
)
self.assertEqual(
worker._get_attention_group_num_need_pulls_for_decode_tp(full_attention_decode_group, 8, 2),
4,
)
def test_build_kv_group2layeridx_splits_uniform_group_by_kv_heads(self):
mla_spec = MLAAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=64,
dtype=torch.float16,
)
qga_spec = FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=64,
head_size_v=64,
dtype=torch.float16,
)
layer_specs = {
"language_model.model.layers.0.self_attn": mla_spec,
"model.layers.32.self_attn": qga_spec,
}
uniform_spec = UniformTypeKVCacheSpecs(block_size=16, kv_cache_specs=layer_specs)
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.vllm_config = MockVllmConfig()
worker.vllm_config.model_config.hf_text_config.num_key_value_heads = 128
worker.vllm_config.model_config.get_total_num_kv_heads = MagicMock(return_value=128)
worker.vllm_config.speculative_config = types.SimpleNamespace(
draft_model_config=types.SimpleNamespace(
hf_text_config=types.SimpleNamespace(num_key_value_heads=8),
get_total_num_kv_heads=MagicMock(return_value=8),
),
)
worker.total_layers = 32
worker.kv_cache_config = MockKVCacheConfig(
kv_cache_groups=[
MockKVCacheGroup(
layer_names=list(layer_specs),
kv_cache_spec=uniform_spec,
)
]
)
kv_group2layeridx = worker._build_kv_group2layeridx()
self.assertEqual(len(kv_group2layeridx), 2)
self.assertEqual(kv_group2layeridx[0][0]["kv_cache_group_id"], 0)
self.assertEqual(kv_group2layeridx[1][0]["kv_cache_group_id"], 0)
self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[0][0]), 1)
self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[1][0]), 8)
self.assertEqual(kv_group2layeridx[0][0]["kv_cache_spec"]["num_kv_heads"], 1)
self.assertEqual(kv_group2layeridx[1][0]["kv_cache_spec"]["num_kv_heads"], 1)
def test_build_kv_group2layeridx_splits_equal_local_heads_by_total_heads(self):
shared_local_spec = FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=64,
head_size_v=64,
dtype=torch.float16,
)
layer_names = [
"model.layers.0.self_attn",
"eagle.model.layers.0.self_attn",
]
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.vllm_config = MockVllmConfig()
worker.vllm_config.model_config.hf_text_config.num_key_value_heads = 16
worker.vllm_config.model_config.get_total_num_kv_heads = MagicMock(return_value=16)
worker.vllm_config.speculative_config = types.SimpleNamespace(
draft_model_config=types.SimpleNamespace(
hf_text_config=types.SimpleNamespace(num_key_value_heads=8),
get_total_num_kv_heads=MagicMock(return_value=8),
),
)
worker.total_layers = 32
worker.kv_cache_config = MockKVCacheConfig(
kv_cache_groups=[
MockKVCacheGroup(
layer_names=layer_names,
kv_cache_spec=shared_local_spec,
)
]
)
kv_group2layeridx = worker._build_kv_group2layeridx()
self.assertEqual(len(kv_group2layeridx), 2)
self.assertEqual(kv_group2layeridx[0][1], [0])
self.assertEqual(kv_group2layeridx[1][1], [32])
self.assertEqual(kv_group2layeridx[0][0]["kv_cache_spec"]["total_num_kv_heads"], 16)
self.assertEqual(kv_group2layeridx[1][0]["kv_cache_spec"]["total_num_kv_heads"], 8)
self.assertEqual(kv_group2layeridx[0][0]["kv_cache_spec"]["num_kv_heads"], 1)
self.assertEqual(kv_group2layeridx[1][0]["kv_cache_spec"]["num_kv_heads"], 1)
worker.kv_group2layeridx = kv_group2layeridx
self.assertTrue(worker._requires_group_aware_attention_transfer())
worker.tp_rank = 0
worker.tp_size = 8
worker._decode_tp_size = 8
worker._prefill_tp_size = 16
worker._prefill_pp_size = 1
worker.use_sparse = False
_, rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls("req-1", prefill_tp_size=16)
pulls = [pull for group_pulls in rank_group_pulls.values() for pull in group_pulls]
target_pulls = [pull for pull in pulls if pull.group_id == 0]
draft_pulls = [pull for pull in pulls if pull.group_id == 1]
self.assertEqual(len(target_pulls), 2)
self.assertTrue(all(pull.num_group_pulls == 2 for pull in target_pulls))
self.assertEqual(len(draft_pulls), 1)
self.assertEqual(draft_pulls[0].num_group_pulls, 1)
def test_hybrid_rank_pulls_use_transfer_group_kv_heads(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.vllm_config = MockVllmConfig()
worker.vllm_config.model_config.is_deepseek_mla = True
worker.tp_rank = 0
worker.tp_size = 4
worker._decode_tp_size = 4
worker._prefill_tp_size = 8
worker._prefill_pp_size = 1
worker.num_key_value_heads = 128
worker.use_sparse = False
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
"kv_cache_group_id": 0,
"kv_cache_spec": {
"total_num_kv_heads": 1,
"model.layers.0.self_attn": {"num_kv_heads": 1},
},
},
[0],
),
1: (
{
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
"kv_cache_group_id": 0,
"kv_cache_spec": {
"total_num_kv_heads": 8,
"model.layers.1.self_attn": {"num_kv_heads": 8},
},
},
[1],
),
}
_, rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls("req-1", prefill_tp_size=8)
pulls = [pull for group_pulls in rank_group_pulls.values() for pull in group_pulls]
mla_pulls = [pull for pull in pulls if pull.group_id == 0]
qga_pulls = [pull for pull in pulls if pull.group_id == 1]
self.assertEqual(len(mla_pulls), 1)
self.assertEqual(mla_pulls[0].num_group_pulls, 1)
self.assertEqual(len(qga_pulls), 2)
self.assertTrue(all(pull.num_group_pulls == 2 for pull in qga_pulls))
def test_hybrid_group_pulls_metadata_filters_groups_per_remote_card(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.vllm_config = MockVllmConfig()
worker.vllm_config.model_config.is_deepseek_mla = True
worker._is_hma_required = True
worker.tp_rank = 0
worker.tp_size = 4
worker._decode_tp_size = 4
worker._prefill_tp_size = 8
worker._prefill_pp_size = 1
worker.num_key_value_heads = 128
worker.use_sparse = False
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_group_id": 0,
"kv_cache_spec": {"num_kv_heads": 1},
},
[0],
),
1: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_group_id": 0,
"kv_cache_spec": {"num_kv_heads": 8},
},
[1],
),
}
worker.pcp_size = 1
worker.dcp_size = 1
req_id = "req-1"
remote_base_port = 30000
chosen_rank_list, expected_rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls(
req_id, prefill_tp_size=8
)
remote_handshake_port_list = [[remote_base_port + rank for rank in chosen_rank_list]]
group_pulls_list = worker._get_group_pulls_metadata(
req_id,
remote_handshake_port_list,
prefill_tp_size=8,
remote_base_port=remote_base_port,
remote_pcp_size=1,
remote_dcp_size=1,
)
self.assertEqual(len(group_pulls_list), 1)
self.assertEqual(len(group_pulls_list[0]), len(chosen_rank_list))
group_ids_by_port = [[group_pull.group_id for group_pull in port_pulls] for port_pulls in group_pulls_list[0]]
expected_group_ids_by_port = [
[group_pull.group_id for group_pull in expected_rank_group_pulls[rank]] for rank in chosen_rank_list
]
self.assertEqual(group_ids_by_port, expected_group_ids_by_port)
self.assertTrue(any(set(group_ids) != {0, 1} for group_ids in group_ids_by_port))
self.assertFalse(all(set(group_ids) == {0, 1} for group_ids in group_ids_by_port))
class TestKVCacheRecvingThreadBasic(unittest.TestCase):
def setUp(self):
self.engine = MagicMock()
self.ready_event = threading.Event()
self.vllm_config = MockVllmConfig()
self.kv_caches: dict[str, Any] = {}
self.thread = KVCacheRecvingThread(
tp_rank=0,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000], [0x2000]],
block_len_per_addr=[[1024], [2048]],
block_stride_per_addr=[[1024], [2048]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches=self.kv_caches,
prefill_pp_layer_partition=None,
)
def test_add_request(self):
test_req: dict[str, Any] = {
"request_id": "req1",
"local_block_ids": [1, 2],
"remote_block_ids": [3, 4],
"remote_engine_id": "remote_engine",
"remote_host": "localhost",
"remote_handshake_port": 6666,
"offset": 0,
"tp_num_need_pulls": 2,
"all_task_done": False,
}
self.thread.add_request(
request_id=test_req["request_id"],
remote_request_id=test_req["request_id"],
local_block_ids=test_req["local_block_ids"],
remote_block_ids=test_req["remote_block_ids"],
group_pulls=[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)],
remote_engine_id=test_req["remote_engine_id"],
remote_host=test_req["remote_host"],
remote_handshake_port=test_req["remote_handshake_port"],
all_task_done=test_req["all_task_done"],
)
queued = self.thread.request_queue.get_nowait()
self.assertEqual(queued["request_id"], "req1")
self.assertEqual(queued["remote_host"], "localhost")
self.assertEqual(queued["num_computed_tokens"], 0)
def test_mark_and_is_failed(self):
self.thread._mark_failed_recv_request("req1", [[10, 20]])
self.assertTrue(self.thread._is_failed_recv_request("req1"))
self.assertIn(10, self.thread.invalid_block_ids)
self.assertIn(20, self.thread.invalid_block_ids)
def test_clear_failed_recv_request(self):
self.thread._mark_failed_recv_request("req2", [[30]])
self.thread._clear_failed_recv_request("req2")
self.assertFalse(self.thread._is_failed_recv_request("req2"))
def test_get_and_clear_invalid_block_ids(self):
self.thread.invalid_block_ids = {1, 2, 3}
result = self.thread.get_and_clear_invalid_block_ids()
self.assertSetEqual(result, {1, 2, 3})
self.assertEqual(self.thread.invalid_block_ids, set())
@patch.object(KVCacheTaskTracker, "get_and_clear_finished_requests")
def test_get_finished_requests(self, mock_tracker):
mock_tracker.return_value = {"req1", "req2"}
result = self.thread.get_and_clear_finished_requests()
self.assertEqual(result, {"req1", "req2"})
def test_executor_workers_bind_kv_cache_device_before_handling_requests(self):
expected_device = torch.device("npu:5")
kv_cache = MagicMock(device=expected_device)
worker_events: defaultdict[int, list[tuple[str, int | str]]] = defaultdict(list)
events_lock = threading.Lock()
both_workers_started = threading.Event()
release_workers = threading.Event()
def record_set_device(device):
device_index = device if isinstance(device, int) else torch.device(device).index
with events_lock:
worker_events[threading.get_ident()].append(("set_device", cast(int, device_index)))
with patch("torch.npu.set_device", side_effect=record_set_device):
thread = KVCacheRecvingThread(
tp_rank=1,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000]],
block_len_per_addr=[[1024]],
block_stride_per_addr=[[1024]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches={"layer.0": (kv_cache, kv_cache)},
prefill_pp_layer_partition=None,
)
def handle_request(req_meta: dict[str, Any]):
with events_lock:
worker_events[threading.get_ident()].append(("handle", req_meta["request_id"]))
handled_worker_count = sum(
any(event == "handle" for event, _ in events) for events in worker_events.values()
)
if handled_worker_count == 2:
both_workers_started.set()
release_workers.wait()
thread._handle_request = handle_request # type: ignore[method-assign]
try:
for index in range(2):
thread._submit_request(
{
"request_id": f"req-{index}",
"remote_host": f"host-{index}",
"remote_handshake_port": 6000 + index,
"all_task_done": True,
}
)
self.assertTrue(both_workers_started.wait(timeout=5.0), "executor did not start two workers")
finally:
release_workers.set()
thread.executor.shutdown(wait=True, cancel_futures=True)
handled_worker_events = [events for events in worker_events.values() if any(e == "handle" for e, _ in events)]
self.assertEqual(len(handled_worker_events), 2)
for events in handled_worker_events:
self.assertEqual(events[0], ("set_device", expected_device.index))
self.assertEqual(events[1][0], "handle")
def test_submit_request_serializes_same_peer_fifo(self):
release_first_request = threading.Event()
first_request_started = threading.Event()
other_peer_started = threading.Event()
handled_requests: list[str] = []
active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
max_active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
state_lock = threading.Lock()
def handle_request(req_meta: dict[str, Any]):
peer_key = (req_meta["remote_host"], req_meta["remote_handshake_port"])
with state_lock:
active_by_peer[peer_key] += 1
max_active_by_peer[peer_key] = max(max_active_by_peer[peer_key], active_by_peer[peer_key])
handled_requests.append(req_meta["request_id"])
if req_meta["request_id"] == "same-peer-1":
first_request_started.set()
self.assertTrue(release_first_request.wait(timeout=2.0))
elif req_meta["request_id"] == "other-peer-1":
other_peer_started.set()
time.sleep(0.01)
with state_lock:
active_by_peer[peer_key] -= 1
self.thread._handle_request = handle_request # type: ignore[method-assign]
same_peer_1 = {
"request_id": "same-peer-1",
"remote_host": "host-a",
"remote_handshake_port": 6000,
"all_task_done": False,
}
same_peer_2 = {
"request_id": "same-peer-2",
"remote_host": "host-a",
"remote_handshake_port": 6000,
"all_task_done": True,
}
other_peer = {
"request_id": "other-peer-1",
"remote_host": "host-b",
"remote_handshake_port": 6001,
"all_task_done": True,
}
try:
self.thread._submit_request(same_peer_1)
self.assertTrue(first_request_started.wait(timeout=1.0))
self.thread._submit_request(same_peer_2)
self.thread._submit_request(other_peer)
self.assertTrue(other_peer_started.wait(timeout=1.0))
time.sleep(0.05)
self.assertNotIn("same-peer-2", handled_requests)
finally:
release_first_request.set()
self.thread.executor.shutdown(wait=True, cancel_futures=True)
self.assertLess(handled_requests.index("same-peer-1"), handled_requests.index("same-peer-2"))
self.assertEqual(max_active_by_peer[("host-a", 6000)], 1)
self.assertEqual(max_active_by_peer[("host-b", 6001)], 1)
def test_peer_handler_yields_after_batch_limit(self):
peer_key = ("host-a", 6000)
requests = [
{
"request_id": f"req-{idx}",
"remote_host": peer_key[0],
"remote_handshake_port": peer_key[1],
}
for idx in range(MAX_REQUESTS_PER_PEER_HANDLER + 1)
]
handled_requests: list[str] = []
self.thread.peer_request_queues[peer_key].extend(requests)
self.thread.active_peer_request_handlers.add(peer_key)
self.thread.executor = MagicMock()
def handle_request(req_meta: dict[str, Any]):
handled_requests.append(req_meta["request_id"])
self.thread._handle_request = handle_request # type: ignore[method-assign]
self.thread._handle_peer_requests(peer_key)
self.assertEqual(handled_requests, [f"req-{idx}" for idx in range(MAX_REQUESTS_PER_PEER_HANDLER)])
self.assertEqual(
[req["request_id"] for req in self.thread.peer_request_queues[peer_key]],
[f"req-{MAX_REQUESTS_PER_PEER_HANDLER}"],
)
self.assertIn(peer_key, self.thread.active_peer_request_handlers)
self.thread.executor.submit.assert_called_once_with(self.thread._handle_peer_requests, peer_key)
class TestSocketManagement(unittest.TestCase):
def setUp(self):
self.engine = MagicMock()
self.ready_event = threading.Event()
self.vllm_config = MockVllmConfig()
self.kv_caches: dict[str, Any] = {}
self.thread = KVCacheRecvingThread(
tp_rank=0,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000], [0x2000]],
block_len_per_addr=[[1024], [2048]],
block_stride_per_addr=[[1024], [2048]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches=self.kv_caches,
prefill_pp_layer_partition=None,
)
self.thread.remote_sockets = defaultdict(deque)
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.zmq.Context")
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.make_zmq_socket")
def test_get_remote_socket(self, mock_make_socket, mock_context):
mock_sock = MagicMock()
mock_make_socket.return_value = mock_sock
test_host = "test_host"
test_port = 12345
sock = self.thread._get_remote_socket(test_host, test_port)
self.assertEqual(sock, mock_sock)
mock_make_socket.assert_called_once()
args, kwargs = mock_make_socket.call_args
self.assertEqual(kwargs.get("path"), "tcp://test_host:12345")
self.assertEqual(kwargs.get("socket_type"), zmq.REQ) # type: ignore
self.assertFalse(kwargs.get("bind", True))
mock_sock.setsockopt.assert_any_call(zmq.SNDTIMEO, int(self.thread.timeout * 1000)) # type: ignore
mock_sock.setsockopt.assert_any_call(zmq.RCVTIMEO, int(self.thread.timeout * 1000)) # type: ignore
def test_return_socket_to_pool(self):
mock_sock = MagicMock()
test_host = "test_host"
test_port = 12345
test_path = make_zmq_path("tcp", test_host, test_port)
self.thread._return_remote_socket(mock_sock, test_host, test_port)
self.assertEqual(len(self.thread.remote_sockets[test_path]), 1)
self.assertEqual(self.thread.remote_sockets[test_path][0], mock_sock)
class TestCoreFunctionality(unittest.TestCase):
def setUp(self):
self.engine = MagicMock()
self.ready_event = threading.Event()
self.mock_queue = MagicMock()
self.vllm_config = MockVllmConfig()
self.kv_caches: dict[str, Any] = {"layer_0": (MagicMock(), MagicMock())}
self.thread = KVCacheRecvingThread(
tp_rank=0,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000], [0x2000]],
block_len_per_addr=[[1024], [2048]],
block_stride_per_addr=[[1024], [2048]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches=self.kv_caches,
prefill_pp_layer_partition=None,
)
self.thread.request_queue = self.mock_queue
self.test_req = {
"request_id": "req1",
"remote_request_id": "req1",
"local_block_ids": [[1, 2]],
"remote_block_ids": [[3, 4]],
"group_pulls": [GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1, is_group_transfer_end=True)],
"remote_engine_id": "remote_engine",
"remote_host": "localhost",
"remote_handshake_port": 6666,
"remote_port_send_num": {6666: 1},
"all_task_done": True,
"remote_block_size": 16,
}
self.thread.kv_group2layeridx = {0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0])}
self.thread.group_compress_ratios = {0: 1}
self.thread.block_size_scale = [[1]]
self.thread.task_tracker = MagicMock()
self.engine.batch_transfer_sync_read.return_value = 0
self.thread.remote_te_port = {"remote_engine": {6666: 7777}}
self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[1024]]
@patch.object(KVCacheRecvingThread, "_transfer_kv_cache_all_groups")
@patch.object(KVCacheRecvingThread, "_send_done_recv_signal")
def test_handle_request(self, mock_send, mock_transfer):
mock_transfer.return_value = None
mock_send.return_value = None
self.thread._handle_request(self.test_req)
mock_transfer.assert_called_once_with(self.test_req)
mock_send.assert_called_once_with("req1", "localhost", 6666, {6666: 1})
cast(Any, self.thread.task_tracker).update_done_task_count.assert_called_once_with("req1")
self.mock_queue.task_done.assert_called_once()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_kv_cache(self, mock_get_meta):
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[1]]}
self.thread._transfer_kv_cache_all_groups(self.test_req)
self.engine.batch_transfer_sync_read.assert_called_once()
call_args, call_kwargs = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[0], "localhost:7777")
self.assertIsInstance(call_args[1], list)
self.assertIsInstance(call_args[2], list)
self.assertIsInstance(call_args[3], list)
self.assertEqual(len(call_args[1]), len(call_args[2]))
self.assertEqual(len(call_args[1]), len(call_args[3]))
mock_get_meta.assert_not_called()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_groups_contiguous_kernel_blocks(self, mock_get_meta):
# Kernel-level ids now arrive pre-expanded from _get_kv_split_metadata; the
# transfer stage only groups contiguous kernels and computes addresses.
req = dict(self.test_req)
req["local_block_ids"] = [[2, 3, 4]]
req["remote_block_ids"] = [[7, 8, 9]]
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.block_size_scale = [[2]]
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[2]]}
self.thread._transfer_kv_cache_all_groups(req)
call_args, _ = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[1], [0x1000 + 2 * 1024])
self.assertEqual(call_args[2], [0x3000 + 7 * 1024])
self.assertEqual(call_args[3], [3 * 1024])
mock_get_meta.assert_not_called()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_prefix_cache_offset_uses_compress_ratio(self, mock_get_meta):
req = dict(self.test_req)
req["local_block_ids"] = [[1, 2]]
req["remote_block_ids"] = [[3, 4]]
req["num_computed_tokens"] = 32
self.thread.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
"kv_cache_spec": {"layer_0": {"compress_ratio": "4"}},
},
[0],
)
}
self.thread.group_compress_ratios = {0: 4}
req["remote_block_ids"] = [[6, 7]]
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.block_size_scale = [[2]]
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[2]]}
self.thread._transfer_kv_cache_all_groups(req)
# compress_ratio / block_size_scale no longer affect the transfer stage:
# kernel-block expansion happens upstream in _get_kv_split_metadata, so the
# block ids [1, 2] / [6, 7] are consumed directly. The two contiguous kernel
# blocks are grouped into a single transfer starting at the first block id.
call_args, _ = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[1], [0x1000 + 1 * 1024])
self.assertEqual(call_args[2], [0x3000 + 6 * 1024])
self.assertEqual(call_args[3], [2 * 1024])
mock_get_meta.assert_not_called()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_prefix_cache_trims_remote_kernel_blocks(self, mock_get_meta):
# Kernel-block expansion/trimming now happens upstream in
# _get_kv_split_metadata, so the remote block ids arrive pre-expanded and
# the transfer stage consumes them directly. remote_block_size_scale is no
# longer applied here; the remote address is base + block_id * remote_block_stride.
req = dict(self.test_req)
req["local_block_ids"] = [[1, 2]]
req["remote_block_ids"] = [[3, 4]]
req["num_computed_tokens"] = 0
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.block_size_scale = [[1]]
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[2]]}
self.thread._transfer_kv_cache_all_groups(req)
call_args, _ = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[1], [0x1000 + 1024])
self.assertEqual(call_args[2], [0x3000 + 3 * 1024])
self.assertEqual(call_args[3], [2 * 1024])
mock_get_meta.assert_not_called()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_kv_cache_uses_block_stride_for_block_offsets(self, mock_get_meta):
req = dict(self.test_req)
req["local_block_ids"] = [[1, 2]]
req["remote_block_ids"] = [[3, 4]]
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[1]]}
self.thread.block_len_per_addr = [[1024]]
self.thread.block_stride_per_addr = [[2048]]
self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[4096]]
self.thread._transfer_kv_cache_all_groups(req)
call_args, _ = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[1], [0x1000 + 1 * 2048, 0x1000 + 2 * 2048])
self.assertEqual(call_args[2], [0x3000 + 3 * 4096, 0x3000 + 4 * 4096])
self.assertEqual(call_args[3], [1024, 1024])
mock_get_meta.assert_not_called()
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
def test_transfer_replicated_indexer_when_regular_kv_shard_is_empty(self, mock_get_meta):
req = dict(self.test_req)
req["local_block_ids"] = [[]]
req["remote_block_ids"] = [[]]
req["local_block_ids_replicate_k"] = ([4, 5],)
req["remote_block_ids_replicate_k"] = ([7, 8],)
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
mock_config.return_value.enable_kv_nz = False
self.thread.enable_sfa_dcp_replicated_indexer = True
self.thread.kv_caches_base_addr["local_engine"][5555] = [[0x1000, 0x2000]]
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000, 0x4000]]}
self.thread.block_size_scale = [[1, 2]]
self.thread.block_len_per_addr = [[1024, 2048]]
self.thread.block_stride_per_addr = [[1024, 2048]]
self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[4096, 8192]]
self.thread._transfer_kv_cache_all_groups(req)
call_args, _ = self.engine.batch_transfer_sync_read.call_args
self.assertEqual(call_args[1], [0x2000 + 4 * 2048])
self.assertEqual(call_args[2], [0x4000 + 7 * 8192])
self.assertEqual(call_args[3], [2 * 2048])
mock_get_meta.assert_not_called()
def test_append_mamba_transfer_meta_uses_block_stride_for_block_offsets(self):
src_list: list[int] = []
dst_list: list[int] = []
length_list: list[int] = []
self.thread._append_mamba_transfer_meta(
src_list,
dst_list,
length_list,
group_spec={"kv_cache_spec_type": "MambaSpec"},
src_layer_base_addr=[0x1000, 0x2000],
dst_layer_base_addr=[0x3000, 0x4000],
block_len=[100, 200],
block_stride=[128, 256],
remote_block_stride=[160, 512],
remote_block_id=3,
local_block_id=2,
tp_num_need_pulls=1,
remote_tp_offset=0,
)
self.assertEqual(src_list, [0x1000 + 2 * 128, 0x2000 + 2 * 256])
self.assertEqual(dst_list, [0x3000 + 3 * 160, 0x4000 + 3 * 512])
self.assertEqual(length_list, [100, 200])
def test_transfer_kv_cache_failure(self):
self.engine.batch_transfer_sync_read.return_value = -1
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000]]}
self.thread.remote_block_size_scale["remote_engine"] = {6666: [[1]]}
with self.assertRaises(RuntimeError):
self.thread._transfer_kv_cache_all_groups(self.test_req)
class TestMetadataHandling(unittest.TestCase):
def setUp(self):
self.engine = MagicMock()
self.ready_event = threading.Event()
self.vllm_config = MockVllmConfig()
self.kv_caches: dict[str, Any] = {}
self.thread = KVCacheRecvingThread(
tp_rank=0,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000], [0x2000]],
block_len_per_addr=[[1024], [2048]],
block_stride_per_addr=[[1024], [2048]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches=self.kv_caches,
prefill_pp_layer_partition=None,
)
self.test_metadata = make_agent_metadata(
engine_id="remote_engine", te_rpc_port=9090, kv_caches_base_addr=[[0x3000], [0x4000]], num_blocks=2
)
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.ensure_zmq_send")
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.ensure_zmq_recv")
def test_get_remote_metadata_success(self, mock_recv, mock_send):
mock_recv.return_value = msgspec.msgpack.encode(self.test_metadata)
with (
patch.object(self.thread, "_get_remote_socket") as mock_get_socket,
patch.object(self.thread, "_return_remote_socket") as mock_return_socket,
):
mock_socket = MagicMock(spec=zmq.Socket) # type: ignore[attr-defined]
mock_get_socket.return_value = mock_socket
self.thread._get_remote_metadata("host1", 5555)
mock_get_socket.assert_called_once_with("host1", 5555)
mock_return_socket.assert_called_once_with(mock_socket, "host1", 5555)
mock_send.assert_called_once_with(mock_socket, self.thread.encoder.encode((GET_META_MSG, "")), "host1:5555")
mock_recv.assert_called_once_with(mock_socket, "host1:5555")
self.assertEqual(self.thread.kv_caches_base_addr["remote_engine"][5555], [[0x3000], [0x4000]])
self.assertEqual(self.thread.remote_block_stride_per_addr["remote_engine"][5555], [[1024]])
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.ensure_zmq_send")
@patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.ensure_zmq_recv",
side_effect=Exception("Network error"),
)
def test_get_remote_metadata_failure(self, mock_recv, mock_send):
with (
patch.object(self.thread, "_get_remote_socket") as mock_get_socket,
patch.object(self.thread, "_return_remote_socket") as mock_return_socket,
):
mock_socket = MagicMock(spec=zmq.Socket) # type: ignore[attr-defined]
mock_get_socket.return_value = mock_socket
with self.assertRaises(Exception) as context:
self.thread._get_remote_metadata("host1", 5555)
self.assertEqual(str(context.exception), "Network error")
mock_socket.close.assert_called_once()
mock_return_socket.assert_not_called()
class TestMainThreadLoop(unittest.TestCase):
def setUp(self):
self.engine = MagicMock()
self.ready_event = threading.Event()
self.vllm_config = MockVllmConfig()
self.kv_caches: dict[str, Any] = {}
self.thread = KVCacheRecvingThread(
tp_rank=0,
tp_size=4,
_prefill_pp_size=1,
engine=self.engine,
local_engine_id="local_engine",
local_handshake_port=5555,
side_channel_port=30000,
local_kv_caches_base_addr=[[0x1000], [0x2000]],
block_len_per_addr=[[1024], [2048]],
block_stride_per_addr=[[1024], [2048]],
ready_event=self.ready_event,
vllm_config=self.vllm_config,
kv_caches=self.kv_caches,
prefill_pp_layer_partition=None,
)
self.thread.request_queue = queue.Queue()
@patch.object(KVCacheRecvingThread, "_handle_request")
def test_run_loop_normal(self, mock_handle):
test_request = {
"request_id": "req1",
"local_block_ids": [1, 2],
"remote_block_ids": [3, 4],
"remote_engine_id": "remote_engine",
"remote_host": "localhost",
"remote_handshake_port": 6666,
"remote_transfer_port": 7777,
"offset": 0,
"tp_num_need_pulls": 2,
"all_task_done": False,
}
self.thread.request_queue.put(test_request)
self.thread.request_queue.put(None)
self.thread.start()
time.sleep(0.1)
self.thread.join(timeout=1.0)
self.assertTrue(self.thread.ready_event.is_set())
mock_handle.assert_called_once_with(test_request)
self.assertTrue(self.thread.request_queue.empty())
class MockVllmConfig:
def __init__(self):
self.model_config = MagicMock()
self.parallel_config = MagicMock()
self.cache_config = MagicMock()
self.kv_transfer_config = MagicMock()
self.scheduler_config = MagicMock(disable_hybrid_kv_cache_manager=True)
self.speculative_config = None
self.model_config.use_mla = False
self.model_config.is_deepseek_mla = False
self.model_config.hf_text_config = types.SimpleNamespace(
num_key_value_heads=8,
num_hidden_layers=32,
head_dim=16,
kv_lora_rank=16,
qk_rope_head_dim=8,
model_type="qwen2",
)
self.model_config.get_num_layers = MagicMock(return_value=32)
self.parallel_config.tensor_parallel_size = 2
self.parallel_config.data_parallel_rank = 0
self.parallel_config.data_parallel_size = 1
self.parallel_config.data_parallel_size_local = 1
self.parallel_config.pipeline_parallel_size = 1
self.parallel_config.data_parallel_rank_local = 0
self.parallel_config.prefill_context_parallel_size = 1
self.parallel_config.decode_context_parallel_size = 1
self.model_config.get_num_layers_by_block_type = MagicMock(return_value=32)
self.cache_config.block_size = 16
self.kv_transfer_config.kv_port = 5000
self.kv_transfer_config.kv_role = "kv_producer"
self.kv_transfer_config.engine_id = "test_engine"
self.kv_transfer_config.get_from_extra_config = MagicMock()
self.kv_transfer_config.get_from_extra_config.side_effect = lambda k, d: {
"prefill": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
"decode": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
}.get(k, d)
self.additional_config = {}
class MockRequest:
def __init__(self, request_id, prompt_token_ids=None, kv_transfer_params=None, status=None):
self.request_id = request_id
self.prompt_token_ids = prompt_token_ids or [1, 2, 3, 4]
self.kv_transfer_params = kv_transfer_params or {}
self.status = status or "running"
self.output_token_ids = [101, 102]
class MockKVCacheGroup:
def __init__(self, layer_names=None, kv_cache_spec=None):
self.layer_names = layer_names or ["model.layers.0.self_attn"]
self.kv_cache_spec = kv_cache_spec or MagicMock()
class MockKVCacheConfig:
def __init__(self, kv_cache_groups=None, num_blocks=10):
self.kv_cache_groups = kv_cache_groups or [MockKVCacheGroup()]
self.num_blocks = num_blocks
class TestKVCacheTaskTracker(unittest.TestCase):
def setUp(self):
self.tracker = KVCacheTaskTracker()
def test_update_done_task_count(self):
self.assertEqual(len(self.tracker.finished_requests), 0)
self.assertEqual(len(self.tracker.delayed_free_requests), 0)
self.assertEqual(len(self.tracker.reqs_to_process), 0)
current_time = time.time()
self.tracker.add_req_to_process("req_1")
self.tracker.add_delayed_request("req_1", current_time)
result = self.tracker.delayed_free_requests
self.assertEqual(len(result), 1)
self.assertEqual(result["req_1"], current_time)
self.tracker.update_done_task_count("req_1")
result_finished = self.tracker.finished_requests
result_delayed = self.tracker.delayed_free_requests
self.assertEqual(result_finished, {"req_1"})
self.assertEqual(len(result_delayed), 0)
self.assertEqual(len(self.tracker.reqs_to_process), 0)
self.tracker.update_done_task_count("req_2")
result_finished = self.tracker.finished_requests
result_delayed = self.tracker.delayed_free_requests
self.assertEqual(result_finished, {"req_1"})
self.assertEqual(len(result_delayed), 0)
self.assertEqual(len(self.tracker.reqs_to_process), 0)
def test_updtate_add_delayed_request(self) -> None:
self.tracker.update_done_task_count("req2")
self.tracker.add_delayed_request("req2", time.time())
result_delayed = self.tracker.delayed_free_requests
self.assertEqual(len(result_delayed), 0)
def test_retrieve_expired_requests(self):
current_time = time.time()
self.tracker.add_req_to_process("req_1")
self.tracker.add_req_to_process("req_2")
self.tracker.add_delayed_request("req_1", current_time - 100000)
self.tracker.add_delayed_request("req_2", current_time)
result = self.tracker._retrieve_expired_requests()
self.assertEqual(
result,
{
"req_1",
},
)
result_delay = self.tracker.delayed_free_requests
self.assertEqual(len(result_delay), 1)
self.assertIn("req_2", result_delay)
def test_duplicate_task_update(self):
self.tracker.add_req_to_process("req1")
self.tracker.update_done_task_count("req1")
self.tracker.update_done_task_count("req1")
self.tracker.update_done_task_count("req1")
finished = self.tracker.get_and_clear_finished_requests()
self.assertEqual(finished, {"req1"})
class TestMooncakeConnectorMetadata(unittest.TestCase):
def test_add_new_req(self):
meta = MooncakeConnectorMetadata()
self.assertEqual(len(meta.requests), 0)
self.assertEqual(len(meta.requests_to_send), 0)
meta.add_new_req(
request_id="req1",
local_block_ids=[1, 2, 3],
local_full_block_ids=[0, 1, 2, 3],
num_external_tokens=48,
kv_transfer_params={
"remote_block_ids": [4, 5, 6],
"remote_engine_id": "remote_engine",
"remote_request_id": "remote_req1",
"remote_host": "localhost",
"remote_port": 5000,
"remote_pcp_size": 1,
"remote_dcp_size": 1,
"remote_ptp_size": 2,
},
)
self.assertEqual(len(meta.requests), 1)
req_meta = meta.requests["req1"]
self.assertIsInstance(req_meta, ReqMeta)
self.assertEqual(req_meta.local_block_ids, [1, 2, 3])
self.assertEqual(req_meta.local_full_block_ids, [0, 1, 2, 3])
self.assertEqual(req_meta.remote_block_ids, [4, 5, 6])
self.assertEqual(req_meta.remote_engine_id, "remote_engine")
self.assertEqual(req_meta.remote_host, "localhost")
self.assertEqual(req_meta.remote_port, 5000)
self.assertEqual(req_meta.remote_ptp_size, 2)
class TestMooncakeConnectorSchedulerMatchedTokens(unittest.TestCase):
def setUp(self):
config = MockVllmConfig()
self.p1 = patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config", new=MagicMock()
)
self.p2 = patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
new=MagicMock(return_value=MagicMock()),
)
self.p1.start()
self.p2.start()
self.addCleanup(self.p1.stop)
self.addCleanup(self.p2.stop)
self.scheduler = MooncakeConnectorScheduler(config, "test_engine", MockKVCacheConfig())
def test_get_num_new_matched_tokens(self):
request = MockRequest("req1")
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
self.assertEqual(tokens, 0)
self.assertFalse(async_flag)
request.kv_transfer_params = {"do_remote_prefill": True}
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
self.assertEqual(tokens, 4)
self.assertTrue(async_flag)
self.assertEqual(request.kv_transfer_params["num_computed_tokens"], 0)
def test_build_connector_meta(self):
request = MockRequest("req1")
self.scheduler._reqs_need_recv["req1"] = (request, [4, 5, 6], [0, 4, 5, 6], 48)
request.kv_transfer_params = {
"remote_block_ids": [1, 2, 3],
"remote_engine_id": "remote",
"remote_request_id": "remote_req1",
"remote_host": "localhost",
"remote_port": 5000,
"remote_pcp_size": 1,
"remote_dcp_size": 1,
"num_computed_tokens": 16,
}
meta = self.scheduler.build_connector_meta(MagicMock())
self.assertIsInstance(meta, MooncakeConnectorMetadata)
self.assertEqual(len(meta.requests), 1)
self.assertEqual(meta.requests["req1"].local_block_ids, [4, 5, 6])
self.assertEqual(meta.requests["req1"].local_full_block_ids, [0, 4, 5, 6])
self.assertEqual(meta.requests["req1"].remote_block_ids, [1, 2, 3])
self.assertEqual(meta.requests["req1"].num_computed_tokens, 16)
self.assertEqual(len(self.scheduler._reqs_need_recv), 0)
class TestHelperFunctions(unittest.TestCase):
def test_group_concurrent_contiguous(self):
src: list[int] = [1, 2, 3, 5, 6]
dst: list[int] = [10, 11, 12, 14, 15]
src_groups, dst_groups = group_concurrent_contiguous(src, dst)
self.assertEqual(len(src_groups), 2)
self.assertEqual(src_groups[0], [1, 2, 3])
self.assertEqual(src_groups[1], [5, 6])
self.assertEqual(dst_groups[0], [10, 11, 12])
self.assertEqual(dst_groups[1], [14, 15])
def test_group_concurrent_contiguous_empty(self):
src: list[int] = []
dst: list[int] = []
src_groups, dst_groups = group_concurrent_contiguous(src, dst)
self.assertEqual(src_groups, [])
self.assertEqual(dst_groups, [])
def test_group_concurrent_contiguous_uses_stride_for_memory_contiguity(self):
src: list[int] = [1, 2]
dst: list[int] = [10, 11]
src_groups, dst_groups = group_concurrent_contiguous(
src,
dst,
src_block_stride=4096,
dst_block_stride=2048,
block_len=1024,
)
self.assertEqual(src_groups, [[1], [2]])
self.assertEqual(dst_groups, [[10], [11]])
def test_split_if_not_byte_contiguous_fast_path(self):
src_groups = [[1, 2]]
dst_groups = [[10, 11]]
src_result, dst_result = split_if_not_byte_contiguous(
src_groups,
dst_groups,
src_block_stride=1024,
dst_block_stride=1024,
block_len=1024,
)
self.assertIs(src_result, src_groups)
self.assertIs(dst_result, dst_groups)
def test_string_to_int64_hash(self):
hash1 = string_to_int64_hash("test_string")
hash2 = string_to_int64_hash("test_string")
self.assertEqual(hash1, hash2)
hash3 = string_to_int64_hash("different_string")
self.assertNotEqual(hash1, hash3)
class TestMooncakeConnectorForScheduler(unittest.TestCase):
def test_scheduler_role(self):
config = MockVllmConfig()
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
self.assertIsNotNone(connector.connector_scheduler)
self.assertIsNone(connector.connector_worker)
@patch.object(MooncakeConnectorScheduler, "get_num_new_matched_tokens")
def test_scheduler_methods(self, mock_method):
config = MockVllmConfig()
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
request = MockRequest("req1")
connector.get_num_new_matched_tokens(request, 0)
mock_method.assert_called_once_with(request, 0)
class MockKVCacheBlocks:
def get_unhashed_block_ids(self):
return [4, 5, 6]
def get_unhashed_block_ids_all_groups(self):
return ([4, 5, 6],)
def get_block_ids(self):
return ([1, 2, 4, 5, 6],)
class MockSchedulerOutput:
pass
class MockForwardContext:
pass
class TestMooncakeConnector(unittest.TestCase):
def setUp(self):
self.config = MockVllmConfig()
os.environ["ASCEND_RT_VISIBLE_DEVICES"] = "0,1"
def test_scheduler_initialization(self):
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(self.config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
self.assertIsNotNone(connector.connector_scheduler)
self.assertIsNone(connector.connector_worker)
@patch.object(MooncakeConnectorScheduler, "get_num_new_matched_tokens")
def test_get_num_new_matched_tokens(self, mock_method):
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(self.config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
request = MockRequest("req1")
connector.get_num_new_matched_tokens(request, 0)
mock_method.assert_called_once_with(request, 0)
@patch.object(MooncakeConnectorScheduler, "update_state_after_alloc")
def test_update_state_after_alloc(self, mock_method):
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(self.config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
request = MockRequest("req1")
blocks = MockKVCacheBlocks()
connector.update_state_after_alloc(request, blocks, 3)
mock_method.assert_called_once_with(request, blocks, 3)
@patch.object(MooncakeConnectorScheduler, "build_connector_meta")
def test_build_connector_meta(self, mock_method):
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(self.config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
scheduler_output = MockSchedulerOutput()
connector.build_connector_meta(scheduler_output)
mock_method.assert_called_once_with(scheduler_output)
@patch.object(MooncakeConnectorScheduler, "request_finished")
def test_request_finished(self, mock_method):
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
connector = MooncakeConnector(self.config, KVConnectorRole.SCHEDULER, MockKVCacheConfig())
request = MockRequest("req1")
connector.request_finished(request, [1, 2, 3])
mock_method.assert_called_once_with(request, ([1, 2, 3],))
class TestMooncakeConnectorScheduler(unittest.TestCase):
def setUp(self):
self.config = MockVllmConfig()
with (
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.init_ascend_config"),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
):
self.scheduler = MooncakeConnectorScheduler(self.config, "test_engine", MockKVCacheConfig())
def _make_remote_decode_request(self, prompt_len: int, request_id: str = "req1"):
return MockRequest(
request_id,
prompt_token_ids=list(range(prompt_len)),
kv_transfer_params={"do_remote_decode": True},
status=RequestStatus.FINISHED_LENGTH_CAPPED,
)
def test_get_num_new_matched_tokens_no_remote_prefill(self):
request = MockRequest("req1")
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
self.assertEqual(tokens, 0)
self.assertFalse(async_flag)
def test_get_num_new_matched_tokens_with_remote_prefill(self):
request = MockRequest("req1", kv_transfer_params={"do_remote_prefill": True})
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
self.assertEqual(tokens, 4)
self.assertTrue(async_flag)
def test_update_state_after_alloc_no_remote_prefill(self):
request = MockRequest("req1")
blocks = MagicMock()
self.scheduler.update_state_after_alloc(request, blocks, 0)
self.assertEqual(len(self.scheduler._reqs_need_recv), 0)
def test_update_state_after_alloc_with_remote_prefill(self):
request = MockRequest(
"req1",
kv_transfer_params={
"do_remote_prefill": True,
"remote_block_ids": [1, 2, 3],
"remote_engine_id": "remote",
"remote_request_id": "remote_req1",
"remote_host": "localhost",
"remote_port": 5000,
},
)
blocks = MockKVCacheBlocks()
self.scheduler.update_state_after_alloc(request, blocks, 3)
self.assertEqual(len(self.scheduler._reqs_need_recv), 1)
self.assertEqual(self.scheduler._reqs_need_recv["req1"][0], request)
self.assertEqual(self.scheduler._reqs_need_recv["req1"][1], ([4, 5, 6],))
self.assertEqual(self.scheduler._reqs_need_recv["req1"][2], ([1, 2, 4, 5, 6],))
def test_request_finished_no_remote_decode(self):
request = MockRequest("req1")
delay_free, params = self.scheduler.request_finished(request, [1, 2, 3])
self.assertFalse(delay_free)
self.assertIsNone(params)
def test_get_transfer_block_ids_trims_attention_mtp_blocks(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=0,
is_state_group=False,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([10, 11, 12, 13, 14],), prompt_len=33)
self.assertEqual(block_ids, ([10, 11, 12],))
def test_get_transfer_block_ids_keeps_state_group(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=0,
is_state_group=True,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([20, 21, 22, 23],), prompt_len=16)
self.assertEqual(block_ids, ([20, 21, 22, 23],))
def test_get_transfer_block_ids_uses_compressed_prompt_len(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=32,
blocks_per_window=0,
is_state_group=False,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([30, 31, 32, 33],), prompt_len=64)
self.assertEqual(block_ids, ([30, 31],))
def test_get_transfer_block_ids_uses_cp_grouped_block_len(self):
self.scheduler.pcp_size = 1
self.scheduler.dcp_size = 4
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=0,
is_state_group=False,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([10, 11, 12, 13, 14],), prompt_len=65)
self.assertEqual(block_ids, ([10, 11],))
def test_get_transfer_block_ids_trims_sliding_window_mtp_blocks(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([40, 41, 42, 43, 44],), prompt_len=48)
self.assertEqual(block_ids, ([40, 41, 42],))
def test_get_swa_transfer_block_ids_clips_sliding_window_group(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
)
]
block_ids = self.scheduler._get_swa_transfer_block_ids(([40, 41, 42, 43, 44],))
self.assertEqual(block_ids, ([42, 43, 44],))
def test_get_swa_transfer_block_ids_drops_zero_from_sliding_window_tail(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=2,
is_state_group=False,
)
]
block_ids = self.scheduler._get_swa_transfer_block_ids(([0, 10],))
self.assertEqual(block_ids, ([10],))
def test_transfer_block_ids_trims_mtp_before_swa_zero_filter(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace( # type: ignore[list-item]
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
)
]
block_ids = self.scheduler._get_transfer_block_ids(([0, 10, 11, 12, 13],), prompt_len=32)
block_ids = self.scheduler._get_swa_transfer_block_ids(block_ids)
self.assertEqual(block_ids, ([10],))
def test_request_finished_trims_mtp_blocks_in_params(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=0,
is_state_group=False,
)
]
request = self._make_remote_decode_request(prompt_len=33, request_id="req_mtp")
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
self.assertTrue(delay_free)
self.assertIsNotNone(params)
assert params is not None
self.assertEqual(params["remote_block_ids"], ([10, 11, 12],))
self.assertEqual(params["num_prompt_blocks"], 3)
self.assertIn("req_mtp", self.scheduler._reqs_need_send)
def test_request_finished_trims_cp_grouped_mtp_blocks_in_params(self):
self.scheduler.pcp_size = 1
self.scheduler.dcp_size = 4
self.scheduler.group_transfer_info = [
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=0,
is_state_group=False,
)
]
request = self._make_remote_decode_request(prompt_len=65, request_id="req_cp_mtp")
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
self.assertTrue(delay_free)
self.assertIsNotNone(params)
assert params is not None
self.assertEqual(params["remote_block_ids"], ([10, 11],))
# num_prompt_blocks stays in no-CP units for worker-side CP distribution.
self.assertEqual(params["num_prompt_blocks"], 5)
self.assertIn("req_cp_mtp", self.scheduler._reqs_need_send)
def test_request_finished_clips_sliding_window_blocks_in_params(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
)
]
request = self._make_remote_decode_request(prompt_len=80, request_id="req_swa")
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
self.assertTrue(delay_free)
self.assertIsNotNone(params)
assert params is not None
self.assertEqual(params["remote_block_ids"], ([12, 13, 14],))
self.assertEqual(params["num_prompt_blocks"], 5)
self.assertIn("req_swa", self.scheduler._reqs_need_send)
def test_request_finished_trims_mtp_before_swa_tail_clip(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
)
]
request = self._make_remote_decode_request(prompt_len=64, request_id="req_mtp_swa")
delay_free, params = self.scheduler.request_finished(request, ([0, 10, 11, 12, 13, 14],))
self.assertTrue(delay_free)
self.assertIsNotNone(params)
assert params is not None
self.assertEqual(params["remote_block_ids"], ([10, 11, 12],))
self.assertEqual(params["num_prompt_blocks"], 4)
self.assertIn("req_mtp_swa", self.scheduler._reqs_need_send)
def test_request_finished_handles_mtp_swa_and_state_groups_together(self):
self.scheduler.group_transfer_info = [
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=0,
is_state_group=False,
),
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=3,
is_state_group=False,
),
types.SimpleNamespace(
tokens_per_block=16,
blocks_per_window=0,
is_state_group=True,
),
]
request = self._make_remote_decode_request(prompt_len=64, request_id="req_mixed_groups")
delay_free, params = self.scheduler.request_finished(
request,
(
[100, 101, 102, 103, 104],
[0, 200, 201, 202, 203, 204],
[300, 301, 302, 303, 304],
),
)
self.assertTrue(delay_free)
self.assertIsNotNone(params)
assert params is not None
self.assertEqual(
params["remote_block_ids"],
(
[100, 101, 102, 103],
[200, 201, 202],
[300, 301, 302, 303, 304],
),
)
self.assertEqual(params["num_prompt_blocks"], 4)
self.assertIn("req_mixed_groups", self.scheduler._reqs_need_send)
class TestUtils(unittest.TestCase):
def test_string_to_int64_hash(self):
h1 = string_to_int64_hash("hello")
h2 = string_to_int64_hash("hello")
h3 = string_to_int64_hash("world")
self.assertEqual(h1, h2)
self.assertNotEqual(h1, h3)
self.assertIsInstance(h1, int)
def test_group_concurrent_contiguous(self):
src: list[int] = [1, 2, 3, 5, 6]
dst: list[int] = [10, 11, 12, 20, 21]
src_g, dst_g = group_concurrent_contiguous(src, dst)
self.assertEqual(src_g, [[1, 2, 3], [5, 6]])
self.assertEqual(dst_g, [[10, 11, 12], [20, 21]])
def test_group_empty(self):
src_g, dst_g = group_concurrent_contiguous([], [])
self.assertEqual(src_g, [])
self.assertEqual(dst_g, [])
def test_zmq_ctx_invalid_type(self):
with self.assertRaises(ValueError), zmq_ctx("INVALID", "tcp://127.0.0.1:5555"):
pass
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.make_zmq_socket")
def test_zmq_ctx_ok(self, mock_make_socket):
mock_socket = MagicMock()
mock_make_socket.return_value = mock_socket
with zmq_ctx(zmq.REQ, "tcp://localhost:1234") as s: # type: ignore
self.assertEqual(s, mock_socket)
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
def test_ensure_zmq_send_success(self, mock_logger):
mock_socket = MagicMock()
ensure_zmq_send(mock_socket, b"hello", "tcp://localhost:1234")
mock_socket.send.assert_called_once_with(b"hello")
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
def test_ensure_zmq_send_retry_and_fail(self, mock_logger):
mock_socket = MagicMock()
mock_socket.send.side_effect = zmq.ZMQError( # type: ignore
"send failed"
)
with self.assertRaises(RuntimeError):
ensure_zmq_send(mock_socket, b"hello", "tcp://localhost:1234", max_retries=2)
self.assertEqual(mock_socket.send.call_count, 2)
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
def test_ensure_zmq_recv_success(self, mock_logger):
mock_socket = MagicMock()
mock_socket.recv.return_value = b"response"
data = ensure_zmq_recv(mock_socket, "tcp://localhost:1234")
self.assertEqual(data, b"response")
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
def test_ensure_zmq_recv_timeout_and_fail(self, mock_logger):
mock_socket = MagicMock()
mock_socket.recv.side_effect = zmq.ZMQError("Receive timeout") # type: ignore
with self.assertRaises(RuntimeError):
ensure_zmq_recv(mock_socket, "tcp://localhost:1234", max_retries=2)
class MockMooncakeAgentMetadata:
def __init__(self, **kwargs):
pass
class MockMooncakeConnectorMetadata:
def __init__(self):
self.requests = {}
class MockKVCacheSendingThread(threading.Thread):
def __init__(self, *args, **kwargs):
super().__init__()
self.daemon = True
self._finished_requests = set()
def get_and_clear_finished_requests(self):
return self._finished_requests
def start(self):
pass
class MockKVCacheRecvingThread(threading.Thread):
def __init__(self, *args, **kwargs):
super().__init__()
self.daemon = True
self._finished_requests = set()
self.add_request = MagicMock()
def get_and_clear_finished_requests(self):
return self._finished_requests
def start(self):
pass
class MockTensor:
def __init__(self, *args, **kwargs):
self.size = MagicMock(return_value=(10, 16, 8, 16))
self.element_size = MagicMock(return_value=4)
self.shape = (10, 16, 8, 16)
self.data_ptr = MagicMock(return_value=0x1000)
mock_logger = MagicMock()
class MockTransferEngine:
def initialize(self, *args, **kwargs):
return 0
def register_memory(self, *args, **kwargs):
return 1
class MockEnvsAscend:
MOONCAKE_CONNECTOR_PROTOCOL = "mock_protocol"
def mock_get_tensor_model_parallel_rank():
return 0
def mock_get_tp_group():
return MagicMock()
def mock_get_ip():
return "127.0.0.1"
def mock_string_to_int64_hash(s):
return hash(s)
def make_cpu_kv_cache(kv_heads: int = 8, head_dim: int = 16):
return (
torch.empty((10, 16, kv_heads, head_dim), device="cpu"),
torch.empty((10, 16, kv_heads, head_dim), device="cpu"),
)
class TestMooncakeConnectorWorker(unittest.TestCase):
def setUp(self):
self.mock_transfer_engine = MagicMock()
self.mock_transfer_engine.get_rpc_port.return_value = 9090
self.mock_transfer_engine.initialize.return_value = 0
self.mock_transfer_engine.register_memory.return_value = 0
self.patches = [
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_tensor_model_parallel_rank",
mock_get_tensor_model_parallel_rank,
),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_tp_group", mock_get_tp_group),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_pp_group",
return_value=_mock_pp_group,
),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_pcp_group",
return_value=_mock_pcp_group,
),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_decode_context_model_parallel_world_size",
return_value=1,
),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_decode_context_model_parallel_rank",
return_value=0,
),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ip", mock_get_ip),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.string_to_int64_hash",
mock_string_to_int64_hash,
),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.global_te.get_transfer_engine",
return_value=self.mock_transfer_engine,
),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.global_te.register_buffer",
return_value=None,
),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.KVCacheSendingThread", MagicMock()),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.KVCacheRecvingThread", MagicMock()),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger", MagicMock()),
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.threading.Event", MagicMock()),
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
patch.object(
MooncakeConnectorWorker,
"_build_kv_group2layeridx",
return_value={
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"layer_names": ["model.layers.0.self_attn"],
},
[0],
)
},
),
]
for p in self.patches:
p.start() # type: ignore
self.vllm_config = MockVllmConfig()
self.engine_id = "test_engine"
self.kv_caches = {"model.layers.0.self_attn": make_cpu_kv_cache()}
def tearDown(self):
for p in self.patches:
p.stop() # type: ignore
def test_register_kv_caches_producer(self):
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.register_kv_caches(self.kv_caches)
self.assertEqual(len(worker.kv_caches), 1)
self.assertIsNotNone(worker.kv_send_thread)
self.assertIsNone(worker.kv_recv_thread)
def test_register_kv_caches_consumer(self):
self.vllm_config.kv_transfer_config.kv_role = "kv_consumer"
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.register_kv_caches(self.kv_caches)
self.assertIsNone(worker.kv_send_thread)
self.assertIsNotNone(worker.kv_recv_thread)
def test_register_kv_caches_mla_case(self):
self.vllm_config.model_config.is_deepseek_mla = True
mla_caches = {"model.layers.0.self_attn": make_cpu_kv_cache(kv_heads=1)}
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.register_kv_caches(mla_caches)
self.assertTrue(worker.use_mla)
self.assertEqual(len(worker.block_len_per_addr[0]), 2)
def test_device_id_selection_with_physical_devices(self):
# Test with physical devices set
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
# Default tp_rank is 0, so device_id should be 10
self.assertIsNotNone(worker.engine)
def test_get_remote_tp_rank(self):
def get_tp_rank(
prefill_tp_size: int,
prefill_pp_size: int,
decode_tp_size: int,
num_kv_heads: int,
tp_num_need_pulls: int,
is_deepseek_mla: bool,
remote_ptp_size: int | None = None,
):
with (
patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config",
return_value=MagicMock(),
),
patch.object(
self.vllm_config.kv_transfer_config,
"get_from_extra_config",
side_effect=lambda k, d=None: {
"prefill": {"tp_size": prefill_tp_size, "dp_size": 1, "pp_size": prefill_pp_size},
"decode": {"tp_size": decode_tp_size, "dp_size": 1, "pp_size": 1},
}.get(k, d),
),
):
self.vllm_config.model_config.hf_text_config.num_key_value_heads = num_kv_heads
self.vllm_config.model_config.is_deepseek_mla = is_deepseek_mla
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.tp_num_need_pulls = tp_num_need_pulls
worker.use_sparse = False
return worker._get_remote_ranks_for_req("test", remote_ptp_size)
self.assertIn(
get_tp_rank(16, 1, 1, 4, 4, False)[0], [[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14], [3, 7, 11, 15]]
)
self.assertIn(get_tp_rank(8, 1, 1, 4, 4, False)[0], [[0, 2, 4, 6], [1, 3, 5, 7]])
self.assertIn(get_tp_rank(4, 1, 1, 4, 4, False)[0], [[0, 1, 2, 3]])
self.assertIn(
get_tp_rank(16, 1, 4, 4, 1, False),
[[[0], [4], [8], [12]], [[1], [5], [9], [13]], [[2], [6], [10], [14]], [[3], [7], [11], [15]]],
)
self.assertIn(get_tp_rank(8, 1, 4, 4, 1, False), [[[0], [2], [4], [6]], [[1], [3], [5], [7]]])
self.assertIn(get_tp_rank(4, 2, 2, 4, 2, False), [[[0, 1, 4, 5], [2, 3, 6, 7]]])
self.assertIn(get_tp_rank(4, 1, 4, 4, 1, False), [[[0], [1], [2], [3]]])
self.assertIn(get_tp_rank(8, 2, 1, 4, 4, False)[0], [[0, 2, 4, 6, 8, 10, 12, 14], [1, 3, 5, 7, 9, 11, 13, 15]])
self.assertIn(get_tp_rank(4, 2, 2, 4, 2, False), [[[0, 1, 4, 5], [2, 3, 6, 7]]])
self.assertIn(get_tp_rank(2, 2, 1, 4, 2, False), [[[0, 1, 2, 3]]])
self.assertIn(get_tp_rank(4, 4, 2, 8, 2, False), [[[0, 1, 4, 5, 8, 9, 12, 13], [2, 3, 6, 7, 10, 11, 14, 15]]])
self.assertIn(get_tp_rank(4, 2, 1, 4, 4, False)[0], [[0, 1, 2, 3, 4, 5, 6, 7]])
self.assertIn(get_tp_rank(4, 4, 1, 4, 4, False)[0], [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]])
self.assertIn(
get_tp_rank(8, 2, 4, 4, 1, False),
[[[0, 8], [2, 10], [4, 12], [6, 14]], [[1, 9], [3, 11], [5, 13], [7, 15]]],
)
self.assertIn(get_tp_rank(4, 2, 4, 4, 4, False), [[[0, 4], [1, 5], [2, 6], [3, 7]]])
self.assertIn(
get_tp_rank(4, 4, 4, 4, 1, False), [[[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14], [3, 7, 11, 15]]]
)
self.assertIn(
get_tp_rank(16, 1, 1, 1, 1, True)[0],
[[0], [1], [2], [3], [4], [5], [6], [7], [8], [9], [10], [11], [12], [13], [14], [15]],
)
self.assertIn(get_tp_rank(4, 1, 4, 1, 1, True), [[[0], [1], [2], [3]]])
self.assertIn(
get_tp_rank(8, 2, 1, 1, 1, True)[0], [[0, 8], [2, 10], [4, 12], [6, 14], [1, 9], [3, 11], [5, 13], [7, 15]]
)
self.assertIn(
get_tp_rank(4, 4, 1, 1, 1, True)[0], [[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14], [3, 7, 11, 15]]
)
self.assertIn(
get_tp_rank(8, 2, 4, 1, 1, True)[0], [[0, 8], [2, 10], [4, 12], [6, 14], [1, 9], [3, 11], [5, 13], [7, 15]]
)
self.assertIn(
get_tp_rank(4, 4, 4, 1, 1, True), [[[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14], [3, 7, 11, 15]]]
)
# check remote ptp size
self.assertListEqual(get_tp_rank(16, 1, 2, 4, 2, False, 8), get_tp_rank(8, 1, 2, 4, 2, False))
self.assertListEqual(get_tp_rank(8, 1, 2, 4, 2, False, 4), get_tp_rank(4, 1, 2, 4, 2, False))
self.assertListEqual(get_tp_rank(4, 1, 2, 4, 1, False, 2), get_tp_rank(2, 1, 2, 4, 1, False))
def test_get_kv_split_metadata(self):
def get_kv_split_metadata(
use_mla,
pcp_size,
dcp_size,
tp_size,
tp_rank,
pcp_rank,
_prefill_tp_size,
remote_pcp_size,
remote_dcp_size,
remote_port,
remote_block_ids,
local_block_ids,
remote_engine_id,
remote_ptp_size=None,
remote_block_size=0,
dcp_rank=0,
):
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = use_mla
worker.pcp_size = pcp_size
worker.dcp_size = dcp_size
worker.tp_size = tp_size
worker.tp_rank = tp_rank
worker.pcp_rank = pcp_rank
worker.dcp_rank = 0
worker._prefill_tp_size = _prefill_tp_size
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.block_size = 16
worker.num_key_value_heads = 1
worker.use_sparse = False
# scale 1 => kernel size == block size (kernel ids == logical ids for equal sizes)
worker.block_size_scale = [[1]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"layer_names": ["model.layers.0.self_attn"],
},
[0],
)
}
meta = types.SimpleNamespace()
meta.remote_pcp_size = remote_pcp_size
meta.remote_dcp_size = remote_dcp_size
meta.remote_ptp_size = remote_ptp_size
meta.remote_port = remote_port
meta.remote_block_ids = (remote_block_ids,)
meta.local_block_ids = (local_block_ids,)
meta.num_external_tokens = pcp_size * dcp_size * len(local_block_ids) * worker.block_size
meta.num_prompt_blocks = pcp_size * dcp_size * len(local_block_ids)
meta.num_computed_tokens = 0
meta.remote_engine_id = remote_engine_id
meta.remote_host = "localhost"
meta.remote_block_size = worker.block_size
meta.remote_multi_nodes_meta_mapping = {}
(
remote_handshake_port_list,
local_block_ids_list,
remote_block_ids_list,
) = worker._get_kv_split_metadata("0", cast(ReqMeta, meta))
return (
remote_handshake_port_list,
[block_ids[0] for block_ids in local_block_ids_list],
[block_ids[0] for block_ids in remote_block_ids_list],
)
self.assertEqual(
get_kv_split_metadata(True, 1, 1, 8, 1, 0, 8, 1, 8, 30000, [1], [1], 0, remote_block_size=32),
(
[[30001], [30002], [30003], [30004], [30005], [30006], [30007], [30000]],
[[], [], [], [], [], [], [], [1]],
[[], [], [], [], [], [], [], [1]],
),
)
self.assertEqual(
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 8, 2, 8, 30000, [1], [1], 0),
(
[
[30001],
[30002],
[30003],
[30004],
[30005],
[30006],
[30007],
[30008],
[30009],
[30010],
[30011],
[30012],
[30013],
[30014],
[30015],
[30000],
],
[[], [], [], [], [], [], [], [], [], [], [], [], [], [], [], [1]],
[[], [], [], [], [], [], [], [], [], [], [], [], [], [], [], [1]],
),
)
self.assertEqual(
get_kv_split_metadata(True, 1, 1, 8, 1, 0, 8, 2, 2, 30000, [1], [1], 0),
([[30001], [30008], [30009], [30000]], [[], [], [], [1]], [[], [], [], [1]]),
)
self.assertEqual(
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 8, 2, 2, 30000, [1], [1], 0),
([[30001], [30008], [30009], [30000]], [[], [], [], [1]], [[], [], [], [1]]),
)
self.assertEqual(
get_kv_split_metadata(True, 1, 2, 8, 1, 0, 8, 2, 2, 30000, [1], [1], 0),
([[30000], [30008]], [[1], []], [[1], []]),
)
self.assertEqual(
get_kv_split_metadata(False, 1, 2, 8, 1, 0, 8, 2, 2, 30000, [1], [1], 0),
([[30000], [30008]], [[1], []], [[1], []]),
)
# D rank0 holds 5 external blocks [1,2,3,4,5]; P stores blocks interleaved
# across 4 cp ranks (cp0: global 0,4,8 -> D local idx 0,2,4 = blocks 1,3,5;
# cp2: global 2,6 -> D local idx 1,3 = blocks 2,4). Expansion now happens in
# _get_kv_split_metadata (scale 1 => kernel == block), so each shard's local
# list is the chunk-selected kernels: shard0 -> [1,3,5], shard1 -> [2,4].
self.assertEqual(
get_kv_split_metadata(True, 1, 2, 8, 0, 0, 8, 2, 2, 30000, [1, 2, 3], [1, 2, 3, 4, 5], 0)[:3],
([[30000], [30008]], [[1, 3, 5], [2, 4]], [[1, 2, 3], [1, 2]]),
)
# P cp size 4 -> D cp size 2: P block ids are per-CP local ids. D rank0
# must write cp0's [1,2,3] into local [1,3,5] and cp2's [1,2,3] into
# local [2,4,6], rather than slicing local blocks contiguously.
self.assertEqual(
get_kv_split_metadata(True, 1, 2, 8, 0, 0, 8, 1, 4, 30000, [1, 2, 3], [1, 2, 3, 4, 5, 6], 0)[:3],
([[30000], [30002]], [[1, 3, 5], [2, 4, 6]], [[1, 2, 3], [1, 2, 3]]),
)
# check remote ptp size
self.assertEqual(
get_kv_split_metadata(True, 1, 1, 8, 1, 0, 8, 1, 8, 30000, [1], [1], 0, 16),
get_kv_split_metadata(True, 1, 1, 8, 1, 0, 16, 1, 8, 30000, [1], [1], 0),
)
self.assertEqual(
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 8, 1, 8, 30000, [1], [1], 0, 16),
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 16, 1, 8, 30000, [1], [1], 0),
)
self.assertEqual(
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 8, 2, 8, 30000, [1], [1], 0, 16),
get_kv_split_metadata(False, 1, 1, 8, 1, 0, 16, 2, 8, 30000, [1], [1], 0),
)
def test_get_kv_split_metadata_unequal_block_size_with_decode_cp(self):
"""Bd=2*Bp with D-side CP: P cp ranks 0,1 -> D rank0; cp ranks 2,3 -> D rank1."""
for dcp_rank in (0, 1):
with self.subTest(dcp_rank=dcp_rank):
with patch(
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_decode_context_model_parallel_rank",
return_value=dcp_rank,
):
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = False
worker.use_sparse = False
worker.pcp_size = 1
worker.dcp_size = 2
worker.pcp_rank = 0
worker.dcp_rank = dcp_rank
worker.tp_size = 2
worker.tp_rank = dcp_rank
worker.block_size = 32
worker.num_key_value_heads = 1
worker._prefill_tp_size = 4
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.side_channel_port = 5000
worker.handshake_port = worker.side_channel_port + worker.tp_rank
# Bd=32 stored as 2 kernels of 16 (scale 2); Bp=16 == 1 kernel.
worker.block_size_scale = [[2]]
worker.kv_group2layeridx = {
0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0]),
}
meta = types.SimpleNamespace(
remote_pcp_size=2,
remote_dcp_size=2,
remote_ptp_size=4,
remote_port=30000,
remote_block_ids=([10, 11, 12, 13],),
local_block_ids=([100],),
num_external_tokens=64,
num_prompt_blocks=4,
num_computed_tokens=0,
remote_block_size=16,
remote_engine_id=f"remote_bs_{dcp_rank}",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_bs", cast(ReqMeta, meta))
self.assertEqual(len(ports), 2)
# local D block 100 (size 32) splits into kernels 200 (first half, start 0)
# and 201 (second half, start 16); each remote block 10 is a single kernel.
self.assertEqual(local_ids, [([200],), ([201],)])
self.assertEqual(remote_ids, [([10],), ([10],)])
if dcp_rank == 0:
self.assertEqual(ports, [[30000], [30001]])
else:
self.assertEqual(ports, [[30004], [30005]])
def test_get_kv_split_metadata_cp_with_prefix_cache_skips_prefix(self):
"""CP + prefix cache hit (P0>0): remote ids must start past the prefix
blocks (remote_first), aligned with local_chunk_token_starts."""
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = True
worker.use_sparse = False
worker.pcp_size = 1
worker.dcp_size = 1
worker.pcp_rank = 0
worker.dcp_rank = 0
worker.tp_size = 8
worker.tp_rank = 0
worker.block_size = 16
worker.num_key_value_heads = 1
worker._prefill_tp_size = 8
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.side_channel_port = 5000
worker.handshake_port = worker.side_channel_port + worker.tp_rank
worker.block_size_scale = [[1]]
worker.kv_group2layeridx = {
0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0]),
}
# 6 prompt blocks, 4 external (P0 = 2 prefix-cached blocks), remote PCP=2.
meta = types.SimpleNamespace(
remote_pcp_size=2,
remote_dcp_size=1,
remote_ptp_size=8,
remote_port=30000,
remote_block_ids=([50, 51, 52],),
local_block_ids=([100, 101, 102, 103],),
num_external_tokens=4 * worker.block_size,
num_prompt_blocks=6,
num_computed_tokens=0,
remote_block_size=16,
remote_engine_id="remote_prefix_cp",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_prefix_cp", cast(ReqMeta, meta))
self.assertEqual(len(ports), 2)
# remote starts at index remote_first=1 (skips the prefix block 50), NOT [:2].
self.assertEqual(remote_ids, [([51, 52],), ([51, 52],)])
# Expansion (scale 1) now selects per-shard kernels via the interleaved token
# starts [[0,32],[16,48]]: shard0 -> blocks 100,102; shard1 -> blocks 101,103.
self.assertEqual(local_ids, [([100, 102],), ([101, 103],)])
def _build_non_cp_worker(self):
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = False
worker.use_sparse = False
worker.pcp_size = 1
worker.dcp_size = 1
worker.pcp_rank = 0
worker.dcp_rank = 0
worker.tp_size = 1
worker.tp_rank = 0
worker.block_size = 16
worker.num_key_value_heads = 1
worker._prefill_tp_size = 1
worker._is_hma_required = False
# No CP, so the remote rank choice is irrelevant to the expansion under test.
worker._get_remote_rank = lambda *a, **k: [0]
return worker
def test_get_kv_split_metadata_non_cp_prefix_skip_and_trim(self):
"""No-CP: remote kernels are expanded, prefix-skipped by num_computed_tokens,
and trimmed to the local count - all inside _get_kv_split_metadata now."""
worker = self._build_non_cp_worker()
worker.block_size_scale = [[1]]
worker.kv_group2layeridx = {0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0])}
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=1,
remote_ptp_size=1,
remote_port=30000,
remote_block_ids=([3, 4, 5],),
local_block_ids=([1, 2],),
num_external_tokens=2 * worker.block_size,
num_prompt_blocks=3,
num_computed_tokens=worker.block_size, # 1 prefix block -> skip first remote kernel
remote_block_size=worker.block_size,
remote_engine_id="e_non_cp",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("r", cast(ReqMeta, meta))
self.assertEqual(ports, [[30000]])
# scale 1: remote kernels [3,4,5] -> skip 1 -> [4,5]; local [1,2]; min -> 2.
self.assertEqual(local_ids, [([1, 2],)])
self.assertEqual(remote_ids, [([4, 5],)])
def test_get_kv_split_metadata_non_cp_uses_compress_ratio(self):
"""No-CP: the per-group compress_ratio scales the prefix-skip offset."""
worker = self._build_non_cp_worker()
# block 16 stored as 2 kernels of 8 (scale 2).
worker.block_size_scale = [[2]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
"kv_cache_spec": {"layer_0": {"compress_ratio": 4}},
},
[0],
)
}
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=1,
remote_ptp_size=1,
remote_port=30000,
remote_block_ids=([3, 4],),
local_block_ids=([1, 2],),
num_external_tokens=2 * worker.block_size,
num_prompt_blocks=3,
# kernel token size = kernel_size(8) * compress_ratio(4) = 32 -> skip 1 kernel.
num_computed_tokens=32,
remote_block_size=worker.block_size,
remote_engine_id="e_compress",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("r", cast(ReqMeta, meta))
self.assertEqual(ports, [[30000]])
# local [1,2] -> kernels [2,3,4,5]; remote [3,4] -> kernels [6,7,8,9] -> skip 1 -> [7,8,9];
# min -> 3 kernels.
self.assertEqual(local_ids, [([2, 3, 4],)])
self.assertEqual(remote_ids, [([7, 8, 9],)])
def _build_worker_for_pd_case(self, case, tp_rank, pcp_rank=0, dcp_rank=0):
with patch.object(
self.vllm_config.kv_transfer_config,
"get_from_extra_config",
side_effect=lambda k, d=None, case=case: {
"prefill": {
"tp_size": case["prefill_tp_size"],
"dp_size": 1,
"pp_size": case["prefill_pp_size"],
},
"decode": {"tp_size": case["decode_tp_size"], "dp_size": 1, "pp_size": 1},
}.get(k, d),
):
self.vllm_config.model_config.is_deepseek_mla = case["use_mla"]
self.vllm_config.model_config.hf_text_config.num_key_value_heads = case["num_key_value_heads"]
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = case["use_mla"]
worker.use_sparse = False
worker.num_key_value_heads = case["num_key_value_heads"]
worker.tp_size = case["decode_tp_size"]
worker.tp_rank = tp_rank
worker.pcp_size = case["pcp_size"]
worker.dcp_size = case["dcp_size"]
worker.pcp_rank = pcp_rank
worker.dcp_rank = dcp_rank
worker.pp_rank = 0
worker._prefill_tp_size = case["prefill_tp_size"]
worker._prefill_pp_size = case["prefill_pp_size"]
worker._decode_tp_size = case["decode_tp_size"]
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.side_channel_port = 5000
worker.handshake_port = worker.side_channel_port + (worker.pp_rank + pcp_rank) * worker.tp_size + tp_rank
worker.block_size_scale = [[1], [1]]
worker.kv_group2layeridx = {
0: ({"kv_cache_spec_type": "FullAttentionSpec"}, [0]),
1: ({"kv_cache_spec_type": "FullAttentionSpec"}, [1]),
}
return worker
def _assert_group_pull_finish_flags(self, ports, group_pulls, expected_group_ids):
self.assertEqual(len(group_pulls), len(ports))
finish_count_by_group = {group_id: 0 for group_id in expected_group_ids}
for pcp_dcp_rank, (remote_ports, port_group_pulls) in enumerate(zip(ports, group_pulls)):
self.assertEqual(len(port_group_pulls), len(remote_ports))
for remote_port_idx, pulls in enumerate(port_group_pulls):
self.assertEqual({pull.group_id for pull in pulls}, expected_group_ids)
for pull in pulls:
self.assertEqual(
pull.is_group_transfer_end,
pull.remote_tp_offset == pull.num_group_pulls - 1,
f"port={remote_ports[remote_port_idx]}, group={pull.group_id}, "
f"offset={pull.remote_tp_offset}, num_pulls={pull.num_group_pulls}",
)
if pull.is_group_transfer_end:
finish_count_by_group[pull.group_id] += 1
if len(remote_ports) == 1:
expected_offset = pcp_dcp_rank % pulls[0].num_group_pulls
else:
expected_offset = remote_port_idx % pulls[0].num_group_pulls
self.assertTrue(all(pull.remote_tp_offset == expected_offset for pull in pulls))
self.assertTrue(
all(count > 0 for count in finish_count_by_group.values()),
f"Each group should have at least one pull-finish marker: {finish_count_by_group}",
)
def _assert_hybrid_group_pull_finish_flags(self, ports, group_pulls, expected_group_ids, expected_finishes):
self.assertEqual(len(group_pulls), len(ports))
finish_count_by_group = {group_id: 0 for group_id in expected_group_ids}
seen_group_ids = set()
for remote_ports, port_group_pulls in zip(ports, group_pulls):
self.assertEqual(len(port_group_pulls), len(remote_ports))
for remote_port, pulls in zip(remote_ports, port_group_pulls):
self.assertTrue(pulls, f"remote port {remote_port} should pull at least one group")
for pull in pulls:
self.assertIn(pull.group_id, expected_group_ids)
seen_group_ids.add(pull.group_id)
self.assertEqual(
pull.is_group_transfer_end,
pull.remote_tp_offset == pull.num_group_pulls - 1,
f"port={remote_port}, group={pull.group_id}, offset={pull.remote_tp_offset}, "
f"num_pulls={pull.num_group_pulls}",
)
if pull.is_group_transfer_end:
finish_count_by_group[pull.group_id] += 1
self.assertEqual(seen_group_ids, expected_group_ids)
self.assertEqual(finish_count_by_group, expected_finishes)
def test_pd_disaggregated_split_cross_covers_prefix_tp_cp_pp(self):
cases: list[dict[str, Any]] = [
{
"name": "gqa_tp_unequal_remote_cp_pp_unequal",
"use_mla": False,
"num_key_value_heads": 2,
"prefill_tp_size": 8,
"decode_tp_size": 4,
"prefill_pp_size": 2,
"remote_pcp_size": 2,
"remote_dcp_size": 2,
"pcp_size": 1,
"dcp_size": 2,
"remote_block_ids": ([10, 11], [10, 11]),
"local_block_ids": ([20, 21], [20, 21]),
"num_prompt_blocks": 6,
"num_external_blocks": 4,
},
{
"name": "mla_tp_unequal_decode_cp_unequal",
"use_mla": True,
"num_key_value_heads": 1,
"prefill_tp_size": 8,
"decode_tp_size": 4,
"prefill_pp_size": 1,
"remote_pcp_size": 1,
"remote_dcp_size": 4,
"pcp_size": 1,
"dcp_size": 2,
"remote_block_ids": ([30, 31, 32], [30, 31, 32]),
"local_block_ids": ([40, 41], [40, 41]),
"num_prompt_blocks": 5,
"num_external_blocks": 4,
},
]
for case in cases:
for tp_rank in range(case["decode_tp_size"]):
for pcp_rank in range(case["pcp_size"]):
for dcp_rank in range(case["dcp_size"]):
with self.subTest(
case=case["name"],
tp_rank=tp_rank,
pcp_rank=pcp_rank,
dcp_rank=dcp_rank,
):
worker = self._build_worker_for_pd_case(case, tp_rank, pcp_rank, dcp_rank)
meta = types.SimpleNamespace(
remote_pcp_size=case["remote_pcp_size"],
remote_dcp_size=case["remote_dcp_size"],
remote_ptp_size=case["prefill_tp_size"],
remote_port=30000,
remote_block_ids=case["remote_block_ids"],
local_block_ids=case["local_block_ids"],
num_external_tokens=case["num_external_blocks"] * worker.block_size,
num_prompt_blocks=case["num_prompt_blocks"],
remote_block_size=worker.block_size,
remote_engine_id=f"remote_{case['name']}_{tp_rank}_{pcp_rank}_{dcp_rank}",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_pd", meta)
group_pulls = worker._get_group_pulls_metadata(
"req_pd",
ports,
case["prefill_tp_size"],
30000,
case["remote_pcp_size"],
case["remote_dcp_size"],
)
self.assertEqual(len(ports), len(local_ids))
self.assertEqual(len(local_ids), len(remote_ids))
# Expansion now happens in _get_kv_split_metadata (scale 1 =>
# kernel == block), so each shard carries only the kernels it
# writes. The rank's external blocks are partitioned across
# shards, so the per-shard local lengths sum to the per-rank
# external block count.
per_rank_external_blocks = case["num_external_blocks"] // (
case["pcp_size"] * case["dcp_size"]
)
self.assertEqual(sum(len(ids[0]) for ids in local_ids), per_rank_external_blocks)
self._assert_group_pull_finish_flags(ports, group_pulls, {0, 1})
def test_pd_disaggregated_hybrid_prefix_tp_and_pp_unequal(self):
for tp_rank in range(2):
with self.subTest(tp_rank=tp_rank):
with patch.object(
self.vllm_config.kv_transfer_config,
"get_from_extra_config",
side_effect=lambda k, d=None: {
"prefill": {"tp_size": 4, "dp_size": 1, "pp_size": 2},
"decode": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
}.get(k, d),
):
self.vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
self.vllm_config.model_config.is_deepseek_mla = False
self.vllm_config.model_config.hf_text_config.num_key_value_heads = 8
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker._is_hma_required = True
worker.use_mla = False
worker.use_sparse = False
worker.num_key_value_heads = 8
worker.tp_size = 2
worker.tp_rank = tp_rank
worker.pcp_size = 1
worker.dcp_size = 1
worker._decode_tp_size = 2
worker._prefill_tp_size = 4
worker._prefill_pp_size = 2
worker.block_size_scale = [[1], [1], [1]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_spec": {"num_kv_heads": 8},
},
[0, 1],
),
1: ({"kv_cache_spec_type": "MambaSpec"}, [2]),
}
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=1,
remote_ptp_size=4,
remote_port=31000,
remote_block_ids=([50, 51, 52], [60, 61, 62]),
local_block_ids=([70, 71], [80, 81, 82]),
num_external_tokens=3 * worker.block_size,
num_prompt_blocks=4,
num_computed_tokens=0,
remote_block_size=worker.block_size,
remote_engine_id=f"remote_hybrid_{tp_rank}",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_hybrid", cast(ReqMeta, meta))
group_pulls = worker._get_group_pulls_metadata(
"req_hybrid", ports, 4, 31000, meta.remote_pcp_size, meta.remote_dcp_size
)
# Attention (group 0) is now expanded + min-trimmed in metadata: D holds 2
# external blocks [70,71], so the 3 remote blocks are trimmed to [50,51].
# Mamba (group 1) keeps the full logical state.
self.assertEqual(local_ids, [([70, 71], [80, 81, 82])])
self.assertEqual(remote_ids, [([50, 51], [60, 61, 62])])
self.assertGreater(len(ports[0]), 1)
self._assert_hybrid_group_pull_finish_flags(
ports,
group_pulls,
expected_group_ids={0, 1},
expected_finishes={0: worker._prefill_pp_size, 1: worker._prefill_pp_size},
)
def test_pd_disaggregated_hybrid_remote_pcp_splits_attention_and_final_mamba_state(self):
for tp_rank in range(2):
with self.subTest(tp_rank=tp_rank):
with patch.object(
self.vllm_config.kv_transfer_config,
"get_from_extra_config",
side_effect=lambda k, d=None: {
"prefill": {"tp_size": 4, "dp_size": 1, "pp_size": 1},
"decode": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
}.get(k, d),
):
self.vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
self.vllm_config.model_config.is_deepseek_mla = False
self.vllm_config.model_config.hf_text_config.num_key_value_heads = 8
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker._is_hma_required = True
worker.use_mla = False
worker.use_sparse = False
worker.num_key_value_heads = 8
worker.tp_size = 2
worker.tp_rank = tp_rank
worker.pcp_size = 1
worker.dcp_size = 1
worker.pcp_rank = 0
worker.dcp_rank = 0
worker._decode_tp_size = 2
worker._prefill_tp_size = 4
worker._prefill_pp_size = 1
worker.side_channel_port = 5000
worker.handshake_port = worker.side_channel_port + tp_rank
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.block_size_scale = [[1], [1], [1]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_spec": {"num_kv_heads": 8},
},
[0, 1],
),
1: ({"kv_cache_spec_type": "MambaSpec"}, [2]),
}
meta = types.SimpleNamespace(
remote_pcp_size=2,
remote_dcp_size=1,
remote_ptp_size=4,
remote_port=31000,
remote_block_ids=([50, 51, 52, 53], [60, 61, 62, 63]),
local_block_ids=([70, 71, 72, 73], [80, 81, 82, 83]),
num_external_tokens=4 * worker.block_size,
num_prompt_blocks=4,
remote_engine_id=f"remote_hybrid_pcp_{tp_rank}",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
remote_block_size=16,
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_hybrid_pcp", cast(ReqMeta, meta))
group_pulls = worker._get_group_pulls_metadata(
"req_hybrid_pcp", ports, 4, 31000, meta.remote_pcp_size, meta.remote_dcp_size
)
self.assertEqual(len(ports), 2)
# Attention (group 0) is expanded in metadata (scale 1): the 4 external
# blocks are interleaved across the 2 PCP shards, 2 kernels each.
self.assertEqual([len(ids[0]) for ids in local_ids], [2, 2])
self.assertEqual([ids[1] for ids in local_ids], [[], [80, 81, 82, 83]])
self.assertEqual([ids[1] for ids in remote_ids], [[], [60, 61, 62, 63]])
self.assertTrue(worker.remote_port_send_num[meta.remote_engine_id])
self._assert_hybrid_group_pull_finish_flags(
ports,
group_pulls,
expected_group_ids={0, 1},
expected_finishes={0: 2, 1: 1},
)
def test_hybrid_no_cp_uses_kv_cache_group_ids_for_split_transfer_groups(self):
with patch.object(
self.vllm_config.kv_transfer_config,
"get_from_extra_config",
side_effect=lambda k, d=None: {
"prefill": {"tp_size": 4, "dp_size": 1, "pp_size": 1},
"decode": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
}.get(k, d),
):
self.vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
self.vllm_config.model_config.is_deepseek_mla = False
self.vllm_config.model_config.hf_text_config.num_key_value_heads = 8
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker._is_hma_required = True
worker.use_mla = False
worker.use_sparse = False
worker.num_key_value_heads = 8
worker.tp_size = 2
worker.tp_rank = 0
worker.pcp_size = 1
worker.dcp_size = 1
worker.pcp_rank = 0
worker.dcp_rank = 0
worker._decode_tp_size = 2
worker._prefill_tp_size = 4
worker._prefill_pp_size = 1
worker.side_channel_port = 5000
worker.handshake_port = worker.side_channel_port + worker.tp_rank
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.block_size_scale = [[1], [1], [1]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_group_id": 0,
"kv_cache_spec": {"num_kv_heads": 1},
},
[0],
),
1: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"kv_cache_group_id": 0,
"kv_cache_spec": {"num_kv_heads": 8},
},
[1],
),
2: (
{
"kv_cache_spec_type": "MambaSpec",
"kv_cache_group_id": 1,
},
[2],
),
}
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=1,
remote_ptp_size=4,
remote_port=31000,
remote_block_ids=([50, 51, 52, 53], [60, 61, 62, 63]),
local_block_ids=([70, 71, 72, 73], [80, 81, 82, 83]),
num_external_tokens=4 * worker.block_size,
num_prompt_blocks=4,
num_computed_tokens=0,
remote_engine_id="remote_hybrid_split_transfer_groups",
remote_host="localhost",
remote_multi_nodes_meta_mapping={},
remote_block_size=16,
)
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_hybrid_split", cast(ReqMeta, meta))
self.assertEqual(len(ports), 1)
self.assertEqual(local_ids, [([70, 71, 72, 73], [80, 81, 82, 83])])
self.assertEqual(remote_ids, [([50, 51, 52, 53], [60, 61, 62, 63])])
def test_get_tp_num_need_pulls(self):
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.num_key_value_heads = 8
worker.vllm_config.model_config.is_deepseek_mla = True
tp_num_need_pulls = worker._get_tp_num_need_pulls(prefill_tp_size=4)
self.assertEqual(tp_num_need_pulls, 1)
worker.vllm_config.model_config.is_deepseek_mla = False
tp_num_need_pulls = worker._get_tp_num_need_pulls(prefill_tp_size=4)
self.assertEqual(tp_num_need_pulls, 2)
tp_num_need_pulls = worker._get_tp_num_need_pulls(prefill_tp_size=None)
self.assertEqual(tp_num_need_pulls, 1)
def test_start_load_kv_puts_replicated_indexer_on_existing_transfer_port(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.kv_send_thread = None
worker.kv_recv_thread = MagicMock()
worker._prefill_tp_size = 4
worker.remote_port_send_num = {"remote_engine": {31001: {"num": 1, "host": "localhost"}}}
worker._get_sfa_replicate_k_block_ids = MagicMock(return_value=(([40],), ([20],)))
worker._get_kv_split_metadata = MagicMock(
return_value=(
[[31001], [31003]],
[([10],), ([11],)],
[([30],), ([31],)],
)
)
worker._get_group_pulls_metadata = MagicMock(
return_value=[
[[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]],
[[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]],
]
)
worker._get_remote_host_info_by_port = MagicMock(return_value=("localhost", "remote_engine"))
meta = types.SimpleNamespace(
remote_request_id="remote_req",
remote_engine_id="remote_engine",
remote_host="localhost",
remote_port=31000,
remote_pcp_size=2,
remote_dcp_size=2,
remote_ptp_size=4,
remote_multi_nodes_meta_mapping={},
remote_block_size=16,
local_block_ids=([10],),
remote_block_ids=([30],),
num_computed_tokens=0,
)
metadata = types.SimpleNamespace(reqs_in_batch=["req"], requests={"req": meta})
worker.start_load_kv(cast(MooncakeConnectorMetadata, metadata))
add_request_calls = worker.kv_recv_thread.add_request.call_args_list
self.assertEqual(len(add_request_calls), 2)
self.assertEqual(add_request_calls[0].kwargs["remote_handshake_port"], 31001)
self.assertEqual(add_request_calls[0].kwargs["local_block_ids_replicate_k"], ([40],))
self.assertEqual(add_request_calls[0].kwargs["remote_block_ids_replicate_k"], ([20],))
self.assertEqual(add_request_calls[1].kwargs["remote_handshake_port"], 31003)
self.assertIsNone(add_request_calls[1].kwargs["local_block_ids_replicate_k"])
self.assertIsNone(add_request_calls[1].kwargs["remote_block_ids_replicate_k"])
def test_get_kv_split_metadata_dp1_remote_port_send_num_uses_absolute_ports(self):
self.vllm_config.kv_transfer_config.kv_port = 30000
self.vllm_config.model_config.is_deepseek_mla = True
self.vllm_config.kv_transfer_config.get_from_extra_config.side_effect = lambda k, d: {
"prefill": {"tp_size": 8, "dp_size": 2, "pp_size": 1},
"decode": {"tp_size": 4, "dp_size": 4, "pp_size": 1},
}.get(k, d)
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
worker.use_mla = True
worker.pcp_size = 1
worker.dcp_size = 1
worker.tp_size = 4
worker.tp_rank = 0
worker.pcp_rank = 0
worker.dcp_rank = 0
worker.side_channel_port = 40000
worker.handshake_port = 40000
worker.local_remote_block_port_mapping = {}
worker.remote_port_send_num = {}
worker.block_size = 16
worker.num_key_value_heads = 1
worker.use_sparse = False
worker.block_size_scale = [[1]]
worker.kv_group2layeridx = {
0: (
{
"kv_cache_spec_type": "FullAttentionSpec",
"layer_names": ["model.layers.0.self_attn"],
},
[0],
)
}
remote_mapping = {
str(offset): {
"host": f"host-{offset}",
"engine_id": f"engine-{offset}",
"handshake_port": 30000 + offset,
}
for offset in range(8, 16)
}
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=8,
remote_ptp_size=8,
remote_port=30008,
remote_block_ids=(list(range(100, 103)),),
local_block_ids=(list(range(200, 224)),),
num_external_tokens=24 * worker.block_size,
num_prompt_blocks=24,
num_computed_tokens=0,
remote_engine_id="remote_engine",
remote_host="localhost",
remote_multi_nodes_meta_mapping=remote_mapping,
remote_block_size=16,
)
ports, _, _ = worker._get_kv_split_metadata("req_dp1", cast(ReqMeta, meta))
remote_port_send_num = worker.remote_port_send_num[meta.remote_engine_id]
self.assertEqual([port for shard in ports for port in shard], list(range(30008, 30016)))
self.assertEqual(set(remote_port_send_num), set(range(30008, 30016)))
self.assertNotIn(30016, remote_port_send_num)
self.assertEqual(remote_port_send_num[30008]["host"], "host-8")
self.assertEqual(remote_port_send_num[30015]["host"], "host-15")
self.assertEqual(
worker._get_remote_host_info_by_port(30008, 30015, "localhost", "remote_engine", remote_mapping),
("host-15", "engine-15"),
)
def test_get_sfa_replicated_indexer_block_ids_uses_full_blocks_for_prefix(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.enable_sfa_dcp_replicated_indexer = True
worker.pcp_size = 1
worker.dcp_size = 2
worker.block_size = 16
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=2,
remote_block_ids=([10, 11],),
local_block_ids=([20],),
local_full_block_ids=([19, 20],),
num_external_tokens=32,
num_prompt_blocks=3,
num_computed_tokens=16,
)
local_ids, remote_ids = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
self.assertEqual(local_ids, ([39, 40],))
self.assertEqual(remote_ids, ([21, 22],))
def test_get_sfa_replicated_indexer_block_ids_ignores_empty_regular_kv_shard(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.enable_sfa_dcp_replicated_indexer = True
worker.pcp_size = 1
worker.dcp_size = 2
worker.block_size = 16
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=2,
remote_block_ids=([10],),
local_block_ids=([],),
local_full_block_ids=([20],),
num_external_tokens=16,
num_prompt_blocks=1,
num_computed_tokens=0,
)
local_ids, remote_ids = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
self.assertEqual(local_ids, ([40],))
self.assertEqual(remote_ids, ([20],))
def test_get_sfa_replicated_indexer_block_ids_requires_full_blocks_for_prefix(self):
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
worker.enable_sfa_dcp_replicated_indexer = True
worker.pcp_size = 1
worker.dcp_size = 2
worker.block_size = 16
meta = types.SimpleNamespace(
remote_pcp_size=1,
remote_dcp_size=2,
remote_block_ids=([10, 11],),
local_block_ids=([20],),
local_full_block_ids=tuple(),
num_external_tokens=32,
num_prompt_blocks=3,
num_computed_tokens=16,
)
with self.assertRaises(AssertionError):
worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
if __name__ == "__main__":
unittest.main()