Skip to content

Commit fd2aa3f

Browse files
authored
[Performance][PD] Parallelize KV cache receive via ThreadPoolExecutor (vllm-project#10548)
### What this PR does / why we need it? ## Summary This PR improves KV cache receive throughput in PD disaggregation by running independent receive work concurrently while preserving FIFO ordering for the same remote peer. The original `KVCacheRecvingThread` processed every transfer task sequentially: ```python request_data = self.request_queue.get() self._handle_request(request_data) ``` That makes unrelated transfers block behind each other. Under concurrent PD traffic, tasks from different prefill ranks can queue behind a slow transfer even though they target independent remote peers and can safely progress in parallel. The final implementation dispatches through a per-peer FIFO scheduler: - requests for the same `(remote_host, remote_handshake_port)` are serialized; - requests for different peers can run concurrently in the existing `ThreadPoolExecutor(max_workers=32)`; - a busy peer yields after a small batch so it cannot monopolize executor workers; - request completion accounting waits for all submitted transfer tasks before marking the request done. ## Main changes 1. **Parallel KV receive with per-peer FIFO** `KVCacheRecvingThread.run()` now records a request as submitted and queues it by remote peer. Each active peer has at most one handler running, so REQ/REP ordering for that peer is preserved while independent peers can transfer concurrently. 2. **Bounded peer handler batch** `MAX_REQUESTS_PER_PEER_HANDLER = 5` prevents one busy peer from holding an executor worker forever when other peers are already waiting. 3. **Thread-safe shared state** - removes the shared `zmq.Poller` receive path and uses per-socket `RCVTIMEO`; - protects `proc_not_transfer_request` with a lock; - closes and discards failed REQ sockets instead of returning invalid sockets to the pool; - protects remote metadata cache population/readout with `remote_metadata_lock`; - tracks pending transfer tasks per request before calling `task_tracker.update_done_task_count()`. ## Latest E2E validation Test date: 2026-06-29. Setup: - 4-node DeepSeek V4 PD deployment, 32 NPUs total - Prefill: DP4TP8 - Decode: TP4DP8 - Requests sent from `dev0` to `localhost:9000`; proxy forwards to prefill/decode - Workload: 500 requests, max concurrency 48, input length 256 - Baseline: PR changes fully reverted to merge-base for comparison - Current: this PR head - Timing data was collected with temporary measurement scaffolding during validation; that scaffolding is not included in the final patch. ### Request-level result | Metric | Baseline | Current | Change | |---|---:|---:|---:| | OK | 500/500 | 500/500 | no errors | | Total time | 12.6s | 12.5s | -0.8% | | Request p50 | 0.96s | 0.97s | roughly flat | | Request p99 | 3.44s | 3.10s | -9.9% | ### KV receive timing Aggregate across decode nodes, 2000 KV receive tasks per run: | Metric | Baseline | Current | Change | |---|---:|---:|---:| | queue p50 | 18.14ms | 0.70ms | improved | | queue p90 | 237.79ms | 1.87ms | improved | | queue p99 | 603.63ms | 2.64ms | -99.56% | | lifecycle p99 | 735.16ms | 512.80ms | -30.25% | | active handler max | 1 | 7 | parallel receive active | Per-node queue p99: | Decode node | Baseline | Current | Change | |---|---:|---:|---:| | dev2 | 602.23ms | 4.11ms | -99.32% | | dev3 | 595.88ms | 2.41ms | -99.60% | The queue tail result is consistent with the earlier validation data: queue p99 drops from hundreds of milliseconds to single-digit milliseconds, showing that the tail was caused by serialized receive handling. ## How was this patch tested? - `pre-commit run ruff-check --all-files` - `pre-commit run ruff-format --all-files` - `git diff --check` - `python3 -m py_compile vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_hybrid_connector.py` - `python3 -m pytest tests/ut/kv_offload/test_mooncake_connector.py tests/ut/kv_offload/test_mooncake_hybrid_connector.py -q` - Result: 86 passed, 16 warnings - Full 4-node E2E restart and proxy benchmark: - baseline: 500/500 OK - current: 500/500 OK - Log scan after E2E found no `ZMQError`, `EFSM`, `Traceback`, duplicate done signal, request loss, or address conflict errors. ### Does this PR introduce any user-facing change? No. This changes internal KV receive scheduling and synchronization only. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@dc68bd8 --------- Signed-off-by: Zheng Shoujian <zheng.shoujian@outlook.com>
1 parent 2a0b0b3 commit fd2aa3f

4 files changed

Lines changed: 480 additions & 94 deletions

File tree

tests/ut/kv_offload/test_mooncake_connector.py

Lines changed: 111 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,7 @@
5858
patch("vllm.distributed.parallel_state._DCP", _mock_dcp_group).start()
5959

6060
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector import ( # noqa: E402
61+
MAX_REQUESTS_PER_PEER_HANDLER,
6162
GroupPull,
6263
KVCacheRecvingThread,
6364
KVCacheSendingThread,
@@ -510,6 +511,99 @@ def test_get_finished_requests(self, mock_tracker):
510511
result = self.thread.get_and_clear_finished_requests()
511512
self.assertEqual(result, {"req1", "req2"})
512513

514+
def test_submit_request_serializes_same_peer_fifo(self):
515+
release_first_request = threading.Event()
516+
first_request_started = threading.Event()
517+
other_peer_started = threading.Event()
518+
handled_requests: list[str] = []
519+
active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
520+
max_active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
521+
state_lock = threading.Lock()
522+
523+
def handle_request(req_meta: dict[str, Any]):
524+
peer_key = (req_meta["remote_host"], req_meta["remote_handshake_port"])
525+
with state_lock:
526+
active_by_peer[peer_key] += 1
527+
max_active_by_peer[peer_key] = max(max_active_by_peer[peer_key], active_by_peer[peer_key])
528+
handled_requests.append(req_meta["request_id"])
529+
530+
if req_meta["request_id"] == "same-peer-1":
531+
first_request_started.set()
532+
self.assertTrue(release_first_request.wait(timeout=2.0))
533+
elif req_meta["request_id"] == "other-peer-1":
534+
other_peer_started.set()
535+
536+
time.sleep(0.01)
537+
with state_lock:
538+
active_by_peer[peer_key] -= 1
539+
540+
self.thread._handle_request = handle_request # type: ignore[method-assign]
541+
same_peer_1 = {
542+
"request_id": "same-peer-1",
543+
"remote_host": "host-a",
544+
"remote_handshake_port": 6000,
545+
"all_task_done": False,
546+
}
547+
same_peer_2 = {
548+
"request_id": "same-peer-2",
549+
"remote_host": "host-a",
550+
"remote_handshake_port": 6000,
551+
"all_task_done": True,
552+
}
553+
other_peer = {
554+
"request_id": "other-peer-1",
555+
"remote_host": "host-b",
556+
"remote_handshake_port": 6001,
557+
"all_task_done": True,
558+
}
559+
560+
try:
561+
self.thread._submit_request(same_peer_1)
562+
self.assertTrue(first_request_started.wait(timeout=1.0))
563+
self.thread._submit_request(same_peer_2)
564+
self.thread._submit_request(other_peer)
565+
566+
self.assertTrue(other_peer_started.wait(timeout=1.0))
567+
time.sleep(0.05)
568+
self.assertNotIn("same-peer-2", handled_requests)
569+
finally:
570+
release_first_request.set()
571+
self.thread.executor.shutdown(wait=True, cancel_futures=True)
572+
573+
self.assertLess(handled_requests.index("same-peer-1"), handled_requests.index("same-peer-2"))
574+
self.assertEqual(max_active_by_peer[("host-a", 6000)], 1)
575+
self.assertEqual(max_active_by_peer[("host-b", 6001)], 1)
576+
577+
def test_peer_handler_yields_after_batch_limit(self):
578+
peer_key = ("host-a", 6000)
579+
requests = [
580+
{
581+
"request_id": f"req-{idx}",
582+
"remote_host": peer_key[0],
583+
"remote_handshake_port": peer_key[1],
584+
}
585+
for idx in range(MAX_REQUESTS_PER_PEER_HANDLER + 1)
586+
]
587+
handled_requests: list[str] = []
588+
self.thread.peer_request_queues[peer_key].extend(requests)
589+
self.thread.active_peer_request_handlers.add(peer_key)
590+
self.thread.executor = MagicMock()
591+
592+
def handle_request(req_meta: dict[str, Any]):
593+
handled_requests.append(req_meta["request_id"])
594+
595+
self.thread._handle_request = handle_request # type: ignore[method-assign]
596+
597+
self.thread._handle_peer_requests(peer_key)
598+
599+
self.assertEqual(handled_requests, [f"req-{idx}" for idx in range(MAX_REQUESTS_PER_PEER_HANDLER)])
600+
self.assertEqual(
601+
[req["request_id"] for req in self.thread.peer_request_queues[peer_key]],
602+
[f"req-{MAX_REQUESTS_PER_PEER_HANDLER}"],
603+
)
604+
self.assertIn(peer_key, self.thread.active_peer_request_handlers)
605+
self.thread.executor.submit.assert_called_once_with(self.thread._handle_peer_requests, peer_key)
606+
513607

514608
class TestSocketManagement(unittest.TestCase):
515609
def setUp(self):
@@ -534,7 +628,6 @@ def setUp(self):
534628
prefill_pp_layer_partition=None,
535629
)
536630
self.thread.remote_sockets = defaultdict(deque)
537-
self.thread.remote_poller = MagicMock()
538631

539632
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.zmq.Context")
540633
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.make_zmq_socket")
@@ -552,7 +645,8 @@ def test_get_remote_socket(self, mock_make_socket, mock_context):
552645
self.assertEqual(kwargs.get("path"), "tcp://test_host:12345")
553646
self.assertEqual(kwargs.get("socket_type"), zmq.REQ) # type: ignore
554647
self.assertFalse(kwargs.get("bind", True))
555-
self.thread.remote_poller.register.assert_called_with(mock_sock, zmq.POLLIN) # type: ignore
648+
mock_sock.setsockopt.assert_any_call(zmq.SNDTIMEO, int(self.thread.timeout * 1000)) # type: ignore
649+
mock_sock.setsockopt.assert_any_call(zmq.RCVTIMEO, int(self.thread.timeout * 1000)) # type: ignore
556650

557651
def test_return_socket_to_pool(self):
558652
mock_sock = MagicMock()
@@ -564,7 +658,6 @@ def test_return_socket_to_pool(self):
564658

565659
self.assertEqual(len(self.thread.remote_sockets[test_path]), 1)
566660
self.assertEqual(self.thread.remote_sockets[test_path][0], mock_sock)
567-
self.thread.remote_poller.register.assert_not_called()
568661

569662

570663
class TestCoreFunctionality(unittest.TestCase):
@@ -798,15 +891,15 @@ def test_get_remote_metadata_success(self, mock_recv, mock_send):
798891
patch.object(self.thread, "_get_remote_socket") as mock_get_socket,
799892
patch.object(self.thread, "_return_remote_socket") as mock_return_socket,
800893
):
801-
mock_socket = MagicMock()
894+
mock_socket = MagicMock(spec=zmq.Socket) # type: ignore[attr-defined]
802895
mock_get_socket.return_value = mock_socket
803896

