diff --git a/tensorrt_llm/executor/proxy.py b/tensorrt_llm/executor/proxy.py index 5a446fc4811e..f923b2d45786 100644 --- a/tensorrt_llm/executor/proxy.py +++ b/tensorrt_llm/executor/proxy.py @@ -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, @@ -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, @@ -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: @@ -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. @@ -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()): diff --git a/tensorrt_llm/executor/worker.py b/tensorrt_llm/executor/worker.py index 6db9f3f5878c..a11b8ecd0341 100644 --- a/tensorrt_llm/executor/worker.py +++ b/tensorrt_llm/executor/worker.py @@ -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, @@ -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") + logger_debug(f"Worker {mpi_rank()} ready to setup backend...\n", "green") try: diff --git a/tests/unittest/executor/test_proxy_fast_death.py b/tests/unittest/executor/test_proxy_fast_death.py index c72ed8c21b1a..d6f4fd952648 100644 --- a/tests/unittest/executor/test_proxy_fast_death.py +++ b/tests/unittest/executor/test_proxy_fast_death.py @@ -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(): @@ -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()