Skip to content

Commit cd31915

Browse files
leoda1ceci3
authored andcommitted
fix: match P/D requests by transfer_id in FlagCX connector (#315)
### PR Category Core ### PR Type Bug Fixes ### Description 1. In disaggregated (P/D) serving, the Prefill and Decode engines assign **different** `request_id`s to the same logical request. The FlagCX connector previously used `request_id` as the key for its send/receive bookkeeping, so the Decode side's `req_id` could never be found in the Prefill side's `reqs_need_send` map, causing KV transfers to silently fail to match (`Request %s not found in reqs_need_send`). 2. This PR introduces a stable, cross-engine `transfer_id` that both sides agree on, rekeys all connector bookkeeping by it, and hardens the transfer path with proper timeouts. 3. This PR will use the FlagCX v0.13.0 interface, while also being compatible with FlagCX v0.9.0 within CICD. ### Related Issues <!-- Link any related issues: Fixes #issue, Closes #issue, or Related to #issue --> ### Changes - **`flagcx_connector.py`**: added a `TransferId` type and threaded `transfer_id` through `RecvReqMeta`, `SendBlockMeta`, `SendReqMeta`, and the connector metadata. `reqs_to_send` / `req_blocks` now carry `(transfer_id, block_ids)`; the worker looks up send metadata by `transfer_id` while still reporting completion/timeout under the Prefill `request_id` (`p_req_id`) that vLLM expects. - **`flagcx_connector.py`**: guarded against a missing `transfer_id` — the scheduler logs a warning and skips the request instead of mis-matching. - **`flagcx_connector.py`**: replaced unbounded / hardcoded 60s waits with bounded waits driven by `_abort_request_timeout` — `send_meta.ready.wait()` uses a deadline and raises a clear `RuntimeError` on timeout, and the receive socket `RCVTIMEO` now scales with `_abort_request_timeout`. - **`examples/disaggregated_serving_xpyd/router.py`**: inject `transfer_id = f"fgx-{request_id}"` into `kv_transfer_params` for both the prefill and decode legs so both engines share the same transfer key. - **`vllm_fl/distributed/device_communicators/flagcx.py`**: fixed ctypes usage for `flagcxGetUniqueId()` / `flagcxCommInitRank()` — pass the unique id object directly instead of `.contents` / `ctypes.byref(...)`. - **`vllm_fl/ops/fused_moe/fused_moe_utils.py`**: provide a default MoE backend priority (`TRITON`, `BATCHED_TRITON`) for out-of-tree platforms. ### Testing <!-- How has this change been tested? Include test commands, hardware used, etc. --> - ### Checklist - [x] I have run the existing tests and they pass - [x] I have added tests for my changes (if applicable) - [x] I have updated the documentation (if applicable) --------- Co-authored-by: ceci3 <ceci3@users.noreply.github.com>
1 parent e8a1fb2 commit cd31915

5 files changed

Lines changed: 115 additions & 44 deletions

File tree

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -99,7 +99,7 @@ In theory, vllm-plugin-FL can support all models available in vLLM, as long as n
9999
```sh
100100
git clone https://github.com/flagos-ai/FlagCX.git
101101
cd FlagCX
102-
git checkout -b v0.9.0
102+
git checkout -b v0.13.0
103103
git submodule update --init --recursive
104104
```
105105

examples/disaggregated_serving_xpyd/router.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -124,7 +124,11 @@ async def lifespan(app):
124124

125125
async def send_to_prefill(prefill_client, endpoint, req_data, request_id):
126126
data = req_data.copy()
127-
data["kv_transfer_params"] = {"do_remote_decode": True, "do_remote_prefill": False}
127+
data["kv_transfer_params"] = {
128+
"do_remote_decode": True,
129+
"do_remote_prefill": False,
130+
"transfer_id": f"fgx-{request_id}",
131+
}
128132
data["stream"] = False
129133
data["max_tokens"] = 1
130134
if "max_completion_tokens" in data:
@@ -153,6 +157,7 @@ async def stream_from_decode(
153157
"do_remote_decode": False,
154158
"remote_host": prefill_client["remote_host"],
155159
"remote_port": prefill_client["side_channel_port"],
160+
"transfer_id": f"fgx-{request_id}",
156161
}
157162
headers = {"X-Request-Id": request_id}
158163
api_key = os.environ.get("OPENAI_API_KEY", "")

vllm_fl/distributed/device_communicators/flagcx.py

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -96,9 +96,17 @@ def __init__(
9696
self.available = True
9797
self.disabled = False
9898

99+
self._legacy_unique_id_api = (
100+
hasattr(self.flagcx, "handler")
101+
and not hasattr(self.flagcx, "devHandle")
102+
)
103+
99104
if self.rank == 0:
100105
# get the unique id from NCCL
101-
self.unique_id = self.flagcx.flagcxGetUniqueId().contents
106+
if self._legacy_unique_id_api:
107+
self.unique_id = self.flagcx.flagcxGetUniqueId().contents
108+
else:
109+
self.unique_id = self.flagcx.flagcxGetUniqueId()
102110
else:
103111
# construct an empty unique id
104112
self.unique_id = flagcxUniqueId()
@@ -133,8 +141,12 @@ def __init__(
133141
device_ctx = torch.cuda.device(self.device)
134142

135143
with device_ctx:
136-
self.comm = self.flagcx.flagcxCommInitRank(
137-
self.world_size, ctypes.byref(self.unique_id), self.rank)
144+
if self._legacy_unique_id_api:
145+
self.comm = self.flagcx.flagcxCommInitRank(
146+
self.world_size, ctypes.pointer(self.unique_id), self.rank)
147+
else:
148+
self.comm = self.flagcx.flagcxCommInitRank(
149+
self.world_size, self.unique_id, self.rank)
138150

139151
stream = current_stream()
140152
# A small all_reduce for warmup.

vllm_fl/distributed/kv_transfer/flagcx_connector.py

Lines changed: 88 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,7 @@
7474

7575
EngineId = str
7676
ReqId = str
77+
TransferId = str
7778

7879
TRANS_DONE = b"trans_done"
7980
TRANS_ERROR = b"trans_error"
@@ -109,8 +110,7 @@ class FlagCXAgentMetadata(
109110
remote_port: int
110111
remote_tp_size: int
111112
remote_tp_rank: int
112-
# req_id -> per KV-cache-group block ids on the Decode (receiver) side.
113-
req_blocks: dict[ReqId, list[list[int]]]
113+
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]]
114114
kv_caches_base_addr: list[int]
115115
block_lens: list[int]
116116
kv_block_lens: list[int]
@@ -124,11 +124,14 @@ class RecvReqMeta:
124124
local_block_ids: list[list[int]]
125125
remote_host: str
126126
remote_port: int
127+
transfer_id: TransferId
127128
remote_tp_size: int = 0
128129

129130

130131
@dataclass
131132
class SendBlockMeta:
133+
p_req_id: ReqId
134+
transfer_id: TransferId
132135
local_block_ids: list[list[int]]
133136
ready: threading.Event
134137
expire_time: float = float("inf")
@@ -138,7 +141,7 @@ class SendBlockMeta:
138141

139142
@dataclass
140143
class SendReqMeta:
141-
reqs: dict[ReqId, SendBlockMeta]
144+
reqs: dict[TransferId, SendBlockMeta]
142145
lock: threading.Lock
143146

144147

@@ -158,7 +161,9 @@ class FlagCXConnectorMetadata(KVConnectorMetadata):
158161
def __init__(self):
159162
self.reqs_to_recv: dict[ReqId, RecvReqMeta] = {}
160163
# Per-request, per KV-cache-group block ids.
161-
self.reqs_to_send: dict[ReqId, list[list[int]]] = {}
164+
self.reqs_to_send: dict[
165+
ReqId, tuple[TransferId, list[list[int]]]
166+
] = {}
162167

163168
def add_new_req(
164169
self,
@@ -167,15 +172,17 @@ def add_new_req(
167172
kv_transfer_params: dict[str, Any],
168173
load_remote_cache: bool = True,
169174
):
175+
transfer_id = kv_transfer_params["transfer_id"]
170176
if load_remote_cache:
171177
self.reqs_to_recv[request_id] = RecvReqMeta(
172178
local_block_ids=local_block_ids,
173179
remote_host=kv_transfer_params["remote_host"],
174180
remote_port=kv_transfer_params["remote_port"],
181+
transfer_id=transfer_id,
175182
remote_tp_size=kv_transfer_params.get("remote_tp_size", 0),
176183
)
177184
else:
178-
self.reqs_to_send[request_id] = local_block_ids
185+
self.reqs_to_send[request_id] = (transfer_id, local_block_ids)
179186

180187

181188
class FlagCXConnector(KVConnectorBase_V1, SupportsHMA):
@@ -334,7 +341,9 @@ def __init__(
334341
]
335342

336343
self._reqs_need_recv: dict[ReqId, tuple[Request, list[list[int]]]] = {}
337-
self._reqs_need_send: dict[ReqId, list[list[int]]] = {}
344+
self._reqs_need_send: dict[
345+
ReqId, tuple["Request", list[list[int]]]
346+
] = {}
338347

339348
def get_sw_clipped_blocks(
340349
self,
@@ -415,7 +424,9 @@ def update_state_after_alloc(
415424

416425
if params.get("do_remote_prefill"):
417426
assert self.kv_role != "kv_producer"
418-
if all(p in params for p in ("remote_host", "remote_port")):
427+
if all(
428+
p in params for p in ("remote_host", "remote_port", "transfer_id")
429+
):
419430
unhashed_block_ids = (
420431
blocks.get_unhashed_block_ids_all_groups()
421432
if num_external_tokens > 0
@@ -431,7 +442,10 @@ def update_state_after_alloc(
431442
params["do_remote_prefill"] = False
432443

433444
elif params.get("do_remote_decode"):
434-
self._reqs_need_send[request.request_id] = []
445+
if not params.get("transfer_id"):
446+
logger.warning("Missing transfer_id in KVTransferParams: %s", params)
447+
else:
448+
self._reqs_need_send[request.request_id] = (request, [])
435449

436450
def build_connector_meta(
437451
self, scheduler_output: SchedulerOutput
@@ -449,11 +463,12 @@ def build_connector_meta(
449463
self._reqs_need_recv.clear()
450464

451465
if self.kv_role != "kv_consumer":
452-
for req_id, block_ids in self._reqs_need_send.items():
466+
for req_id, (req, block_ids) in self._reqs_need_send.items():
467+
assert req.kv_transfer_params is not None
453468
meta.add_new_req(
454469
request_id=req_id,
455470
local_block_ids=block_ids,
456-
kv_transfer_params={},
471+
kv_transfer_params=req.kv_transfer_params,
457472
load_remote_cache=False,
458473
)
459474
self._reqs_need_send.clear()
@@ -464,7 +479,7 @@ def request_finished(
464479
self, request: "Request", block_ids: tuple[list[int], ...]
465480
) -> tuple[bool, dict[str, Any] | None]:
466481
params = request.kv_transfer_params
467-
if not params:
482+
if not params or not params.get("transfer_id"):
468483
return False, None
469484

470485
if params.get("do_remote_prefill"):
@@ -483,8 +498,9 @@ def request_finished(
483498
delay_free_blocks = any(len(group) > 0 for group in block_ids)
484499

485500
if delay_free_blocks:
486-
self._reqs_need_send[request.request_id] = self.get_sw_clipped_blocks(
487-
block_ids
501+
self._reqs_need_send[request.request_id] = (
502+
request,
503+
self.get_sw_clipped_blocks(block_ids),
488504
)
489505

490506
return delay_free_blocks, dict(
@@ -493,6 +509,7 @@ def request_finished(
493509
remote_host=self.side_channel_host,
494510
remote_port=self.side_channel_port,
495511
remote_tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
512+
transfer_id=params["transfer_id"],
496513
)
497514

498515

@@ -920,19 +937,30 @@ def _send_kv_to_decode(self, meta: FlagCXAgentMetadata) -> None:
920937

921938
ready_reqs: list[tuple[ReqId, SendBlockMeta]] = []
922939
with self.reqs_need_send.lock:
923-
for req_id in meta.req_blocks:
924-
send_meta = self.reqs_need_send.reqs.get(req_id)
940+
for d_req_id, (transfer_id, _) in meta.req_blocks.items():
941+
send_meta = self.reqs_need_send.reqs.get(transfer_id)
925942
if send_meta is None:
926-
logger.warning("Request %s not found in reqs_need_send", req_id)
927-
return
943+
send_meta = SendBlockMeta(
944+
p_req_id="",
945+
transfer_id=transfer_id,
946+
local_block_ids=[],
947+
ready=threading.Event(),
948+
)
949+
self.reqs_need_send.reqs[transfer_id] = send_meta
928950
# Mark it as not expired. We will send it now.
929951
send_meta.expire_time = float("inf")
930952
send_meta.need_send = need
931-
ready_reqs.append((req_id, send_meta))
953+
ready_reqs.append((d_req_id, send_meta))
932954

933955
# Wait until the scheduler has committed each request's blocks.
934-
for _, send_meta in ready_reqs:
935-
send_meta.ready.wait()
956+
deadline = time.perf_counter() + self._abort_request_timeout
957+
for d_req_id, send_meta in ready_reqs:
958+
remaining = deadline - time.perf_counter()
959+
if remaining <= 0 or not send_meta.ready.wait(timeout=remaining):
960+
raise RuntimeError(
961+
"Timed out waiting for prefill blocks of transfer "
962+
f"{send_meta.transfer_id} (decode request {d_req_id})."
963+
)
936964

937965
remote_session = f"{meta.remote_hostname}:{meta.remote_port}"
938966
conn = self._get_conn(remote_session)
@@ -980,11 +1008,15 @@ def _send_kv_to_decode(self, meta: FlagCXAgentMetadata) -> None:
9801008

9811009
finished: list[ReqId] = []
9821010
with self.reqs_need_send.lock:
983-
for req_id, send_meta in ready_reqs:
1011+
for _, send_meta in ready_reqs:
9841012
send_meta.sent += 1
9851013
if send_meta.sent >= max(send_meta.need_send, 1):
986-
self.reqs_need_send.reqs.pop(req_id, None)
987-
finished.append(req_id)
1014+
self.reqs_need_send.reqs.pop(send_meta.transfer_id, None)
1015+
if not send_meta.p_req_id:
1016+
raise RuntimeError(
1017+
f"Missing Prefill request ID for {send_meta.transfer_id}."
1018+
)
1019+
finished.append(send_meta.p_req_id)
9881020
if finished:
9891021
with self.finished_sending_reqs.lock:
9901022
self.finished_sending_reqs.set.update(finished)
@@ -1011,7 +1043,7 @@ def _build_transfer_params(
10111043
group_specs = self.kv_cache_config.kv_cache_groups
10121044

10131045
for d_req_id, send_meta in ready_reqs:
1014-
remote_block_ids_per_group = agent_meta.req_blocks[d_req_id]
1046+
_, remote_block_ids_per_group = agent_meta.req_blocks[d_req_id]
10151047
if not remote_block_ids_per_group or all(
10161048
len(g) == 0 for g in remote_block_ids_per_group
10171049
):
@@ -1129,7 +1161,7 @@ def _receiver_loop_fn(self, loop: asyncio.AbstractEventLoop):
11291161
async def _receive_kv(
11301162
self,
11311163
path: str,
1132-
req_blocks: dict[ReqId, list[list[int]]],
1164+
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]],
11331165
):
11341166
req_ids = list(req_blocks.keys())
11351167

@@ -1152,7 +1184,9 @@ async def _receive_kv(
11521184
sock: zmq.asyncio.Socket = make_zmq_socket(
11531185
self.async_zmq_ctx, path, zmq.REQ, bind=False, linger=0
11541186
)
1155-
sock.setsockopt(zmq.RCVTIMEO, 60000)
1187+
sock.setsockopt(
1188+
zmq.RCVTIMEO, (self._abort_request_timeout + 60) * 1000
1189+
)
11561190

11571191
try:
11581192
await sock.send(encoded_data)
@@ -1193,13 +1227,21 @@ def start_load_kv(self, metadata: FlagCXConnectorMetadata):
11931227

11941228
if self.kv_role != "kv_consumer":
11951229
with self.reqs_need_send.lock:
1196-
for req_id, block_ids in metadata.reqs_to_send.items():
1197-
send_meta = self.reqs_need_send.reqs.get(req_id)
1230+
for p_req_id, (
1231+
transfer_id,
1232+
block_ids,
1233+
) in metadata.reqs_to_send.items():
1234+
send_meta = self.reqs_need_send.reqs.get(transfer_id)
11981235
if send_meta is None:
11991236
send_meta = SendBlockMeta(
1200-
local_block_ids=[], ready=threading.Event()
1237+
p_req_id=p_req_id,
1238+
transfer_id=transfer_id,
1239+
local_block_ids=[],
1240+
ready=threading.Event(),
12011241
)
1202-
self.reqs_need_send.reqs[req_id] = send_meta
1242+
self.reqs_need_send.reqs[transfer_id] = send_meta
1243+
else:
1244+
send_meta.p_req_id = p_req_id
12031245
# Non-empty means request_finished() has committed the
12041246
# per-group block ids; arm the send.
12051247
if block_ids:
@@ -1223,7 +1265,9 @@ async def _group_kv_pull(
12231265
loop; populates ``_pull_pending`` and launches the per-peer pulls
12241266
without an intervening await, so it is atomic w.r.t. other coroutines.
12251267
"""
1226-
kv_pulls: dict[str, dict[ReqId, list[list[int]]]] = defaultdict(dict)
1268+
kv_pulls: dict[
1269+
str, dict[ReqId, tuple[TransferId, list[list[int]]]]
1270+
] = defaultdict(dict)
12271271
for req_id, meta in reqs_to_recv.items():
12281272
remote_tp_size = meta.remote_tp_size or self.tp_size
12291273
target_p_ranks = self.kv_topo.handshake_target_ranks(remote_tp_size)
@@ -1232,7 +1276,10 @@ async def _group_kv_pull(
12321276
path = make_zmq_path(
12331277
"tcp", meta.remote_host, meta.remote_port + p_rank
12341278
)
1235-
kv_pulls[path][req_id] = meta.local_block_ids
1279+
kv_pulls[path][req_id] = (
1280+
meta.transfer_id,
1281+
meta.local_block_ids,
1282+
)
12361283
for path, req_blocks in kv_pulls.items():
12371284
asyncio.ensure_future(self._receive_kv(path, req_blocks))
12381285

@@ -1261,15 +1308,17 @@ def get_finished(self) -> tuple[set[str] | None, set[str] | None]:
12611308
now = time.perf_counter()
12621309
with self.reqs_need_send.lock:
12631310
expired = [
1264-
rid
1265-
for rid, sm in self.reqs_need_send.reqs.items()
1311+
transfer_id
1312+
for transfer_id, sm in self.reqs_need_send.reqs.items()
12661313
if sm.expire_time < now
12671314
]
1268-
for rid in expired:
1269-
logger.warning("Request %s send timed out, freeing blocks", rid)
1270-
del self.reqs_need_send.reqs[rid]
1271-
if expired:
1272-
finished_sending.update(expired)
1315+
for transfer_id in expired:
1316+
send_meta = self.reqs_need_send.reqs.pop(transfer_id)
1317+
logger.warning(
1318+
"Transfer %s send timed out, freeing blocks", transfer_id
1319+
)
1320+
if send_meta.p_req_id:
1321+
finished_sending.add(send_meta.p_req_id)
12731322

12741323
return finished_sending or None, finished_recving or None
12751324

vllm_fl/ops/fused_moe/fused_moe_utils.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,11 @@ def _move_to_back(
8585
_AVAILABLE_BACKENDS = [UnquantizedMoeBackend.XPU]
8686
elif current_platform.is_cpu():
8787
_AVAILABLE_BACKENDS = [UnquantizedMoeBackend.CPU]
88+
elif current_platform.is_out_of_tree():
89+
_AVAILABLE_BACKENDS = [
90+
UnquantizedMoeBackend.TRITON,
91+
UnquantizedMoeBackend.BATCHED_TRITON,
92+
]
8893
return _AVAILABLE_BACKENDS
8994

9095
## Adopt from select_unquantized_moe_backend

0 commit comments

Comments
 (0)