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
64 changes: 49 additions & 15 deletions tensorrt_llm/executor/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ def _check_collective_rpc_guard(

class GenerationExecutorProxy(GenerationExecutor):
READY_SIGNAL = b"READY"
WORKER_PROCESS_IDENTITIES_SIGNAL = b"WORKER_PROCESS_IDENTITIES"

def __init__(
self,
Expand Down Expand Up @@ -633,6 +634,9 @@ def mpi_done_callback(future: concurrent.futures.Future):
k: v
for k, v in worker_kwargs.items() if k != 'tokenizer'
}
worker_process_identities_signal = (
self.WORKER_PROCESS_IDENTITIES_SIGNAL
if self._can_monitor_worker_processes() else None)

self.mpi_futures = self.mpi_session.submit(
worker_main,
Expand All @@ -641,23 +645,14 @@ def mpi_done_callback(future: concurrent.futures.Future):
tracer_init_kwargs=tracer_init_kwargs,
_torch_model_class_mapping=MODEL_CLASS_MAPPING,
ready_signal=GenerationExecutorProxy.READY_SIGNAL,
worker_process_identities_signal=worker_process_identities_signal,
)
for fut in self.mpi_futures:
fut.add_done_callback(mpi_done_callback)

self.workers_started = True

while True:
if self.worker_init_status_queue.poll(1):
status = self.worker_init_status_queue.get()
# Send ACK to the worker
self.worker_init_status_queue.put("ACK")
logger.info("get signal from executor worker")
break
if any(fut.done() for fut in self.mpi_futures):
logger.error("Executor worker died during initialization.")
raise RuntimeError("Executor worker died during initialization")
self._handle_background_error()
status = self._wait_for_executor_workers_ready()

ready_signal, error_trace = status[:2]
if ready_signal != GenerationExecutorProxy.READY_SIGNAL:
Expand All @@ -669,7 +664,41 @@ def mpi_done_callback(future: concurrent.futures.Future):
raise RuntimeError(
"Executor worker returned error") from ready_signal

self._register_worker_processes(status)
def _wait_for_executor_workers_ready(self) -> tuple:
"""Wait for worker readiness while monitoring published processes."""
worker_processes_registered = False

while True:
if self.worker_init_status_queue.poll(1):
status = self.worker_init_status_queue.get()
# Send ACK to the worker
self.worker_init_status_queue.put("ACK")
logger.info("get signal from executor worker")

signal = status[0]
if signal == self.WORKER_PROCESS_IDENTITIES_SIGNAL:
if len(status) != 3:
raise RuntimeError(
"Executor worker returned invalid process identities"
)
self._register_worker_processes(status)
worker_processes_registered = True
continue

# Backward compatibility for workers that only publish their
# identities together with READY.
if (signal == self.READY_SIGNAL
and not worker_processes_registered):
self._register_worker_processes(status)
return status

if (self._check_mpi_workers() or self._check_remote_worker_death()):
error = self._fatal_error or RuntimeError(
"MPI worker exited unexpectedly")
message = f"Executor worker died during initialization: {error}"
logger.error(message)
raise RuntimeError(message) from error
self._handle_background_error()

def _register_worker_processes(self, status: tuple) -> None:
"""Register identities returned by locally spawned MPI workers.
Expand All @@ -678,12 +707,17 @@ def _register_worker_processes(self, status: tuple) -> None:
reference with a factory, so identify pool-backed sessions by excluding
the external communication session types.
"""
if not isinstance(
self.mpi_session,
(MpiCommSession, RemoteMpiCommSessionClient)) and len(status) == 3:
if self._can_monitor_worker_processes() and len(status) == 3:
worker_process_identities: List[WorkerProcessIdentity] = status[2]
self._worker_process_monitor.register(worker_process_identities)

def _can_monitor_worker_processes(self) -> bool:
"""Return whether the session uses locally spawned MPI workers."""
return not isinstance(
self.mpi_session,
(MpiCommSession, RemoteMpiCommSessionClient),
)

def _abort_all_requests(self):
# The results can be finished during this loop, so self._results may be changed.
for result in list(self._results.values()):
Expand Down
16 changes: 16 additions & 0 deletions tensorrt_llm/executor/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,7 @@ def worker_main(
_torch_model_class_mapping: Optional[dict] = None,
postproc_worker_config: Optional[PostprocWorkerConfig] = None,
ready_signal: Optional[str] = None,
worker_process_identities_signal: Optional[bytes] = None,
is_llm_executor: Optional[
bool] = True, # whether it's the main executor instance
hf_model_dir: Optional[Path] = None,
Expand Down Expand Up @@ -341,6 +342,21 @@ def notify_proxy_threads_to_quit():
mpi_comm().barrier()
worker_process_identities = mpi_comm().allgather(
capture_worker_process_identity(mpi_rank()))

# Publish the process identities before backend construction begins. Model
# construction can load weights for several minutes, and an externally
# killed worker may not complete its MPI future. Registering the workers at
# this point lets the proxy observe such a death while it is still waiting
# for the READY signal.
if is_leader and worker_process_identities_signal is not None:
identities_msg = (worker_process_identities_signal, None,
worker_process_identities)
if not worker_init_status_queue.notify_with_retry(identities_msg):
# The failed status queue cannot report its own failure. Let this
# escape through the MPI future so the proxy can observe it.
raise RuntimeError(
"Failed to deliver worker process identities to proxy")
Comment thread
chienchunhung marked this conversation as resolved.

logger_debug(f"Worker {mpi_rank()} ready to setup backend...\n", "green")

try:
Expand Down
201 changes: 201 additions & 0 deletions tests/unittest/executor/test_proxy_fast_death.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,10 @@

from tensorrt_llm.executor import EngineDeadError
from tensorrt_llm.executor import proxy as proxy_module
from tensorrt_llm.executor import worker as worker_module
from tensorrt_llm.executor.proxy import GenerationExecutorProxy
from tensorrt_llm.executor.result import GenerationResult
from tensorrt_llm.executor.worker_process_monitor import WorkerProcessIdentity


def test_engine_dead_error_is_importable_and_carries_root_cause():
Expand Down Expand Up @@ -117,6 +119,205 @@ def test_register_worker_processes_with_session_reuse_factory(monkeypatch):
proxy._worker_process_monitor.register.assert_called_once_with(identities)


class _FakeWorkerInitStatusQueue:
"""Worker status queue that returns a fixed message sequence."""

def __init__(self, messages):
self._messages = list(messages)
self.acks = []

def poll(self, timeout):
return bool(self._messages)

def get(self):
return self._messages.pop(0)

def put(self, message):
self.acks.append(message)


class _DeadAfterRegistrationMonitor:
"""Report a worker death only after the proxy registers identities."""

def __init__(self):
self.identities = []

def register(self, identities):
self.identities = identities

def find_dead_worker(self):
return self.identities[0] if self.identities else None


def test_worker_death_before_ready_is_reported_from_registered_identity():
"""A pre-READY worker death must not depend on its MPI future finishing."""
identity = WorkerProcessIdentity(
rank=3, pid=12345, start_time=67890, hostname="localhost", pid_namespace=1
)
identity_status = (GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL, None, [identity])

proxy = _bare_proxy()
proxy.mpi_session = object()
proxy.worker_init_status_queue = _FakeWorkerInitStatusQueue([identity_status])
proxy._worker_process_monitor = _DeadAfterRegistrationMonitor()
proxy.mpi_futures = [_Future()] # Deliberately remains pending.
proxy._fatal_error = None
proxy.doing_shutdown = False
proxy.pre_shutdown = _Mock()
proxy._handle_background_error = _Mock()

with pytest.raises(RuntimeError, match=r"rank 3 \(pid 12345\) exited unexpectedly"):
proxy._wait_for_executor_workers_ready()

assert not proxy.mpi_futures[0].done()
assert proxy.worker_init_status_queue.acks == ["ACK"]
proxy.pre_shutdown.assert_called_once_with()


def test_remote_worker_death_before_ready_is_reported():
"""Remote worker death must be polled before the error monitor starts."""
error = RuntimeError("remote worker died")

proxy = _bare_proxy()
proxy.mpi_session = _Mock()
proxy.mpi_session.check_worker_error.return_value = error
proxy.worker_init_status_queue = _FakeWorkerInitStatusQueue([])
proxy._worker_process_monitor = _Mock()
proxy._worker_process_monitor.find_dead_worker.return_value = None
proxy.mpi_futures = []
proxy._fatal_error = None
proxy._error_queue = _queue.Queue()
proxy.doing_shutdown = False
proxy.pre_shutdown = _Mock()
proxy._handle_background_error = _Mock(
side_effect=AssertionError("wait loop continued after remote worker death")
)

with pytest.raises(RuntimeError, match="remote worker died"):
proxy._wait_for_executor_workers_ready()

proxy.mpi_session.check_worker_error.assert_called_once_with()
assert proxy.worker_init_status_queue.acks == []
assert proxy.mpi_futures == []
assert proxy._fatal_error is error
proxy.pre_shutdown.assert_called_once_with()
proxy._handle_background_error.assert_not_called()


def test_worker_identities_are_registered_once_before_ready():
identity = WorkerProcessIdentity(
rank=0, pid=12345, start_time=67890, hostname="localhost", pid_namespace=1
)
identity_status = (GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL, None, [identity])
ready_status = (GenerationExecutorProxy.READY_SIGNAL, None, [identity])

proxy = _bare_proxy()
proxy.mpi_session = object()
proxy.worker_init_status_queue = _FakeWorkerInitStatusQueue([identity_status, ready_status])
proxy._worker_process_monitor = _Mock()
proxy.mpi_futures = [_Future()]
proxy._handle_background_error = _Mock()

assert proxy._wait_for_executor_workers_ready() == ready_status
proxy._worker_process_monitor.register.assert_called_once_with([identity])
assert proxy.worker_init_status_queue.acks == ["ACK", "ACK"]


def test_worker_error_status_is_not_treated_as_identities():
error = RuntimeError("backend construction failed")
error_status = (error, "traceback", [object()])

proxy = _bare_proxy()
proxy.mpi_session = object()
proxy.worker_init_status_queue = _FakeWorkerInitStatusQueue([error_status])
proxy._worker_process_monitor = _Mock()
proxy.mpi_futures = [_Future()]
proxy._handle_background_error = _Mock()

assert proxy._wait_for_executor_workers_ready() == error_status
proxy._worker_process_monitor.register.assert_not_called()
assert proxy.worker_init_status_queue.acks == ["ACK"]


def test_worker_publishes_identities_before_backend_construction(monkeypatch):
identity = WorkerProcessIdentity(
rank=0, pid=12345, start_time=67890, hostname="localhost", pid_namespace=1
)
events = []

class _FakeComm:
def barrier(self):
pass

def allgather(self, captured_identity):
assert captured_identity == identity
return [identity]

class _FakeInitStatusQueue:
succeeds = True

def notify_with_retry(self, message):
events.append(("notify", message[0]))
return self.succeeds

class _FailingWorker:
def __init__(self, *args, **kwargs):
events.append(("construct", None))
raise RuntimeError("expected construction failure")

init_status_queue = _FakeInitStatusQueue()

def make_ipc_queue(*args, name, **kwargs):
if name == "worker_init_status_queue":
return init_status_queue
return _Mock()

fake_comm = _FakeComm()
monkeypatch.setattr(worker_module, "mpi_comm", lambda: fake_comm)
monkeypatch.setattr(worker_module, "mpi_rank", lambda: 0)
monkeypatch.setattr(worker_module, "capture_worker_process_identity", lambda rank: identity)
monkeypatch.setattr(worker_module, "set_mpi_session_cpp", lambda comm: None)
monkeypatch.setattr(worker_module, "IpcQueue", make_ipc_queue)
monkeypatch.setattr(worker_module, "FusedIpcQueue", lambda *args, **kwargs: _Mock())

worker_queues = _Mock(
frontend_result_queue_addrs=None,
request_queue_addr=("request", b"key"),
worker_init_status_queue_addr=("status", b"key"),
resource_governor_queue_addr=None,
result_queue_addr=("result", b"key"),
)
worker_module.worker_main(
engine=object(),
worker_queues=worker_queues,
log_level=worker_module.logger.level,
worker_cls=_FailingWorker,
ready_signal=GenerationExecutorProxy.READY_SIGNAL,
worker_process_identities_signal=(GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL),
)

assert events[:2] == [
("notify", GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL),
("construct", None),
]

events.clear()
init_status_queue.succeeds = False
with pytest.raises(RuntimeError, match="Failed to deliver worker process identities to proxy"):
worker_module.worker_main(
engine=object(),
worker_queues=worker_queues,
log_level=worker_module.logger.level,
worker_cls=_FailingWorker,
ready_signal=GenerationExecutorProxy.READY_SIGNAL,
worker_process_identities_signal=(
GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL
),
)

assert events == [("notify", GenerationExecutorProxy.WORKER_PROCESS_IDENTITIES_SIGNAL)]


def test_result_step_raises_on_engine_dead():
res = GenerationResult.__new__(GenerationResult)
res.queue = _queue.Queue()
Expand Down
Loading