Skip to content

Commit 44f4f06

Browse files
shipiyouniaocodex
andcommitted
[Bugfix] Release worker RPC payload before next dequeue
Execute each worker RPC in a separate stack frame so deserialized arguments and outputs are released before the next message is deserialized. Preserve exception handling behavior and add deterministic lifetime regression coverage. Fixes #43639 Co-authored-by: OpenAI Codex <codex@openai.com> Signed-off-by: 石皮幼鸟 <2960474346@qq.com>
1 parent aeece10 commit 44f4f06

2 files changed

Lines changed: 90 additions & 20 deletions

File tree

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,67 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3+
4+
import weakref
5+
from types import SimpleNamespace
6+
from typing import Any
7+
8+
import pytest
9+
10+
from vllm.v1.executor.multiproc_executor import WorkerProc
11+
12+
13+
class _ExitWorkerLoop(RuntimeError):
14+
pass
15+
16+
17+
class _RpcPayload:
18+
pass
19+
20+
21+
class _PayloadLifetimeCheckingQueue:
22+
def __init__(self) -> None:
23+
self.payload_ref: weakref.ReferenceType[_RpcPayload] | None = None
24+
self.dequeue_count = 0
25+
26+
def dequeue(self, *, indefinite: bool):
27+
assert indefinite
28+
self.dequeue_count += 1
29+
if self.dequeue_count == 1:
30+
payload = _RpcPayload()
31+
self.payload_ref = weakref.ref(payload)
32+
return "consume", (payload,), {}, None
33+
34+
assert self.payload_ref is not None
35+
assert self.payload_ref() is None
36+
raise _ExitWorkerLoop
37+
38+
39+
def test_worker_rpc_payload_released_before_next_dequeue():
40+
queue = _PayloadLifetimeCheckingQueue()
41+
worker_proc: Any = WorkerProc.__new__(WorkerProc)
42+
worker_proc.rpc_broadcast_mq = queue
43+
worker_proc.rank = 0
44+
worker_proc.worker = SimpleNamespace(consume=lambda payload: payload)
45+
worker_proc.handle_output = lambda output: None
46+
47+
with pytest.raises(_ExitWorkerLoop):
48+
worker_proc.worker_busy_loop()
49+
50+
assert queue.dequeue_count == 2
51+
52+
53+
def test_execute_worker_rpc_returns_worker_exception():
54+
def fail():
55+
raise RuntimeError("test error")
56+
57+
worker_proc: Any = WorkerProc.__new__(WorkerProc)
58+
worker_proc.rank = 0
59+
worker_proc.worker = SimpleNamespace(fail=fail)
60+
outputs: list[Any] = []
61+
worker_proc.handle_output = outputs.append
62+
63+
worker_proc._execute_worker_rpc(("fail", (), {}, None))
64+
65+
assert len(outputs) == 1
66+
assert isinstance(outputs[0], RuntimeError)
67+
assert str(outputs[0]) == "test error"

vllm/v1/executor/multiproc_executor.py

Lines changed: 23 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -1013,28 +1013,31 @@ def worker_busy_loop(self):
10131013
"""Main busy loop for Multiprocessing Workers"""
10141014
assert self.rpc_broadcast_mq is not None
10151015
while True:
1016-
method, args, kwargs, output_rank = self.rpc_broadcast_mq.dequeue(
1017-
indefinite=True
1018-
)
1019-
try:
1020-
if isinstance(method, str):
1021-
func = getattr(self.worker, method)
1022-
elif isinstance(method, bytes):
1023-
func = partial(cloudpickle.loads(method), self.worker)
1016+
self._execute_worker_rpc(self.rpc_broadcast_mq.dequeue(indefinite=True))
1017+
1018+
def _execute_worker_rpc(
1019+
self,
1020+
rpc_request: tuple[str | bytes, tuple[Any, ...], dict[str, Any], int | None],
1021+
) -> None:
1022+
"""Execute one RPC in a separate frame from the dequeue loop."""
1023+
method, args, kwargs, output_rank = rpc_request
1024+
try:
1025+
if isinstance(method, str):
1026+
func = getattr(self.worker, method)
1027+
elif isinstance(method, bytes):
1028+
func = partial(cloudpickle.loads(method), self.worker)
10241029

1025-
output = func(*args, **kwargs)
1030+
output = func(*args, **kwargs)
10261031

1027-
if output_rank is None or self.rank == output_rank:
1028-
self.handle_output(output)
1029-
except Exception as e:
1030-
# Notes have been introduced in python 3.11
1031-
if hasattr(e, "add_note"):
1032-
e.add_note(traceback.format_exc())
1033-
logger.exception("WorkerProc hit an exception.")
1034-
# exception might not be serializable, so we convert it to
1035-
# string, only for logging purpose.
1036-
if output_rank is None or self.rank == output_rank:
1037-
self.handle_output(e)
1032+
if output_rank is None or self.rank == output_rank:
1033+
self.handle_output(output)
1034+
except Exception as e:
1035+
# Notes have been introduced in python 3.11
1036+
if hasattr(e, "add_note"):
1037+
e.add_note(traceback.format_exc())
1038+
logger.exception("WorkerProc hit an exception.")
1039+
if output_rank is None or self.rank == output_rank:
1040+
self.handle_output(e)
10381041

10391042
@staticmethod
10401043
def setup_proc_title_and_log_prefix(enable_ep: bool) -> None:

0 commit comments

Comments
 (0)