Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 67 additions & 0 deletions tests/v1/executor/test_multiproc_executor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import weakref
from types import SimpleNamespace
from typing import Any

import pytest

from vllm.v1.executor.multiproc_executor import WorkerProc


class _ExitWorkerLoop(RuntimeError):
pass


class _RpcPayload:
pass


class _PayloadLifetimeCheckingQueue:
def __init__(self) -> None:
self.payload_ref: weakref.ReferenceType[_RpcPayload] | None = None
self.dequeue_count = 0

def dequeue(self, *, indefinite: bool):
assert indefinite
self.dequeue_count += 1
if self.dequeue_count == 1:
payload = _RpcPayload()
self.payload_ref = weakref.ref(payload)
return "consume", (payload,), {}, None

assert self.payload_ref is not None
assert self.payload_ref() is None
raise _ExitWorkerLoop


def test_worker_rpc_payload_released_before_next_dequeue():
queue = _PayloadLifetimeCheckingQueue()
worker_proc: Any = WorkerProc.__new__(WorkerProc)
worker_proc.rpc_broadcast_mq = queue
worker_proc.rank = 0
worker_proc.worker = SimpleNamespace(consume=lambda payload: payload)
worker_proc.handle_output = lambda output: None

with pytest.raises(_ExitWorkerLoop):
worker_proc.worker_busy_loop()

assert queue.dequeue_count == 2


def test_execute_worker_rpc_returns_worker_exception():
def fail():
raise RuntimeError("test error")

worker_proc: Any = WorkerProc.__new__(WorkerProc)
worker_proc.rank = 0
worker_proc.worker = SimpleNamespace(fail=fail)
outputs: list[Any] = []
worker_proc.handle_output = outputs.append

worker_proc._execute_worker_rpc(("fail", (), {}, None))

assert len(outputs) == 1
assert isinstance(outputs[0], RuntimeError)
assert str(outputs[0]) == "test error"
45 changes: 25 additions & 20 deletions vllm/v1/executor/multiproc_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -1013,28 +1013,33 @@ def worker_busy_loop(self):
"""Main busy loop for Multiprocessing Workers"""
assert self.rpc_broadcast_mq is not None
while True:
method, args, kwargs, output_rank = self.rpc_broadcast_mq.dequeue(
indefinite=True
)
try:
if isinstance(method, str):
func = getattr(self.worker, method)
elif isinstance(method, bytes):
func = partial(cloudpickle.loads(method), self.worker)
self._execute_worker_rpc(self.rpc_broadcast_mq.dequeue(indefinite=True))

def _execute_worker_rpc(
self,
rpc_request: tuple[str | bytes, tuple[Any, ...], dict[str, Any], int | None],
) -> None:
"""Execute one RPC in a separate frame from the dequeue loop."""
method, args, kwargs, output_rank = rpc_request
try:
if isinstance(method, str):
func = getattr(self.worker, method)
elif isinstance(method, bytes):
func = partial(cloudpickle.loads(method), self.worker)

output = func(*args, **kwargs)
output = func(*args, **kwargs)

if output_rank is None or self.rank == output_rank:
self.handle_output(output)
except Exception as e:
# Notes have been introduced in python 3.11
if hasattr(e, "add_note"):
e.add_note(traceback.format_exc())
logger.exception("WorkerProc hit an exception.")
# exception might not be serializable, so we convert it to
# string, only for logging purpose.
Comment thread
shipiyouniao marked this conversation as resolved.
if output_rank is None or self.rank == output_rank:
self.handle_output(e)
if output_rank is None or self.rank == output_rank:
self.handle_output(output)
except Exception as e:
# Notes have been introduced in python 3.11
if hasattr(e, "add_note"):
e.add_note(traceback.format_exc())
logger.exception("WorkerProc hit an exception.")
# enqueue_output converts the exception to a FAILURE response
# containing its string representation before transport.
if output_rank is None or self.rank == output_rank:
self.handle_output(e)

@staticmethod
def setup_proc_title_and_log_prefix(enable_ep: bool) -> None:
Expand Down
Loading