3173 lines
133 KiB
Python
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()
|