804897
self.thread._get_remote_metadata("host1", 5555)
805898

806899
mock_get_socket.assert_called_once_with("host1", 5555)
807900
mock_return_socket.assert_called_once_with(mock_socket, "host1", 5555)
808901
mock_send.assert_called_once_with(mock_socket, self.thread.encoder.encode((GET_META_MSG, "")), "host1:5555")
809-
mock_recv.assert_called_once_with(mock_socket, self.thread.remote_poller, "host1:5555")
902+
mock_recv.assert_called_once_with(mock_socket, "host1:5555")
810903
self.assertEqual(self.thread.kv_caches_base_addr["remote_engine"][5555], [[0x3000], [0x4000]])
811904
self.assertEqual(self.thread.remote_block_stride_per_addr["remote_engine"][5555], [[1024]])
812905

@@ -820,14 +913,15 @@ def test_get_remote_metadata_failure(self, mock_recv, mock_send):
820913
patch.object(self.thread, "_get_remote_socket") as mock_get_socket,
821914
patch.object(self.thread, "_return_remote_socket") as mock_return_socket,
822915
):
823-
mock_socket = MagicMock()
916+
mock_socket = MagicMock(spec=zmq.Socket) # type: ignore[attr-defined]
824917
mock_get_socket.return_value = mock_socket
825918

