Skip to content
Merged
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
133 changes: 131 additions & 2 deletions tests/diffusion/test_multiproc_engine_concurrency.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,10 @@
import asyncio
import multiprocessing as mp
import queue
import signal
import threading
import time
import weakref
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock

Expand Down Expand Up @@ -1330,6 +1332,7 @@ def test_worker_joins_share_one_global_deadline(self, monkeypatch):
class FakeProcess:
def __init__(self, name):
self.name = name
self.pid = 123
self.alive = True
self.terminated = False
self.join_timeouts = []
Expand Down Expand Up @@ -1358,6 +1361,132 @@ def terminate(self):
assert first.terminated and second.terminated
assert not first.is_alive() and not second.is_alive()

@pytest.mark.parametrize("cooperative", [False, True])
def test_real_workers_exit_and_are_reaped(self, monkeypatch, cooperative):
from vllm_omni.diffusion.executor import multiproc_executor as executor_module

ctx = mp.get_context("fork")
ready, stop = ctx.Event(), ctx.Event()

def run_worker():
signal.signal(signal.SIGTERM, signal.SIG_IGN)
ready.set()
stop.wait()

dead = ctx.Process(target=lambda: None)
worker = ctx.Process(target=run_worker)
monkeypatch.setattr(executor_module, "_WORKER_SHUTDOWN_GRACE_S", 0.2)
monkeypatch.setattr(executor_module, "_WORKER_TERMINATE_GRACE_S", 0.1)
monkeypatch.setattr(executor_module, "_WORKER_KILL_GRACE_S", 5.0, raising=False)
dead.start()
worker.start()
try:
dead.join(5)
assert not dead.is_alive()
assert ready.wait(5), "worker did not install its signal handler"
mq = Mock()
if cooperative:
mq.enqueue.side_effect = lambda *args, **kwargs: stop.set()
else:
mq.enqueue.side_effect = OSError("shutdown queue unavailable")
cleaner = executor_module._ExecutorShutdownCleaner(mq, 2, [dead, worker])

cleaner()

assert not worker.is_alive()
assert worker.exitcode == (0 if cooperative else -signal.SIGKILL)
assert cleaner.processes == []
assert cleaner.broadcast_mq is None
cleaner()
finally:
for proc in (dead, worker):
if proc.is_alive():
proc.kill()
proc.join(5)
proc.close()

@pytest.mark.parametrize("failed_action", ["terminate", "kill"])
def test_kill_phase_shares_deadline_and_continues_after_os_errors(self, monkeypatch, failed_action):
from vllm_omni.diffusion.executor import multiproc_executor as executor_module

first, second = Mock(pid=101), Mock(pid=102)
first.name, second.name = "first", "second"
for proc in (first, second):
proc.is_alive.side_effect = [True, True, True, False]
getattr(first, failed_action).side_effect = OSError("signal failed")
first.join.side_effect = [OSError("join failed"), None, None]
monotonic = Mock(side_effect=[100, 100, 110, 120, 120, 124, 130, 130, 134])
monkeypatch.setattr(executor_module, "time", SimpleNamespace(monotonic=monotonic))
cleaner = executor_module._ExecutorShutdownCleaner(processes=[first, second])

cleaner()

first.kill.assert_called_once()
second.kill.assert_called_once()
assert [call.args[0] for call in first.join.call_args_list] == [15, 5, 5]
assert [call.args[0] for call in second.join.call_args_list] == [5, 1, 1]
assert cleaner.processes == []

def test_executor_retains_survivor_for_explicit_shutdown_retry(self, monkeypatch):
from vllm_omni.diffusion.executor import multiproc_executor as executor_module

survivor = Mock(pid=123)
survivor.name = "surviving-worker"
survivor.is_alive.return_value = True
cleaner = executor_module._ExecutorShutdownCleaner(processes=[survivor])
executor, _, _ = _make_executor()
executor._shutdown_cleaner = cleaner
executor._processes = [survivor]
executor._finalizer = weakref.finalize(executor, cleaner)
executor._pump_stop = threading.Event()
executor._futures_lock = threading.RLock()
executor._rpc_futures = {}
executor._output_futures = {}
executor._batch_split_map = {}
log_error = Mock()
monkeypatch.setattr(executor_module.logger, "error", log_error)

executor.shutdown()

assert not executor._finalizer.alive
assert executor._shutdown_cleaner is cleaner
assert executor._processes == [survivor]
log_error.assert_called_once()
assert log_error.call_args.args[1] == [("surviving-worker", 123)]
survivor.is_alive.return_value = False

executor.shutdown()
executor.shutdown()

assert executor._shutdown_cleaner is None
assert executor._processes == []

def test_concurrent_cleanup_does_not_repeat_process_operations(self):
from vllm_omni.diffusion.executor import multiproc_executor as executor_module

joining, release = threading.Event(), threading.Event()
proc = Mock(pid=123)
proc.is_alive.return_value = True

def join(timeout):
joining.set()
assert release.wait(5)
proc.is_alive.return_value = False

proc.join.side_effect = join
cleaner = executor_module._ExecutorShutdownCleaner(processes=[proc])
thread = threading.Thread(target=cleaner)
thread.start()
try:
assert joining.wait(5)
cleaner()
proc.join.assert_called_once()
finally:
release.set()
thread.join(5)
assert not thread.is_alive()
assert cleaner.processes == []


