1173 lines
46 KiB
Python
1173 lines
46 KiB
Python
import contextlib
|
|
import importlib.util
|
|
import os
|
|
import sys
|
|
import threading
|
|
import types
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import torch
|
|
import zmq
|
|
|
|
fake_engine = types.ModuleType("mooncake.engine")
|
|
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
|
|
sys.modules["mooncake.engine"] = fake_engine
|
|
fake_torch_npu = types.ModuleType("torch_npu")
|
|
fake_torch_npu.__spec__ = importlib.util.spec_from_loader("torch_npu", loader=None)
|
|
fake_torch_npu.npu = MagicMock() # type: ignore[attr-defined]
|
|
fake_torch_npu.npu.current_device = MagicMock(return_value=0) # type: ignore[attr-defined]
|
|
fake_torch_npu.npu.Stream = MagicMock # type: ignore[attr-defined]
|
|
fake_torch_npu.npu_fusion_attention = MagicMock() # type: ignore[attr-defined]
|
|
sys.modules.setdefault("torch_npu", fake_torch_npu)
|
|
torch.npu = fake_torch_npu.npu # type: ignore[attr-defined]
|
|
fake_uvloop = types.ModuleType("uvloop")
|
|
fake_uvloop.__spec__ = importlib.util.spec_from_loader("uvloop", loader=None)
|
|
sys.modules.setdefault("uvloop", fake_uvloop)
|
|
|
|
# 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.
|
|
# We save the removed modules so we can restore them after our imports
|
|
# complete, 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)
|
|
|
|
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector import ( # noqa: E402
|
|
KVCacheRecvingLayerThread,
|
|
KVCacheSendingLayerThread,
|
|
KVConnectorRole,
|
|
LayerMetadata,
|
|
MooncakeAgentMetadata,
|
|
MooncakeLayerwiseConnector,
|
|
MooncakeLayerwiseConnectorMetadata,
|
|
MooncakeLayerwiseConnectorScheduler,
|
|
MooncakeLayerwiseConnectorWorker,
|
|
ReqMeta,
|
|
SendReqInfo,
|
|
SendTask,
|
|
ensure_zmq_recv,
|
|
ensure_zmq_send,
|
|
group_concurrent_contiguous,
|
|
string_to_int64_hash,
|
|
zmq_ctx,
|
|
)
|
|
|
|
# Restore the mocked modules so other test files still work correctly.
|
|
# For keys that our real import loaded, overwrite with the saved mock.
|
|
for _k, _v in _saved_modules.items():
|
|
sys.modules[_k] = _v
|
|
|
|
GET_META_MSG = b"get_meta_msg"
|
|
DONE_SENDING_MSG = b"done_sending_msg"
|
|
|
|
|
|
def _make_layer_metadata(**overrides):
|
|
defaults = dict(
|
|
tensor_group_idx=[0],
|
|
kv_caches_base_addr=[1000, 2000],
|
|
block_len=[1024],
|
|
block_size_scale=[1],
|
|
)
|
|
defaults.update(overrides)
|
|
return LayerMetadata(**defaults)
|
|
|
|
|
|
def _make_mock_kv_cache_config(block_size=16):
|
|
kv_cache_spec = MagicMock()
|
|
kv_cache_spec.block_size = block_size
|
|
group_spec = MagicMock()
|
|
group_spec.kv_cache_spec = kv_cache_spec
|
|
group_spec.layer_names = ["layer0"]
|
|
kv_cache_config = MagicMock()
|
|
kv_cache_config.kv_cache_groups = [group_spec]
|
|
return kv_cache_config
|
|
|
|
|
|
class TestKVCacheSendingLayerThread(unittest.TestCase):
|
|
def setUp(self):
|
|
self.engine = MagicMock()
|
|
self.engine.register_memory.return_value = 0
|
|
self.engine.batch_transfer_sync_write.return_value = 1
|
|
fake_stream = MagicMock(name="FakeStream")
|
|
fake_stream.synchronize = MagicMock()
|
|
|
|
self.first_kv_cache = torch.zeros((2, 2, 2, 8), dtype=torch.float32, device="cpu")
|
|
|
|
self.ready_event = threading.Event()
|
|
|
|
self.fake_k_buffer = MagicMock()
|
|
self.fake_v_buffer = MagicMock()
|
|
fake_resharding_stream = MagicMock()
|
|
|
|
self.layer_metadata = {
|
|
"layer0": _make_layer_metadata(
|
|
tensor_group_idx=[0],
|
|
kv_caches_base_addr=[1000, 2000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
"layer1": _make_layer_metadata(
|
|
tensor_group_idx=[0],
|
|
kv_caches_base_addr=[3000, 4000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
"layer2": _make_layer_metadata(
|
|
tensor_group_idx=[0],
|
|
kv_caches_base_addr=[5000, 6000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
}
|
|
|
|
self.vllm_config = MagicMock()
|
|
self.vllm_config.cache_config.mamba_cache_mode = None
|
|
self.vllm_config.speculative_config = None
|
|
|
|
self.kv_cache_config = _make_mock_kv_cache_config()
|
|
self.kv_cache_specs = [MagicMock(block_size=16)]
|
|
|
|
self.key = torch.zeros((4, 8), dtype=torch.float32)
|
|
self.value = torch.zeros((4, 8), dtype=torch.float32)
|
|
self.thread = KVCacheSendingLayerThread(
|
|
engine=self.engine,
|
|
vllm_config=self.vllm_config,
|
|
kv_cache_config=self.kv_cache_config,
|
|
kv_cache_specs=self.kv_cache_specs,
|
|
attn_resharding_group_idx=set(),
|
|
total_layers=3,
|
|
ready_event=self.ready_event,
|
|
tp_size=1,
|
|
tp_rank=0,
|
|
pd_head_ratio=1,
|
|
num_head_replica=1,
|
|
layer_metadata=self.layer_metadata,
|
|
use_mla=True,
|
|
use_attn_mamba_hybrid=False,
|
|
k_buffer=self.fake_k_buffer,
|
|
v_buffer=self.fake_v_buffer,
|
|
enable_kv_quant=False,
|
|
enable_c8_quant=False,
|
|
resharding_stream=fake_resharding_stream,
|
|
callback_func=MagicMock(),
|
|
)
|
|
|
|
self.req_meta_base = ReqMeta(
|
|
local_block_ids=[[5, 8]],
|
|
token_ids=[1, 2, 3],
|
|
remote_block_ids=[[10, 20]],
|
|
remote_block_size=[[16]],
|
|
remote_engine_id="remote_engine",
|
|
remote_host="127.0.0.1",
|
|
remote_port=7777,
|
|
remote_te_rpc_port=6000,
|
|
remote_layer_metadata={
|
|
"layer0": _make_layer_metadata(
|
|
kv_caches_base_addr=[4000, 8000],
|
|
block_len=[64, 64],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
},
|
|
metaserver="http://dummy",
|
|
remote_tp_size=8,
|
|
remote_pcp_size=1,
|
|
remote_dcp_size=1,
|
|
chunk_finish=False,
|
|
)
|
|
|
|
@patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.npu_stream_switch",
|
|
side_effect=lambda *_args, **_kwargs: contextlib.nullcontext(),
|
|
)
|
|
@patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.torch.Tensor.data_ptr",
|
|
autospec=True,
|
|
return_value=0x200000,
|
|
)
|
|
@patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.align_memory",
|
|
side_effect=lambda x, _align: x,
|
|
)
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.torch.npu.synchronize")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.group_concurrent_contiguous")
|
|
def test_transfer_pd_gt1_uses_buffers_and_calls_engine(
|
|
self, mock_group, _mock_sync, _mock_align, _mock_dataptr, mock_stream_switch
|
|
):
|
|
fake_resharding_stream = MagicMock()
|
|
|
|
layer_metadata = {
|
|
"layer0": _make_layer_metadata(
|
|
tensor_group_idx=[0],
|
|
kv_caches_base_addr=[1111, 2222],
|
|
block_len=[64, 64],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
}
|
|
|
|
vllm_config = MagicMock()
|
|
vllm_config.cache_config.mamba_cache_mode = None
|
|
vllm_config.speculative_config = None
|
|
|
|
kv_cache_config = _make_mock_kv_cache_config()
|
|
kv_cache_specs = [MagicMock(block_size=16)]
|
|
|
|
thread = KVCacheSendingLayerThread(
|
|
engine=self.engine,
|
|
vllm_config=vllm_config,
|
|
kv_cache_config=kv_cache_config,
|
|
kv_cache_specs=kv_cache_specs,
|
|
attn_resharding_group_idx=set(),
|
|
total_layers=2,
|
|
ready_event=self.ready_event,
|
|
tp_size=1,
|
|
tp_rank=0,
|
|
pd_head_ratio=2,
|
|
num_head_replica=1,
|
|
layer_metadata=layer_metadata,
|
|
use_mla=False,
|
|
use_attn_mamba_hybrid=False,
|
|
k_buffer=self.fake_k_buffer,
|
|
v_buffer=self.fake_v_buffer,
|
|
enable_kv_quant=False,
|
|
enable_c8_quant=False,
|
|
resharding_stream=fake_resharding_stream,
|
|
callback_func=MagicMock(),
|
|
)
|
|
|
|
req_meta = self.req_meta_base
|
|
req_meta.remote_block_ids = [[10, 20]]
|
|
req_meta.remote_layer_metadata = {
|
|
"layer0": _make_layer_metadata(
|
|
kv_caches_base_addr=[4000, 8000],
|
|
block_len=[64, 64],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
}
|
|
|
|
mock_group.return_value = ([[10, 11], [20, 21]], [])
|
|
key = torch.zeros((1, 8), dtype=torch.float32)
|
|
value = torch.zeros((1, 8), dtype=torch.float32)
|
|
|
|
send_task = SendTask(
|
|
send_request={"req1": req_meta},
|
|
wait_event=MagicMock(),
|
|
k_cache=key,
|
|
v_cache=value,
|
|
layer_idx=0,
|
|
layer_name="layer0",
|
|
group_rearrange_block_ids=[[5, 8]],
|
|
)
|
|
|
|
thread._transfer_kv_cache(send_task)
|
|
|
|
self.engine.batch_transfer_sync_write.assert_called_once()
|
|
session_id, src_list, dst_list, length_list = self.engine.batch_transfer_sync_write.call_args[0]
|
|
self.assertEqual(session_id, "127.0.0.1:6000")
|
|
|
|
self.assertEqual(len(src_list), 4)
|
|
self.assertEqual(len(dst_list), 4)
|
|
self.assertEqual(len(length_list), 4)
|
|
|
|
for L in length_list:
|
|
self.assertGreater(L, 0)
|
|
self.assertEqual(L % 64, 0)
|
|
|
|
remote_block_len = 64
|
|
expected_offsets = [10 * remote_block_len, 20 * remote_block_len]
|
|
self.assertEqual(dst_list[0] - 4000, expected_offsets[0])
|
|
self.assertEqual(dst_list[1] - 4000, expected_offsets[1])
|
|
self.assertEqual(dst_list[2] - 8000, expected_offsets[0])
|
|
self.assertEqual(dst_list[3] - 8000, expected_offsets[1])
|
|
|
|
def test_transfer_skips_when_no_local_blocks(self):
|
|
req_meta = self.req_meta_base
|
|
req_meta.local_block_ids = [[]]
|
|
send_task = SendTask(
|
|
send_request={"req2": req_meta},
|
|
wait_event=MagicMock(),
|
|
k_cache=torch.zeros((1, 8)),
|
|
v_cache=torch.zeros((1, 8)),
|
|
layer_idx=0,
|
|
layer_name="layer0",
|
|
group_rearrange_block_ids=[[]],
|
|
)
|
|
self.thread._transfer_kv_cache(send_task)
|
|
self.engine.batch_transfer_sync_write.assert_not_called()
|
|
|
|
@patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.group_concurrent_contiguous",
|
|
side_effect=group_concurrent_contiguous,
|
|
)
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.torch.npu.synchronize")
|
|
def test_callback_invoked_on_final_layer(self, _mock_sync, _mock_group):
|
|
req_meta = self.req_meta_base
|
|
req_meta.chunk_finish = True
|
|
req_meta.local_block_ids = [[5, 6]]
|
|
req_meta.remote_block_ids = [[10, 11]]
|
|
req_meta.remote_layer_metadata = {
|
|
"layer0": _make_layer_metadata(
|
|
kv_caches_base_addr=[7000, 8000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
"layer1": _make_layer_metadata(
|
|
kv_caches_base_addr=[9000, 10000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
"layer2": _make_layer_metadata(
|
|
kv_caches_base_addr=[11000, 12000],
|
|
block_len=[1024, 2048],
|
|
block_size_scale=[1, 1],
|
|
),
|
|
}
|
|
|
|
key = torch.zeros((1, 8), dtype=torch.float32)
|
|
value = torch.zeros((1, 8), dtype=torch.float32)
|
|
|
|
send_task = SendTask(
|
|
send_request={"req5": req_meta},
|
|
wait_event=MagicMock(),
|
|
k_cache=key,
|
|
v_cache=value,
|
|
layer_idx=2,
|
|
layer_name="layer2",
|
|
group_rearrange_block_ids=[[]],
|
|
)
|
|
self.thread._transfer_kv_cache(send_task)
|
|
|
|
self.thread.callback_func.assert_called_once()
|
|
|
|
|
|
class TestKVCacheRecvingLayerThread(unittest.TestCase):
|
|
def setUp(self):
|
|
self.meta = MooncakeAgentMetadata(
|
|
te_rpc_port=6000,
|
|
layer_metadata={"layer0": _make_layer_metadata()},
|
|
)
|
|
self.ready_event = threading.Event()
|
|
|
|
def test_get_and_clear_done_requests(self):
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=0,
|
|
side_channel_port=5555,
|
|
tp_size=2,
|
|
pd_head_ratio=1,
|
|
local_engine_id="engineA",
|
|
metadata=self.meta,
|
|
ready_event=self.ready_event,
|
|
)
|
|
|
|
with th.lock:
|
|
th.done_requests.update({"r1", "r2"})
|
|
got = th.get_and_clear_done_requests()
|
|
self.assertEqual(got, {"r1", "r2"})
|
|
|
|
got2 = th.get_and_clear_done_requests()
|
|
self.assertEqual(got2, set())
|
|
|
|
def test_get_and_clear_failed_requests(self):
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=0,
|
|
side_channel_port=5555,
|
|
tp_size=2,
|
|
pd_head_ratio=1,
|
|
local_engine_id="engineA",
|
|
metadata=self.meta,
|
|
ready_event=self.ready_event,
|
|
)
|
|
|
|
with th.lock:
|
|
th.failed_requests.update({"r1", "r2"})
|
|
got = th.get_and_clear_failed_requests()
|
|
self.assertEqual(got, {"r1", "r2"})
|
|
|
|
got2 = th.get_and_clear_failed_requests()
|
|
self.assertEqual(got2, set())
|
|
|
|
def test_update_failed_task_aggregates_by_pd_head_ratio(self):
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=0,
|
|
side_channel_port=5555,
|
|
tp_size=2,
|
|
pd_head_ratio=2,
|
|
local_engine_id="engineA",
|
|
metadata=self.meta,
|
|
ready_event=self.ready_event,
|
|
)
|
|
|
|
with th.lock:
|
|
th.task_tracker["reqX"] = set()
|
|
th.request_map = MagicMock()
|
|
|
|
th.update_failed_task("reqX")
|
|
with th.lock:
|
|
self.assertNotIn("reqX", th.task_tracker)
|
|
self.assertIn("reqX", th.failed_requests)
|
|
|
|
def test_update_done_task_aggregates_by_pd_head_ratio(self):
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=0,
|
|
side_channel_port=5555,
|
|
tp_size=2,
|
|
pd_head_ratio=2,
|
|
local_engine_id="engineA",
|
|
metadata=self.meta,
|
|
ready_event=self.ready_event,
|
|
)
|
|
|
|
with th.lock:
|
|
th.task_tracker["reqX"] = set()
|
|
|
|
th.update_done_task("reqX", 2, "path1")
|
|
with th.lock:
|
|
self.assertIn("reqX", th.task_tracker)
|
|
self.assertNotIn("reqX", th.done_requests)
|
|
|
|
th.update_done_task("reqX", 2, "path2")
|
|
with th.lock:
|
|
self.assertNotIn("reqX", th.task_tracker)
|
|
self.assertIn("reqX", th.done_requests)
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_ip", return_value="127.0.0.1")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.make_zmq_socket")
|
|
@patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.make_zmq_path",
|
|
side_effect=lambda proto, host, port: f"{proto}://{host}:{port}",
|
|
)
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.msgspec.msgpack.Decoder")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.msgspec.msgpack.Encoder")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.zmq_ctx")
|
|
def test_run_loop_handles_meta_done_invalid_unexpected_and_ack(
|
|
self, mock_zmq_ctx, mock_Encoder, mock_Decoder, _mock_make_path, _mock_make_sock, _mock_get_ip, mock_logger
|
|
):
|
|
enc_inst = MagicMock()
|
|
enc_inst.encode.return_value = b"ENCODED_META"
|
|
mock_Encoder.return_value = enc_inst
|
|
|
|
dec_inst = MagicMock()
|
|
dec_inst.decode.side_effect = [
|
|
(GET_META_MSG,),
|
|
(DONE_SENDING_MSG, "reqA", 1, "path1"),
|
|
(b"weird_msg",),
|
|
]
|
|
mock_Decoder.return_value = dec_inst
|
|
|
|
sock = MagicMock()
|
|
|
|
sock.recv_multipart.side_effect = [
|
|
[b"ID", b"SOME_PAYLOAD"],
|
|
[b"ID", b"SOME_PAYLOAD2"],
|
|
[b"ONLY_ID"],
|
|
[b"ID", b"SOME_PAYLOAD3"],
|
|
SystemExit,
|
|
]
|
|
|
|
cm = MagicMock()
|
|
cm.__enter__.return_value = sock
|
|
mock_zmq_ctx.return_value = cm
|
|
|
|
ready_event = threading.Event()
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=1,
|
|
side_channel_port=6000,
|
|
tp_size=2,
|
|
pd_head_ratio=1,
|
|
local_engine_id="engineZ",
|
|
metadata=self.meta,
|
|
ready_event=ready_event,
|
|
)
|
|
|
|
with th.lock:
|
|
th.task_tracker["reqA"] = set()
|
|
|
|
with self.assertRaises(SystemExit):
|
|
th.run()
|
|
|
|
self.assertTrue(ready_event.is_set())
|
|
|
|
self.assertGreaterEqual(sock.send_multipart.call_count, 2)
|
|
calls = [c.args for c in sock.send_multipart.call_args_list]
|
|
|
|
meta_call = calls[0]
|
|
self.assertEqual(meta_call[0][0], b"ID")
|
|
self.assertEqual(meta_call[0][1], b"")
|
|
self.assertEqual(meta_call[0][2], b"ENCODED_META")
|
|
|
|
ack_call = calls[1]
|
|
self.assertEqual(ack_call[0][0], b"ID")
|
|
self.assertEqual(ack_call[0][1], b"")
|
|
self.assertEqual(ack_call[0][2], b"ACK")
|
|
|
|
self.assertTrue(mock_logger.error.called)
|
|
|
|
finished = th.get_and_clear_done_requests()
|
|
self.assertIn("reqA", finished)
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_ip", return_value="127.0.0.1")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.msgspec.msgpack.Decoder")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.msgspec.msgpack.Encoder")
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.zmq_ctx")
|
|
def test_run_loop_pd_head_ratio_gt1_requires_multiple_done(
|
|
self, mock_zmq_ctx, mock_Encoder, mock_Decoder, _mock_get_ip, _mock_logger
|
|
):
|
|
enc_inst = MagicMock()
|
|
enc_inst.encode.return_value = b"ENC"
|
|
mock_Encoder.return_value = enc_inst
|
|
|
|
dec_inst = MagicMock()
|
|
dec_inst.decode.side_effect = [
|
|
(DONE_SENDING_MSG, "reqB", 2, "path1"),
|
|
(DONE_SENDING_MSG, "reqB", 2, "path2"),
|
|
]
|
|
mock_Decoder.return_value = dec_inst
|
|
|
|
sock = MagicMock()
|
|
sock.recv_multipart.side_effect = [
|
|
[b"ID", b"PAY1"],
|
|
[b"ID", b"PAY2"],
|
|
SystemExit,
|
|
]
|
|
cm = MagicMock()
|
|
cm.__enter__.return_value = sock
|
|
mock_zmq_ctx.return_value = cm
|
|
|
|
th = KVCacheRecvingLayerThread(
|
|
tp_rank=0,
|
|
side_channel_port=5555,
|
|
tp_size=2,
|
|
pd_head_ratio=2,
|
|
local_engine_id="engineY",
|
|
metadata=self.meta,
|
|
ready_event=self.ready_event,
|
|
)
|
|
with th.lock:
|
|
th.task_tracker["reqB"] = set()
|
|
with self.assertRaises(SystemExit):
|
|
th.run()
|
|
finished = th.get_and_clear_done_requests()
|
|
self.assertIn("reqB", finished)
|
|
|
|
|
|
class MockVllmConfig:
|
|
def __init__(self):
|
|
self.model_config = MagicMock()
|
|
self.parallel_config = MagicMock()
|
|
self.cache_config = MagicMock()
|
|
self.kv_transfer_config = MagicMock()
|
|
self.speculative_config = None
|
|
self.quant_config = None
|
|
self.model_config.use_mla = True
|
|
self.parallel_config.tensor_parallel_size = 2
|
|
self.parallel_config.data_parallel_rank_local = 0
|
|
self.parallel_config.data_parallel_size_local = 1
|
|
self.parallel_config.data_parallel_size = 1
|
|
self.parallel_config.data_parallel_rank = 0
|
|
self.parallel_config.prefill_context_parallel_size = 1
|
|
self.parallel_config.decode_context_parallel_size = 1
|
|
self.cache_config.block_size = 16
|
|
self.cache_config.mamba_cache_mode = None
|
|
self.model_config.hf_config.num_key_value_heads = 1
|
|
self.model_config.get_num_layers = MagicMock(return_value=1)
|
|
self.model_config.get_total_num_kv_heads = MagicMock(return_value=1)
|
|
self.model_config.hf_text_config = MagicMock()
|
|
self.model_config.hf_text_config.model_type = "default"
|
|
|
|
self.kv_transfer_config.engine_id = "test_engine"
|
|
self.kv_transfer_config.kv_port = 5000
|
|
self.kv_transfer_config.is_kv_producer = True
|
|
self.kv_transfer_config.is_kv_consumer = False
|
|
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},
|
|
"decode": {"tp_size": 2, "dp_size": 1},
|
|
}.get(k, d)
|
|
|
|
|
|
class MockKVCacheConfig:
|
|
def __init__(self, block_size=16):
|
|
kv_cache_spec = MagicMock()
|
|
kv_cache_spec.block_size = block_size
|
|
group_spec = MagicMock()
|
|
group_spec.kv_cache_spec = kv_cache_spec
|
|
group_spec.layer_names = ["encoder.layer.0"]
|
|
self.kv_cache_groups = [group_spec]
|
|
self.kv_cache_tensors = []
|
|
self.num_blocks = 10
|
|
|
|
|
|
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.prompt_embeds = None
|
|
self.kv_transfer_params = kv_transfer_params or {}
|
|
self.status = status or "running"
|
|
self.output_token_ids = [101, 102]
|
|
self.num_computed_tokens = 0
|
|
self.num_prompt_tokens = len(self.prompt_token_ids)
|
|
self.max_tokens = 16
|
|
|
|
self.all_token_ids = list(self.prompt_token_ids)
|
|
self._all_token_ids = list(self.prompt_token_ids)
|
|
|
|
|
|
class TestMooncakeLayerwiseConnectorMetadata(unittest.TestCase):
|
|
def test_add_new_req(self):
|
|
meta = MooncakeLayerwiseConnectorMetadata()
|
|
self.assertEqual(len(meta.requests), 0)
|
|
|
|
meta.add_new_req(
|
|
request_id="req1",
|
|
local_block_ids=[[1, 2, 3]],
|
|
kv_transfer_params={
|
|
"remote_block_ids": [[4, 5, 6]],
|
|
"remote_block_size": [[16]],
|
|
"remote_engine_id": "remote_engine",
|
|
"remote_host": "localhost",
|
|
"remote_port": 5000,
|
|
},
|
|
)
|
|
|
|
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.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)
|
|
|
|
|
|
class TestMooncakeLayerwiseConnectorSchedulerMatchedTokens(unittest.TestCase):
|
|
def setUp(self):
|
|
config = MockVllmConfig()
|
|
kv_cache_config = MockKVCacheConfig()
|
|
self.scheduler = MooncakeLayerwiseConnectorScheduler(config, kv_cache_config, "test_engine")
|
|
|
|
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)
|
|
|
|
def test_get_num_new_matched_tokens_hybrid_excludes_last_token(self):
|
|
self.scheduler.need_truncate = True
|
|
request = MockRequest("req1", prompt_token_ids=list(range(17)), kv_transfer_params={"do_remote_prefill": True})
|
|
|
|
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
|
|
|
|
self.assertEqual(tokens, 16)
|
|
self.assertTrue(async_flag)
|
|
|
|
def test_get_num_new_matched_tokens_hybrid_truncates_prefill_request(self):
|
|
self.scheduler.need_truncate = True
|
|
request = MockRequest("req1", prompt_token_ids=list(range(4)), kv_transfer_params={"do_remote_decode": True})
|
|
|
|
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
|
|
|
|
self.assertEqual(tokens, 0)
|
|
self.assertFalse(async_flag)
|
|
self.assertEqual(request.prompt_token_ids, [0, 1, 2])
|
|
self.assertEqual(request._all_token_ids, [0, 1, 2])
|
|
self.assertEqual(request.num_prompt_tokens, 3)
|
|
self.assertEqual(request.max_tokens, 1)
|
|
self.assertTrue(request.kv_transfer_params["_p_side_truncated"])
|
|
|
|
def test_build_connector_meta(self):
|
|
self.scheduler.vllm_config.kv_transfer_config.is_kv_consumer = True
|
|
request = MockRequest("req1")
|
|
|
|
self.scheduler._reqs_need_recv["req1"] = (request, [], [[4, 5, 6]])
|
|
request.kv_transfer_params = {
|
|
"remote_block_ids": [[1, 2, 3]],
|
|
"remote_block_size": [[16]],
|
|
"remote_engine_id": "remote",
|
|
"remote_host": "localhost",
|
|
"remote_port": 5000,
|
|
}
|
|
|
|
meta = self.scheduler.build_connector_meta(MagicMock())
|
|
self.assertIsInstance(meta, MooncakeLayerwiseConnectorMetadata)
|
|
self.assertEqual(len(meta.requests), 1)
|
|
self.assertEqual(meta.requests["req1"].local_block_ids, [[4, 5, 6]])
|
|
self.assertEqual(meta.requests["req1"].remote_block_ids, [[1, 2, 3]])
|
|
self.assertEqual(len(self.scheduler._reqs_need_recv), 0)
|
|
|
|
def test_update_state_after_alloc_hybrid_trims_remote_block_with_only_last_token(self):
|
|
self.scheduler.need_truncate = True
|
|
request = MockRequest(
|
|
"req1",
|
|
prompt_token_ids=list(range(17)),
|
|
kv_transfer_params={"do_remote_prefill": True, "metaserver": "http://meta"},
|
|
)
|
|
blocks = _MockBlocks(unhashed=[], block_ids_tuple=([4, 5],))
|
|
self.scheduler.executor.submit = MagicMock()
|
|
|
|
self.scheduler.update_state_after_alloc(request, blocks, num_external_tokens=16)
|
|
|
|
_, kwargs = self.scheduler.executor.submit.call_args
|
|
self.assertEqual(kwargs["message"]["remote_block_ids"], ([4],))
|
|
|
|
|
|
class _MockBlocks:
|
|
def __init__(self, unhashed, block_ids_tuple=None):
|
|
self._unhashed = list(unhashed)
|
|
self._block_ids_tuple = block_ids_tuple if block_ids_tuple is not None else ([1, 2],)
|
|
|
|
def get_unhashed_block_ids(self):
|
|
return list(self._unhashed)
|
|
|
|
def get_block_ids(self):
|
|
return self._block_ids_tuple
|
|
|
|
|
|
class _MockSchedulerOutput:
|
|
def __init__(
|
|
self,
|
|
cached_req_ids=None,
|
|
cached_new_block_ids=None,
|
|
cached_num_computed=None,
|
|
new_reqs=None,
|
|
num_sched=None,
|
|
scheduled_spec_decode_tokens=None,
|
|
):
|
|
self.scheduled_cached_reqs = SimpleNamespace(
|
|
req_ids=cached_req_ids or [],
|
|
new_block_ids=cached_new_block_ids or [],
|
|
num_computed_tokens=cached_num_computed or [],
|
|
)
|
|
self.scheduled_spec_decode_tokens = scheduled_spec_decode_tokens or {}
|
|
self.scheduled_new_reqs = new_reqs or []
|
|
self.num_scheduled_tokens = num_sched or {}
|
|
|
|
|
|
class TestMooncakeLayerwiseConnectorScheduler_More(unittest.TestCase):
|
|
def setUp(self):
|
|
self.config = MockVllmConfig()
|
|
self.kv_cache_config = MockKVCacheConfig()
|
|
self.scheduler = MooncakeLayerwiseConnectorScheduler(self.config, self.kv_cache_config, "test_engine")
|
|
|
|
def test_get_num_new_matched_tokens_with_prefill_block_aligned(self):
|
|
req = MockRequest(
|
|
"req_prefill", prompt_token_ids=list(range(32)), kv_transfer_params={"do_remote_prefill": True}
|
|
)
|
|
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(req, num_computed_tokens=16)
|
|
self.assertEqual(tokens, 16)
|
|
self.assertTrue(async_flag)
|
|
|
|
def test_update_state_after_alloc_prefill_records_and_resets_flag(self):
|
|
req = MockRequest("req_u1", prompt_token_ids=list(range(24)), kv_transfer_params={"do_remote_prefill": True})
|
|
req.num_computed_tokens = 0
|
|
blocks = _MockBlocks(unhashed=[4, 5, 6], block_ids_tuple=([[4, 5, 6]],))
|
|
|
|
self.scheduler.update_state_after_alloc(req, blocks, num_external_tokens=8)
|
|
self.assertIn("req_u1", self.scheduler._reqs_need_recv)
|
|
record = self.scheduler._reqs_need_recv["req_u1"]
|
|
self.assertIs(record[0], req)
|
|
self.assertEqual(record[1], [])
|
|
self.assertEqual(record[2], ([[4, 5, 6]],))
|
|
self.assertFalse(req.kv_transfer_params.get("do_remote_prefill", True))
|
|
|
|
def test_update_state_after_alloc_decode_records_send_layerwise(self):
|
|
req = MockRequest(
|
|
"req_u2",
|
|
prompt_token_ids=list(range(10)),
|
|
kv_transfer_params={"do_remote_decode": True, "remote_block_ids": [], "remote_cached_tokens": 0},
|
|
)
|
|
blocks = _MockBlocks(unhashed=[], block_ids_tuple=([[7, 8, 9]],))
|
|
self.scheduler.update_state_after_alloc(req, blocks, num_external_tokens=0)
|
|
self.assertIn("req_u2", self.scheduler._reqs_need_send_layerwise)
|
|
info = self.scheduler._reqs_need_send_layerwise["req_u2"]
|
|
self.assertEqual(info.local_block_ids, [[[7, 8, 9]]])
|
|
self.assertIs(info.request, req)
|
|
|
|
def test_build_connector_meta_consumes_reqs_need_recv_and_clears(self):
|
|
self.scheduler.vllm_config.kv_transfer_config.is_kv_consumer = True
|
|
req = MockRequest(
|
|
"req_b1",
|
|
kv_transfer_params={
|
|
"remote_block_ids": [[1, 2]],
|
|
"remote_block_size": [[16]],
|
|
"remote_engine_id": "E",
|
|
"remote_host": "H",
|
|
"remote_port": 5555,
|
|
"remote_te_rpc_port": 6000,
|
|
"remote_layer_metadata": {"layer0": _make_layer_metadata()},
|
|
},
|
|
)
|
|
self.scheduler._reqs_need_recv["req_b1"] = (req, [], [[100, 101]])
|
|
meta = self.scheduler.build_connector_meta(_MockSchedulerOutput())
|
|
self.assertIsInstance(meta, MooncakeLayerwiseConnectorMetadata)
|
|
self.assertIn("req_b1", meta.requests)
|
|
self.assertEqual(meta.requests["req_b1"].local_block_ids, [[100, 101]])
|
|
self.assertEqual(len(self.scheduler._reqs_need_recv), 0)
|
|
|
|
def test_build_connector_meta_accumulates_cached_blocks(self):
|
|
req_meta = MagicMock(spec=SendReqInfo)
|
|
req_meta.local_block_ids = [[1, 2, 3]]
|
|
req_meta.local_transferred_tokens = 50
|
|
req_meta.local_computed_tokens = 75
|
|
req_meta.request = MagicMock()
|
|
req_meta.extend_local_block_ids = MagicMock()
|
|
req_meta.update_computed_tokens = MagicMock()
|
|
req_meta.update_transferred_tokens = MagicMock()
|
|
req_meta.unpack = MagicMock(
|
|
return_value=(
|
|
req_meta.local_block_ids,
|
|
req_meta.local_transferred_tokens,
|
|
req_meta.local_computed_tokens,
|
|
req_meta.request,
|
|
)
|
|
)
|
|
|
|
self.scheduler._reqs_need_send_layerwise["req_b2"] = req_meta
|
|
|
|
out = _MockSchedulerOutput(
|
|
cached_req_ids=["req_b2"],
|
|
cached_new_block_ids=[([[3, 4]],)],
|
|
cached_num_computed=[4],
|
|
new_reqs=[],
|
|
num_sched={},
|
|
)
|
|
meta = self.scheduler.build_connector_meta(out)
|
|
self.assertEqual(len(meta.requests), 0)
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.group_concurrent_contiguous")
|
|
def test_build_connector_meta_emits_when_tokens_reach_total(self, mock_group_concurrent_contiguous):
|
|
send_req_info = MagicMock(spec=SendReqInfo)
|
|
send_req_info.local_block_ids = [[1, 2, 3]]
|
|
send_req_info.local_transferred_tokens = 50
|
|
send_req_info.local_computed_tokens = 75
|
|
send_req_info.request = MagicMock()
|
|
send_req_info.request.kv_transfer_params = {
|
|
"remote_block_ids": [[4, 5]],
|
|
"remote_block_size": [[16]],
|
|
"remote_cached_tokens": 100,
|
|
}
|
|
send_req_info.request.all_token_ids = list(range(80))
|
|
send_req_info.extend_local_block_ids = MagicMock()
|
|
send_req_info.update_computed_tokens = MagicMock()
|
|
send_req_info.update_transferred_tokens = MagicMock()
|
|
send_req_info.unpack = MagicMock(
|
|
return_value=(
|
|
send_req_info.local_block_ids,
|
|
send_req_info.local_transferred_tokens,
|
|
send_req_info.local_computed_tokens,
|
|
send_req_info.request,
|
|
)
|
|
)
|
|
|
|
self.scheduler._reqs_need_send_layerwise["req_b3"] = send_req_info
|
|
out = _MockSchedulerOutput(
|
|
cached_req_ids=["req_b3"],
|
|
cached_new_block_ids=[([[50]],)],
|
|
cached_num_computed=[8],
|
|
new_reqs=[MagicMock(req_id="other", num_computed_tokens=0)],
|
|
num_sched={"req_b3": 4},
|
|
)
|
|
meta = self.scheduler.build_connector_meta(out)
|
|
send_req_info.extend_local_block_ids.assert_called_once_with(([[50]],))
|
|
self.assertIn("req_b3", meta.requests)
|
|
|
|
def test_request_finished_returns_false_none(self):
|
|
ok, params = self.scheduler.request_finished(MockRequest("req_fin"), [1, 2])
|
|
self.assertFalse(ok)
|
|
self.assertIsNone(params)
|
|
|
|
|
|
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_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)
|
|
|
|
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_layerwise_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_layerwise_connector.logger")
|
|
def test_ensure_zmq_send_success(self, _):
|
|
mock_socket = MagicMock()
|
|
path = "127.0.0.1:12345"
|
|
ensure_zmq_send(mock_socket, b"hello", path)
|
|
mock_socket.send.assert_called_once_with(b"hello")
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger")
|
|
def test_ensure_zmq_send_retry_and_fail(self, _):
|
|
mock_socket = MagicMock()
|
|
path = "127.0.0.1:12345"
|
|
mock_socket.send.side_effect = zmq.ZMQError( # type: ignore
|
|
"send failed"
|
|
)
|
|
with self.assertRaises(RuntimeError):
|
|
ensure_zmq_send(mock_socket, b"hello", path, max_retries=2)
|
|
self.assertEqual(mock_socket.send.call_count, 2)
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger")
|
|
def test_ensure_zmq_recv_success(self, _):
|
|
mock_socket = MagicMock()
|
|
mock_socket.recv.return_value = b"response"
|
|
mock_poller = MagicMock()
|
|
mock_poller.poll.return_value = [
|
|
(mock_socket, zmq.POLLIN) # type: ignore
|
|
]
|
|
path = "127.0.0.1:12345"
|
|
data = ensure_zmq_recv(mock_socket, mock_poller, path)
|
|
self.assertEqual(data, b"response")
|
|
|
|
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger")
|
|
def test_ensure_zmq_recv_timeout_and_fail(self, _):
|
|
mock_socket = MagicMock()
|
|
mock_poller = MagicMock()
|
|
mock_poller.poll.return_value = []
|
|
path = "127.0.0.1:12345"
|
|
with self.assertRaises(RuntimeError):
|
|
ensure_zmq_recv(mock_socket, mock_poller, path, timeout=0.01, max_retries=2)
|
|
|
|
|
|
class TestMooncakeLayerwiseConnectorForScheduler(unittest.TestCase):
|
|
def _make_config(self):
|
|
config = MockVllmConfig()
|
|
kv_cache_config = MockKVCacheConfig()
|
|
return config, kv_cache_config
|
|
|
|
def test_scheduler_role(self):
|
|
config, kv_cache_config = self._make_config()
|
|
connector = MooncakeLayerwiseConnector(config, KVConnectorRole.SCHEDULER, kv_cache_config)
|
|
self.assertIsNotNone(connector.connector_scheduler)
|
|
self.assertIsNone(connector.connector_worker)
|
|
|
|
@patch.object(MooncakeLayerwiseConnectorScheduler, "get_num_new_matched_tokens")
|
|
def test_scheduler_methods(self, mock_method):
|
|
config, kv_cache_config = self._make_config()
|
|
connector = MooncakeLayerwiseConnector(config, KVConnectorRole.SCHEDULER, kv_cache_config)
|
|
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]
|
|
|
|
|
|
class MockSchedulerOutput:
|
|
pass
|
|
|
|
|
|
class MockForwardContext:
|
|
pass
|
|
|
|
|
|
class TestMooncakeLayerwiseConnector(unittest.TestCase):
|
|
def setUp(self):
|
|
self.config = MockVllmConfig()
|
|
self.kv_cache_config = MockKVCacheConfig()
|
|
os.environ["ASCEND_RT_VISIBLE_DEVICES"] = "0,1"
|
|
|
|
def test_scheduler_initialization(self):
|
|
connector = MooncakeLayerwiseConnector(self.config, KVConnectorRole.SCHEDULER, self.kv_cache_config)
|
|
self.assertIsNotNone(connector.connector_scheduler)
|
|
self.assertIsNone(connector.connector_worker)
|
|
|
|
@patch.object(MooncakeLayerwiseConnectorScheduler, "get_num_new_matched_tokens")
|
|
def test_get_num_new_matched_tokens(self, mock_method):
|
|
connector = MooncakeLayerwiseConnector(self.config, KVConnectorRole.SCHEDULER, self.kv_cache_config)
|
|
request = MockRequest("req1")
|
|
connector.get_num_new_matched_tokens(request, 0)
|
|
mock_method.assert_called_once_with(request, 0)
|
|
|
|
@patch.object(MooncakeLayerwiseConnectorScheduler, "update_state_after_alloc")
|
|
def test_update_state_after_alloc(self, mock_method):
|
|
connector = MooncakeLayerwiseConnector(self.config, KVConnectorRole.SCHEDULER, self.kv_cache_config)
|
|
request = MockRequest("req1")
|
|
blocks = MockKVCacheBlocks()
|
|
connector.update_state_after_alloc(request, blocks, 3)
|
|
mock_method.assert_called_once_with(request, blocks, 3)
|
|
|
|
@patch.object(MooncakeLayerwiseConnectorScheduler, "build_connector_meta")
|
|
def test_build_connector_meta(self, mock_method):
|
|
connector = MooncakeLayerwiseConnector(self.config, KVConnectorRole.SCHEDULER, self.kv_cache_config)
|
|
scheduler_output = MockSchedulerOutput()
|
|
connector.build_connector_meta(scheduler_output)
|
|
mock_method.assert_called_once_with(scheduler_output)
|
|
|
|
@patch.object(MooncakeLayerwiseConnectorScheduler, "request_finished")
|
|
def test_request_finished(self, mock_method):
|
|
connector = MooncakeLayerwiseConnector(self.config, KVConnectorRole.SCHEDULER, self.kv_cache_config)
|
|
request = MockRequest("req1")
|
|
connector.request_finished(request, [1, 2, 3])
|
|
mock_method.assert_called_once_with(request, [1, 2, 3])
|
|
|
|
|
|
class TestMooncakeLayerwiseConnectorWorker(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("torch.Tensor.size", return_value=(10, 16, 8, 16)),
|
|
patch("torch.Tensor.element_size", return_value=4),
|
|
patch("torch.Tensor.data_ptr", return_value=0x1000),
|
|
patch("math.prod", return_value=128),
|
|
patch("random.Random"),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_tensor_model_parallel_rank",
|
|
return_value=0,
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_tp_group",
|
|
return_value=None,
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_ip",
|
|
return_value="127.0.0.1",
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.string_to_int64_hash",
|
|
side_effect=lambda s: hash(s),
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.global_te.get_transfer_engine",
|
|
return_value=self.mock_transfer_engine,
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.global_te.register_buffer",
|
|
return_value=None,
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.KVCacheSendingLayerThread",
|
|
MagicMock(),
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.KVCacheRecvingLayerThread",
|
|
MagicMock(),
|
|
),
|
|
patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.logger", MagicMock()),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.threading.Event", MagicMock()
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_ascend_config",
|
|
return_value=SimpleNamespace(pd_tp_ratio=1, num_head_replica=1, pd_head_ratio=1),
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_pcp_group",
|
|
),
|
|
patch(
|
|
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_layerwise_connector.get_decode_context_model_parallel_rank",
|
|
return_value=0,
|
|
),
|
|
]
|
|
|
|
for p in self.patches:
|
|
p.start() # type: ignore
|
|
|
|
self.vllm_config = MockVllmConfig()
|
|
self.engine_id = "test_engine"
|
|
mock_k = MagicMock()
|
|
mock_k.shape = (10, 16, 8, 16)
|
|
mock_k.data_ptr.return_value = 0x1000
|
|
mock_k.element_size.return_value = 4
|
|
mock_v = MagicMock()
|
|
mock_v.shape = (10, 16, 8, 16)
|
|
mock_v.data_ptr.return_value = 0x2000
|
|
mock_v.element_size.return_value = 4
|
|
self.kv_caches = {"encoder.layer.0": (mock_k, mock_v)}
|
|
self.vllm_config.parallel_config.tensor_parallel_size = 1
|
|
self.vllm_config.parallel_config.prefill_context_parallel_size = 1
|
|
self.vllm_config.parallel_config.decode_context_parallel_size = 1
|
|
self.vllm_config.parallel_config.data_parallel_rank = 0
|
|
self.vllm_config.kv_transfer_config.kv_port = 1234
|
|
|
|
self.kv_cache_config = MockKVCacheConfig()
|
|
|
|
def tearDown(self):
|
|
for p in self.patches:
|
|
p.stop() # type: ignore
|
|
|
|
def test_register_kv_caches_producer(self):
|
|
self.vllm_config.kv_transfer_config.is_kv_producer = True
|
|
self.vllm_config.kv_transfer_config.is_kv_consumer = False
|
|
worker = MooncakeLayerwiseConnectorWorker(self.vllm_config, self.kv_cache_config, self.engine_id)
|
|
worker.register_kv_caches(self.kv_caches)
|
|
self.assertEqual(len(worker.layer_metadata), 1)
|
|
self.assertIsNotNone(worker.kv_send_layer_thread)
|
|
self.assertIsNone(worker.kv_recv_layer_thread)
|
|
|
|
def test_register_kv_caches_consumer(self):
|
|
self.vllm_config.kv_transfer_config.is_kv_producer = False
|
|
self.vllm_config.kv_transfer_config.is_kv_consumer = True
|
|
worker = MooncakeLayerwiseConnectorWorker(self.vllm_config, self.kv_cache_config, self.engine_id)
|
|
worker.register_kv_caches(self.kv_caches)
|
|
self.assertEqual(len(worker.layer_metadata), 1)
|
|
self.assertIsNone(worker.kv_send_layer_thread)
|
|
self.assertIsNotNone(worker.kv_recv_layer_thread)
|
|
|
|
def test_register_kv_caches_mla_case(self):
|
|
mla_cache1 = MagicMock()
|
|
mla_cache1.size.return_value = (10, 16, 1, 16)
|
|
mla_cache1.shape = (10, 16, 1, 16)
|
|
mla_cache1.data_ptr.return_value = 0x1000
|
|
mla_cache1.element_size.return_value = 4
|
|
mla_cache2 = MagicMock()
|
|
mla_cache2.size.return_value = (10, 16, 1, 8)
|
|
mla_cache2.shape = (10, 16, 1, 8)
|
|
mla_cache2.data_ptr.return_value = 0x2000
|
|
mla_cache2.element_size.return_value = 4
|
|
mla_caches = {"encoder.layer.0": (mla_cache1, mla_cache2)}
|
|
worker = MooncakeLayerwiseConnectorWorker(self.vllm_config, self.kv_cache_config, self.engine_id)
|
|
worker.register_kv_caches(mla_caches)
|
|
self.assertTrue(worker.use_mla)
|
|
self.assertEqual(len(worker.layer_metadata["encoder.layer.0"].block_len), 2)
|