826919
with self.assertRaises(Exception) as context:
827920
self.thread._get_remote_metadata("host1", 5555)
828921

829922
self.assertEqual(str(context.exception), "Network error")
830-
mock_return_socket.assert_called_once()
923+
mock_socket.close.assert_called_once()
924+
mock_return_socket.assert_not_called()
831925

832926

833927
class TestMainThreadLoop(unittest.TestCase):
@@ -1335,7 +1429,7 @@ def test_request_finished_no_remote_decode(self):
13351429

13361430
def test_get_transfer_block_ids_trims_attention_mtp_blocks(self):
13371431
self.scheduler.group_transfer_info = [
1338-
types.SimpleNamespace(
1432+
types.SimpleNamespace( # type: ignore[list-item]
13391433
tokens_per_block=16,
13401434
blocks_per_window=0,
13411435
is_state_group=False,
@@ -1348,7 +1442,7 @@ def test_get_transfer_block_ids_trims_attention_mtp_blocks(self):
13481442

13491443
def test_get_transfer_block_ids_keeps_state_group(self):
13501444
self.scheduler.group_transfer_info = [
1351-
types.SimpleNamespace(
1445+
types.SimpleNamespace( # type: ignore[list-item]
13521446
tokens_per_block=16,
13531447
blocks_per_window=0,
13541448
is_state_group=True,
@@ -1361,7 +1455,7 @@ def test_get_transfer_block_ids_keeps_state_group(self):
13611455

13621456
def test_get_transfer_block_ids_uses_compressed_prompt_len(self):
13631457
self.scheduler.group_transfer_info = [
1364-
types.SimpleNamespace(
1458+
types.SimpleNamespace( # type: ignore[list-item]
13651459
tokens_per_block=32,
13661460
blocks_per_window=0,
13671461
is_state_group=False,
@@ -1374,7 +1468,7 @@ def test_get_transfer_block_ids_uses_compressed_prompt_len(self):
13741468

13751469
def test_get_transfer_block_ids_trims_sliding_window_mtp_blocks(self):
13761470
self.scheduler.group_transfer_info = [
1377-
types.SimpleNamespace(
1471+
types.SimpleNamespace( # type: ignore[list-item]
13781472
tokens_per_block=16,
13791473
blocks_per_window=3,
13801474
is_state_group=False,
@@ -1387,7 +1481,7 @@ def test_get_transfer_block_ids_trims_sliding_window_mtp_blocks(self):
13871481

13881482
def test_get_swa_transfer_block_ids_clips_sliding_window_group(self):
13891483
self.scheduler.group_transfer_info = [
1390-
types.SimpleNamespace(
1484+
types.SimpleNamespace( # type: ignore[list-item]
13911485
tokens_per_block=16,
13921486
blocks_per_window=3,
13931487
is_state_group=False,
@@ -1400,7 +1494,7 @@ def test_get_swa_transfer_block_ids_clips_sliding_window_group(self):
14001494

14011495
def test_get_swa_transfer_block_ids_drops_zero_from_sliding_window_tail(self):
14021496
self.scheduler.group_transfer_info = [
1403-
types.SimpleNamespace(
1497+
types.SimpleNamespace( # type: ignore[list-item]
14041498
tokens_per_block=16,
14051499
blocks_per_window=2,
14061500
is_state_group=False,
@@ -1413,7 +1507,7 @@ def test_get_swa_transfer_block_ids_drops_zero_from_sliding_window_tail(self):
14131507

14141508
def test_transfer_block_ids_trims_mtp_before_swa_zero_filter(self):
14151509
self.scheduler.group_transfer_info = [
1416-
types.SimpleNamespace(
1510+
types.SimpleNamespace( # type: ignore[list-item]
14171511
tokens_per_block=16,
14181512
blocks_per_window=3,
14191513
is_state_group=False,
@@ -1578,20 +1672,15 @@ def test_ensure_zmq_send_retry_and_fail(self, mock_logger):
15781672
def test_ensure_zmq_recv_success(self, mock_logger):
15791673
mock_socket = MagicMock()
15801674
mock_socket.recv.return_value = b"response"
1581-
mock_poller = MagicMock()
1582-
mock_poller.poll.return_value = [
1583-
(mock_socket, zmq.POLLIN) # type: ignore
1584-
]
1585-
data = ensure_zmq_recv(mock_socket, mock_poller, "tcp://localhost:1234")
1675+
data = ensure_zmq_recv(mock_socket, "tcp://localhost:1234")
15861676
self.assertEqual(data, b"response")
15871677

15881678
@patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.logger")
15891679
def test_ensure_zmq_recv_timeout_and_fail(self, mock_logger):
15901680
mock_socket = MagicMock()
1591-
mock_poller = MagicMock()
1592-
mock_poller.poll.return_value = []
1681+
mock_socket.recv.side_effect = zmq.ZMQError("Receive timeout") # type: ignore
15931682
with self.assertRaises(RuntimeError):
1594-
ensure_zmq_recv(mock_socket, mock_poller, "tcp://localhost:1234", timeout=0.01, max_retries=2)
1683+
ensure_zmq_recv(mock_socket, "tcp://localhost:1234", max_retries=2)
15951684

15961685

15971686
class MockMooncakeAgentMetadata:

tests/ut/kv_offload/test_mooncake_hybrid_connector.py

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,11 @@
11
import sys
2+
import threading
3+
import time
24
import types
35
import unittest
6+
from collections import defaultdict, deque
7+
from concurrent.futures import ThreadPoolExecutor
8+
from typing import Any
49
from unittest.mock import MagicMock
510

611
fake_engine = types.ModuleType("mooncake.engine")
@@ -10,6 +15,8 @@
1015
from vllm.v1.request import RequestStatus # noqa: E402
1116

1217
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_hybrid_connector import ( # noqa: E402
18+
MAX_REQUESTS_PER_PEER_HANDLER,
19+
KVCacheRecvingThread,
1320
MooncakeConnectorScheduler,
1421
)
1522

@@ -33,6 +40,114 @@ def __init__(
3340
self.output_token_ids = [101]
3441

3542

43+
class TestHybridKVCacheRecvingThreadDispatch(unittest.TestCase):
44+
def _make_thread(self):
45+
thread = object.__new__(KVCacheRecvingThread)
46+
thread.executor = ThreadPoolExecutor(max_workers=2)
47+
thread.peer_request_queues = defaultdict(deque)
48+
thread.active_peer_request_handlers = set()
49+
thread.peer_request_queues_lock = threading.Lock()
50+
thread.request_task_counts = defaultdict(int)
51+
thread.finished_request_markers = set()
52+
thread.request_task_counts_lock = threading.Lock()
53+
return thread
54+
55+
def test_submit_request_serializes_same_peer_fifo(self):
56+
thread = self._make_thread()
57+
release_first_request = threading.Event()
58+
first_request_started = threading.Event()
59+
other_peer_started = threading.Event()
60+
handled_requests: list[str] = []
61+
active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
62+
max_active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
63+
state_lock = threading.Lock()
64+
65+
def handle_request(req_meta: dict[str, Any]):
66+
peer_key = (req_meta["remote_host"], req_meta["remote_handshake_port"])
67+
with state_lock:
68+
active_by_peer[peer_key] += 1
69+
max_active_by_peer[peer_key] = max(max_active_by_peer[peer_key], active_by_peer[peer_key])
70+
handled_requests.append(req_meta["request_id"])
71+
72+
if req_meta["request_id"] == "same-peer-1":
73+
first_request_started.set()
74+
self.assertTrue(release_first_request.wait(timeout=2.0))
75+
elif req_meta["request_id"] == "other-peer-1":
76+
other_peer_started.set()
77+
78+
time.sleep(0.01)
79+
with state_lock:
80+
active_by_peer[peer_key] -= 1
81+
82+
thread._handle_request = handle_request # type: ignore[method-assign]
83+
same_peer_1 = {
84+
"request_id": "same-peer-1",
85+
"remote_host": "host-a",
86+
"remote_handshake_port": 6000,
87+
"all_task_done": False,
88+
}
89+
same_peer_2 = {
90+
"request_id": "same-peer-2",
91+
"remote_host": "host-a",
92+
"remote_handshake_port": 6000,
93+
"all_task_done": True,
94+
}
95+
other_peer = {
96+
"request_id": "other-peer-1",
97+
"remote_host": "host-b",
98+
"remote_handshake_port": 6001,
99+
"all_task_done": True,
100+
}
101+
102+
try:
103+
thread._submit_request(same_peer_1)
104+
self.assertTrue(first_request_started.wait(timeout=1.0))
105+
thread._submit_request(same_peer_2)
106+
thread._submit_request(other_peer)
107+
108+
self.assertTrue(other_peer_started.wait(timeout=1.0))
109+
time.sleep(0.05)
110+
self.assertNotIn("same-peer-2", handled_requests)
111+
finally:
112+
release_first_request.set()
113+
thread.executor.shutdown(wait=True, cancel_futures=True)
114+
115+
self.assertLess(handled_requests.index("same-peer-1"), handled_requests.index("same-peer-2"))
116+
self.assertEqual(max_active_by_peer[("host-a", 6000)], 1)
117+
self.assertEqual(max_active_by_peer[("host-b", 6001)], 1)
118+
119+
def test_peer_handler_yields_after_batch_limit(self):
120+
thread = self._make_thread()
121+
peer_key = ("host-a", 6000)
122+
requests = [
123+
{
124+
"request_id": f"req-{idx}",
125+
"remote_host": peer_key[0],
126+
"remote_handshake_port": peer_key[1],
127+
}
128+
for idx in range(MAX_REQUESTS_PER_PEER_HANDLER + 1)
129+
]
130+
handled_requests: list[str] = []
131+
thread.peer_request_queues[peer_key].extend(requests)
132+
thread.active_peer_request_handlers.add(peer_key)
133+
thread.executor = MagicMock()
134+
135+
def handle_request(req_meta: dict[str, Any]):
136+
handled_requests.append(req_meta["request_id"])
137+
138+
thread._handle_request = handle_request # type: ignore[method-assign]
139+
140+
thread._handle_peer_requests(peer_key)
141+
142+
self.assertEqual(handled_requests, [f"req-{idx}" for idx in range(MAX_REQUESTS_PER_PEER_HANDLER)])
143+
self.assertEqual(
144+
[req["request_id"] for req in thread.peer_request_queues[peer_key]],
145+
[f"req-{MAX_REQUESTS_PER_PEER_HANDLER}"],
146+
)
147+
self.assertIn(peer_key, thread.active_peer_request_handlers)
148+
thread.executor.submit.assert_called_once_with(thread._handle_peer_requests, peer_key)
149+
150+
36151
class TestMooncakeHybridConnectorScheduler(unittest.TestCase):
37152
def _make_scheduler(self):
38153
scheduler = object.__new__(MooncakeConnectorScheduler)

0 commit comments

Comments
 (0)