# ───────── monitor thread & death sentinel integration tests ─────────

Expand Down Expand Up @@ -1406,7 +1535,7 @@ def test_worker_monitor_sets_is_failed_and_calls_callbacks_on_death(self):
executor._result_mq = None
executor._shutdown_cleaner = None
# Use a no-op so shutdown() doesn't crash on None resources.
executor._finalizer = lambda: None
executor._finalizer = weakref.finalize(executor, lambda: None)
# ------------------------------------------------------------------
# Attributes added by remove_bubble_v2 (async D2H); shutdown() iterates
# over them, so they need to exist even when constructed via __new__.
Expand Down Expand Up @@ -1441,7 +1570,7 @@ def test_worker_monitor_noop_when_already_closed(self):
executor._broadcast_mq = None
executor._result_mq = None
executor._shutdown_cleaner = None
executor._finalizer = lambda: None
executor._finalizer = weakref.finalize(executor, lambda: None)

proc = _make_short_lived_process()
executor._processes = [proc]
Expand Down
71 changes: 54 additions & 17 deletions vllm_omni/diffusion/executor/multiproc_executor.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

from __future__ import annotations

import concurrent.futures
Expand All @@ -10,7 +13,7 @@
import time
import weakref
from collections.abc import Callable
from dataclasses import dataclass
from dataclasses import dataclass, field
from multiprocessing.synchronize import Event
from typing import TYPE_CHECKING, Any, cast

Expand Down Expand Up @@ -42,6 +45,7 @@
_DLO_DP_WAVE_TIMEOUT_S = float(os.environ.get("VLLM_OMNI_DLO_DP_WAVE_TIMEOUT", 600.0))
_WORKER_SHUTDOWN_GRACE_S = 15.0
_WORKER_TERMINATE_GRACE_S = 5.0
_WORKER_KILL_GRACE_S = 5.0
_RESULT_PUMP_JOIN_TIMEOUT_S = 2.0


Expand Down Expand Up @@ -80,33 +84,61 @@ class _ExecutorShutdownCleaner:
broadcast_mq: MessageQueue | None = None
num_workers: int = 0
processes: list[mp.Process] | None = None
_lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False)

def __call__(self) -> None:
"""Clean up background resources."""
# The worker monitor and an explicit shutdown may race. A retry must
# not signal/join the same Process objects while cleanup is in flight.
if not self._lock.acquire(blocking=False):
return
try:
self._cleanup()
finally:
self._lock.release()

def _cleanup(self) -> None:
if self.broadcast_mq is not None:
try:
for _ in range(self.num_workers):
self.broadcast_mq.enqueue(SHUTDOWN_MESSAGE, timeout=1.0)

self.broadcast_mq = None
except Exception as exc:
logger.warning("Failed to send shutdown signal: %s", exc)
finally:
self.broadcast_mq = None

if self.processes:
join_deadline = time.monotonic() + _WORKER_SHUTDOWN_GRACE_S
for proc in self.processes:
if not proc.is_alive():
continue
proc.join(max(0.0, join_deadline - time.monotonic()))

alive = [proc for proc in self.processes if proc.is_alive()]
for proc in alive:
logger.warning("Terminating diffusion worker %s after timeout", proc.name)
proc.terminate()
for action, grace in (
(None, _WORKER_SHUTDOWN_GRACE_S),
("terminate", _WORKER_TERMINATE_GRACE_S),
("kill", _WORKER_KILL_GRACE_S),
):
if not alive:
break
if action is not None:
for proc in alive:
try:
logger.warning("Calling %s on diffusion worker %s (pid=%s)", action, proc.name, proc.pid)
getattr(proc, action)()
except OSError:
logger.exception("Failed to %s diffusion worker %s (pid=%s)", action, proc.name, proc.pid)

terminate_deadline = time.monotonic() + _WORKER_TERMINATE_GRACE_S
for proc in alive:
proc.join(max(0.0, terminate_deadline - time.monotonic()))
deadline = time.monotonic() + grace
for proc in alive:
try:
proc.join(max(0.0, deadline - time.monotonic()))
except OSError:
logger.exception("Failed to join diffusion worker %s (pid=%s)", proc.name, proc.pid)
alive = [proc for proc in alive if proc.is_alive()]

self.processes = alive
if alive:
logger.error(
"Diffusion worker cleanup incomplete after kill: %s; retaining processes for shutdown retry",
[(proc.name, proc.pid) for proc in alive],
)


class MultiprocDiffusionExecutor(DiffusionExecutor):
Expand Down Expand Up @@ -1006,8 +1038,12 @@ def check_health(self) -> None:
def shutdown(self) -> None:
self._closed = True
self._pump_stop.set()
cleaner = self._shutdown_cleaner
try:
self._finalizer()
if self._finalizer.alive:
self._finalizer()
elif cleaner is not None:
cleaner()
finally:
pump_threads = getattr(self, "_result_pump_threads", [])
for thread in pump_threads:
Expand All @@ -1031,5 +1067,6 @@ def shutdown(self) -> None:
self._rpc_futures.clear()
self._output_futures.clear()
self._batch_split_map.clear()
self._shutdown_cleaner = None
self._processes = []
self._processes = (cleaner.processes or []) if cleaner is not None else []
if not self._processes:
self._shutdown_cleaner = None
Loading