-
Notifications
You must be signed in to change notification settings - Fork 2.8k
[https://nvbugs/6697099][fix] Keep MPI progress alive on the PP sample-state relay thread #19643
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1,6 +1,7 @@ | ||
| # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # SPDX-License-Identifier: Apache-2.0 | ||
|
|
||
| import atexit | ||
| import dataclasses | ||
| import datetime | ||
| import math | ||
|
|
@@ -12,9 +13,9 @@ | |
| from concurrent.futures import Future, ThreadPoolExecutor | ||
| from contextlib import contextmanager, nullcontext | ||
| from enum import IntEnum | ||
| from queue import Queue | ||
| from typing import (TYPE_CHECKING, Dict, Iterable, Iterator, List, Literal, | ||
| Optional, Tuple, Union) | ||
| from queue import Empty, Queue | ||
| from typing import (TYPE_CHECKING, Callable, Dict, Iterable, Iterator, List, | ||
| Literal, Optional, Tuple, Union) | ||
|
|
||
| import torch | ||
| from strenum import StrEnum | ||
|
|
@@ -122,6 +123,206 @@ | |
|
|
||
| from ..moe.fused_moe.communication.base import CheckpointableCommunication | ||
|
|
||
| # Idle-time MPI progress for pipeline parallelism. | ||
| # | ||
| # MPI only makes progress on threads that call into MPI. The PP executor drives | ||
| # every host-side PP message from the executor thread, so while that thread is | ||
| # blocked in a CUDA synchronization the process makes no MPI progress. A | ||
| # point-to-point send that the transport placed in the rendezvous protocol | ||
| # only completes once the sending process re-enters MPI, so with pp_size >= 3 | ||
| # one stranded send can close a dependency cycle across the ranks and | ||
| # deadlock the pipeline. | ||
| # | ||
| # The sample-state relay thread, which exists whenever pp_size > 1, breaks the | ||
| # cycle: instead of parking in an untimed Queue.get(), it polls | ||
| # executed_batch_queue and issues a non-matching MPI_Iprobe between polls on | ||
| # its own duplicated communicator. This keeps a progress source alive | ||
| # whenever the relay thread can acquire the GIL; a native call that holds the | ||
| # GIL across a whole CUDA synchronization can still starve it. | ||
| # | ||
| # Each PyExecutor is pumped from start_worker() until shutdown(). Setting | ||
| # TLLM_PP_MPI_PROGRESS_POLL_MS to a value <= 0 disables the pump. | ||
| MPI_PROGRESS_POLL_MS_ENV_VAR_NAME = "TLLM_PP_MPI_PROGRESS_POLL_MS" | ||
| # Matches CPython's default thread switch interval, so the relay thread does | ||
| # not ask for the GIL more often than the interpreter already switches. | ||
| DEFAULT_MPI_PROGRESS_POLL_MS = 5.0 | ||
| # Bounds for a user-supplied period: a tiny period spins a core and competes | ||
| # with the executor thread for the GIL, and a huge one can overflow | ||
| # Queue.get(timeout=...). | ||
| MIN_MPI_PROGRESS_POLL_MS = 0.5 | ||
| MAX_MPI_PROGRESS_POLL_MS = 1000.0 | ||
|
|
||
|
|
||
| def _resolve_mpi_progress_poll_interval_ms() -> float: | ||
| """Resolve the idle-time MPI progress poll period, in milliseconds. | ||
|
|
||
| Returns the default when the variable is unset or not a finite number. A | ||
| value <= 0 is returned unchanged and means "disabled"; any other value is | ||
| clamped into ``[MIN_MPI_PROGRESS_POLL_MS, MAX_MPI_PROGRESS_POLL_MS]``. | ||
|
|
||
| Messages are logged at ERROR because only the leader rank lowers the log | ||
| level, so a WARNING would be dropped on every other rank by default. | ||
| """ | ||
| raw = os.environ.get(MPI_PROGRESS_POLL_MS_ENV_VAR_NAME) | ||
| if raw is None: | ||
| return DEFAULT_MPI_PROGRESS_POLL_MS | ||
| try: | ||
| poll_interval_ms = float(raw) | ||
| except ValueError: | ||
| poll_interval_ms = math.nan | ||
| if not math.isfinite(poll_interval_ms): | ||
| logger.error( | ||
| f"Ignoring malformed {MPI_PROGRESS_POLL_MS_ENV_VAR_NAME}={raw!r}; " | ||
| f"using {DEFAULT_MPI_PROGRESS_POLL_MS} ms.") | ||
| return DEFAULT_MPI_PROGRESS_POLL_MS | ||
| if poll_interval_ms <= 0: | ||
| # Disabled; the caller logs it. | ||
| return poll_interval_ms | ||
| clamped_ms = min(max(poll_interval_ms, MIN_MPI_PROGRESS_POLL_MS), | ||
| MAX_MPI_PROGRESS_POLL_MS) | ||
| if clamped_ms != poll_interval_ms: | ||
| logger.error( | ||
| f"{MPI_PROGRESS_POLL_MS_ENV_VAR_NAME}={raw!r} is outside " | ||
| f"[{MIN_MPI_PROGRESS_POLL_MS}, {MAX_MPI_PROGRESS_POLL_MS}] ms; " | ||
| f"using {clamped_ms} ms instead.") | ||
| return clamped_ms | ||
|
|
||
|
|
||
| def _make_mpi_progress_pump( | ||
| comm, stop_event: threading.Event, quiesced_event: threading.Event | ||
| ) -> Optional[Tuple[Callable[[], bool], float]]: | ||
| """Build the idle-time MPI progress pump for the sample-state relay thread. | ||
|
|
||
| Args: | ||
| comm: The communicator owned by the relay thread. It is passed in | ||
| rather than looked up with ``mpi_comm()``, which on a thread | ||
| without a thread-local binding returns the executor thread's | ||
| communicator. | ||
| stop_event: Set to retire the pump. | ||
| quiesced_event: Set by the pump whenever it returns ``False``, that | ||
| is, once it will issue no further MPI call. The process-exit hook | ||
| waits on it before MPI is finalized. | ||
|
|
||
| Returns: | ||
| A ``(pump, poll_interval_s)`` pair, or ``None`` when the pump is | ||
| disabled or unavailable. ``pump()`` never raises and returns | ||
| ``False`` once it must not be called again. | ||
|
|
||
| The probe is an ``MPI_Iprobe`` on ``PPCommTag.MPI_PROGRESS_PROBE``, a tag | ||
| that is never sent, so it always misses and only drives the progress | ||
| engine. It must stay a plain ``Iprobe``: ``Improbe``/``Mprobe`` remove the | ||
| matched message from the queue and would race pkl5's ``Probe``-then- | ||
| ``Recv`` sequence on the same communicator. For the same reason the pump | ||
| never calls ``Test``/``Wait`` on requests owned by the executor thread. | ||
| """ | ||
| poll_interval_ms = _resolve_mpi_progress_poll_interval_ms() | ||
| if poll_interval_ms <= 0: | ||
| # Logged at ERROR so that a disabled pump is visible on every rank. | ||
| logger.error( | ||
| f"{MPI_PROGRESS_POLL_MS_ENV_VAR_NAME}={poll_interval_ms}: the " | ||
| f"pipeline parallelism MPI progress pump is off and the " | ||
| f"sample-state relay thread keeps its untimed wait " | ||
| f"(nvbugs/6385771).") | ||
| return None | ||
| if mpi_disabled(): | ||
| return None | ||
|
|
||
| from mpi4py import MPI | ||
|
|
||
| # The pump makes the relay thread call MPI concurrently with the executor | ||
| # thread, which requires MPI_THREAD_MULTIPLE. | ||
| provided = MPI.Query_thread() | ||
| if provided < MPI.THREAD_MULTIPLE: | ||
| logger.error( | ||
| f"MPI reports thread level {provided}, below MPI_THREAD_MULTIPLE " | ||
| f"({int(MPI.THREAD_MULTIPLE)}); the pipeline parallelism MPI " | ||
| f"progress pump stays off. A rank that blocks in a CUDA " | ||
| f"synchronization while one of its point-to-point sends is still " | ||
| f"in flight can then deadlock the pipeline (nvbugs/6385771).") | ||
| return None | ||
|
|
||
| assert comm is not None, ("the pipeline parallelism MPI progress pump " | ||
| "needs the relay thread's communicator") | ||
|
|
||
| # Resolve the raw, non-consuming MPI.Comm.Iprobe up front, so that a | ||
| # missing entry point is reported at startup instead of disarming the | ||
| # pump on its first tick. | ||
| iprobe = getattr(comm, "Iprobe", None) | ||
| if iprobe is None: | ||
| logger.error( | ||
| f"{type(comm).__name__} has no Iprobe; the pipeline parallelism " | ||
| f"MPI progress pump stays off. A rank that blocks in a CUDA " | ||
| f"synchronization while one of its point-to-point sends is still " | ||
| f"in flight can then deadlock the pipeline (nvbugs/6385771).") | ||
| return None | ||
|
|
||
| any_source = MPI.ANY_SOURCE | ||
| probe_tag = int(PPCommTag.MPI_PROGRESS_PROBE) | ||
|
|
||
| def pump() -> bool: | ||
| """Drive one progress tick. Never raises; returns still-armed.""" | ||
| # Nothing may escape: if the relay thread dies, nothing drains the | ||
| # bounded executed_batch_queue under | ||
| # TLLM_PP_ASYNC_BROADCAST_SAMPLE_STATE=1 and the executor thread | ||
| # blocks. A torn-down MPI can also fail with non-MPI exceptions, so | ||
| # any failure simply disarms the pump. | ||
| try: | ||
| if stop_event.is_set(): | ||
| quiesced_event.set() | ||
| return False | ||
| # Rejects ticks after MPI_Finalize has completed. A finalize that | ||
| # is still running is covered by the process-exit hook registered | ||
| # in start_worker(). | ||
| if MPI.Is_finalized(): | ||
| quiesced_event.set() | ||
| return False | ||
| iprobe(source=any_source, tag=probe_tag) | ||
| return True | ||
| except Exception as e: # noqa: BLE001 - see the comment above | ||
| try: | ||
| logger.error( | ||
| f"The pipeline parallelism MPI progress pump failed and " | ||
| f"is now off: {e!r}. This rank is back on the pre-fix " | ||
| f"path, where blocking in a CUDA synchronization with a " | ||
| f"send in flight can deadlock the pipeline " | ||
| f"(nvbugs/6385771).") | ||
| except Exception: # noqa: BLE001 - logging must not kill the relay | ||
| pass | ||
| # The probe has returned, so no MPI call is in flight here. | ||
| quiesced_event.set() | ||
| return False | ||
|
|
||
| return pump, poll_interval_ms / 1000.0 | ||
|
|
||
|
|
||
| def _make_mpi_progress_exit_hook(stop_event: threading.Event, | ||
| quiesced_event: threading.Event, | ||
| poll_interval_s: float) -> Callable[[], None]: | ||
| """Build the process-exit hook that retires the pump before MPI_Finalize. | ||
|
|
||
| Some paths never reach ``PyExecutor.shutdown()``, for example non-zero | ||
| ranks with ``orchestrator_type='rpc'``, a failed executor construction or | ||
| the hang-detector branch, so the daemon relay thread could still be inside | ||
| ``MPI_Iprobe`` when mpi4py finalizes MPI at exit. An ``atexit`` callback | ||
| registered from ``start_worker()`` runs before that finalization, whether | ||
| mpi4py finalizes through ``atexit`` (LIFO) or at the C level. | ||
|
|
||
| The hook sets ``stop_event`` and waits, bounded, for ``quiesced_event``. | ||
| The wait is skipped on the normal path, where ``shutdown()`` has already | ||
| set it. The hook holds only the two Events and a float, never the | ||
| executor. | ||
| """ | ||
| # Long enough for the relay thread to observe the flag at its poll | ||
| # cadence, short enough not to delay process exit noticeably. | ||
| quiesce_timeout_s = min(max(20.0 * poll_interval_s, 0.1), 1.0) | ||
|
|
||
| def _retire_mpi_progress_pump() -> None: | ||
| stop_event.set() | ||
| quiesced_event.wait(timeout=quiesce_timeout_s) | ||
|
|
||
| return _retire_mpi_progress_pump | ||
|
|
||
|
|
||
| _UNBOUNDED_STATS_MAX_LEN = -1 | ||
|
|
||
|
|
||
|
|
@@ -1039,6 +1240,19 @@ def on_detected(): | |
| self.worker_started = False | ||
| self.worker_lock = threading.Lock() | ||
| self._broadcast_mpi_comm = None | ||
| # Idle-time MPI progress pump for the sample-state relay thread; see | ||
| # _make_mpi_progress_pump. The (pump, poll interval in seconds) pair | ||
| # is built by start_worker() and stays None when the pump is disabled | ||
| # or unavailable. It is only written on the main thread; retiring the | ||
| # pump goes through the Events instead. _pp_mpi_progress_stop requests | ||
| # the pump to retire, and _pp_mpi_progress_quiesced acknowledges that | ||
| # the relay thread will issue no further MPI call. | ||
| self._pp_mpi_progress: Optional[Tuple[Callable[[], bool], float]] = None | ||
| self._pp_mpi_progress_stop = threading.Event() | ||
| self._pp_mpi_progress_quiesced = threading.Event() | ||
| # Cleared by the relay thread if the timed Queue.get() itself fails; | ||
| # see _wait_for_executed_batch. | ||
| self._pp_mpi_progress_timed_wait_ok = True | ||
| # Secondary MPI communicator and listener thread for multi-rank | ||
| # sleep/wakeup control messages. Both are None until start_worker() | ||
| # calls Dup() (a collective) on the main thread. | ||
|
|
@@ -1430,6 +1644,29 @@ def start_worker(self): | |
| "Create new MPI comm for broadcast sample state thread to avoid deadlock." | ||
| ) | ||
| self._broadcast_mpi_comm = mpi_comm().Dup() | ||
| # Build the idle-time MPI progress pump on the main thread | ||
| # and hand it the relay thread's own communicator, never | ||
| # the executor thread's. | ||
| assert self._broadcast_mpi_comm is not None | ||
| assert self._broadcast_mpi_comm is not mpi_comm() | ||
| # Fresh Events rather than clear(): an executor may be | ||
| # started again after shutdown(), and a flag left set by | ||
| # the previous run would disarm the new pump. | ||
| self._pp_mpi_progress_stop = threading.Event() | ||
| self._pp_mpi_progress_quiesced = threading.Event() | ||
| self._pp_mpi_progress_timed_wait_ok = True | ||
| self._pp_mpi_progress = _make_mpi_progress_pump( | ||
| self._broadcast_mpi_comm, self._pp_mpi_progress_stop, | ||
| self._pp_mpi_progress_quiesced) | ||
| if self._pp_mpi_progress is not None: | ||
| # Retire the pump before MPI_Finalize on exit paths | ||
| # that never reach shutdown(). The hook must not hold | ||
| # a reference to self. | ||
| atexit.register( | ||
| _make_mpi_progress_exit_hook( | ||
| self._pp_mpi_progress_stop, | ||
| self._pp_mpi_progress_quiesced, | ||
| self._pp_mpi_progress[1])) | ||
| broadcast_sample_state_loop = self._broadcast_sample_state_loop | ||
| if is_trace_enabled("TLLM_TRACE_EXECUTOR_LOOP"): | ||
| broadcast_sample_state_loop = trace_func( | ||
|
|
@@ -1718,11 +1955,20 @@ def shutdown(self): | |
| # Since the whole process will shutdown after this `shutdown` call, | ||
| # All threads and memory pools will be freed properly. | ||
| logger.error("Hang detected, shutting down immediately.") | ||
| # Stop the relay thread from probing while the hang detector | ||
| # aborts MPI. A probe already in flight is not interrupted. | ||
| self._pp_mpi_progress_stop.set() | ||
| return | ||
| self.worker_thread.join() | ||
| if self.dist.pp_size > 1: | ||
| # The worker thread has exited, so no PP traffic needs progress. | ||
| self._pp_mpi_progress_stop.set() | ||
| self.executed_batch_queue.put(None) | ||
| self.broadcast_sample_state_handler.join() | ||
| # The relay thread exits through the sentinel without another | ||
| # pump tick, so acknowledge the quiesce here; otherwise the exit | ||
| # hook would wait out its full timeout. | ||
| self._pp_mpi_progress_quiesced.set() | ||
| # Signal non-rank-0 sleep/wakeup listener threads to exit. This runs | ||
| # after the worker thread has joined, which guarantees that the non-rank-0 | ||
| # executor loops have already processed the shutdown broadcast and are | ||
|
|
@@ -3092,6 +3338,51 @@ def handle_executed_batches(executed_batch_num: int): | |
| self._handle_executed_batch(executed_batch) | ||
| self.unhandled_batch_counter -= 1 | ||
|
|
||
| def _wait_for_executed_batch(self) -> Optional[BatchStatePP]: | ||
| """Wait for the next batch to relay, keeping MPI progress alive. | ||
|
|
||
| While the executor thread is blocked in a CUDA synchronization, this | ||
| thread is the process's only MPI progress source, so it polls the | ||
| queue and runs the pump between polls; see _make_mpi_progress_pump. | ||
|
|
||
| Without a pump this is a plain untimed ``Queue.get()``. Once the pump | ||
| retires, the timed wait is kept. If the timed wait itself fails, this | ||
| falls back to the untimed wait for the rest of this executor's life | ||
| rather than letting the relay thread die. | ||
| """ | ||
| progress = self._pp_mpi_progress | ||
| if progress is None or not self._pp_mpi_progress_timed_wait_ok: | ||
| return self.executed_batch_queue.get() | ||
| pump, poll_interval_s = progress | ||
| while True: | ||
| try: | ||
| return self.executed_batch_queue.get(timeout=poll_interval_s) | ||
| except Empty: | ||
| pass | ||
| except Exception as e: # noqa: BLE001 - see below | ||
| # Must not escape: a dead relay thread would block the | ||
| # executor thread on the bounded executed_batch_queue under | ||
| # TLLM_PP_ASYNC_BROADCAST_SAMPLE_STATE=1. | ||
| self._pp_mpi_progress_timed_wait_ok = False | ||
| self._pp_mpi_progress_stop.set() | ||
| self._pp_mpi_progress_quiesced.set() | ||
| try: | ||
| logger.error( | ||
| f"The pipeline parallelism MPI progress pump's timed " | ||
| f"wait failed: {e!r}. This rank falls back to the " | ||
| f"pre-fix untimed wait, where blocking in a CUDA " | ||
| f"synchronization with a send in flight can deadlock " | ||
| f"the pipeline (nvbugs/6385771).") | ||
| except Exception: # noqa: BLE001 - must not kill the relay | ||
| pass | ||
| return self.executed_batch_queue.get() | ||
| # pump() never raises. | ||
| if not pump(): | ||
| # Retire the pump for good. Not logged: pump() has already | ||
| # logged any unexpected disarm, and this is also the normal | ||
| # shutdown path. | ||
| self._pp_mpi_progress_stop.set() | ||
|
|
||
|
Comment on lines
+3341
to
+3385
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🩺 Stability & Availability | 🟡 Minor | ⚡ Quick win Add unit tests for the pump failure paths and the timed-wait fallback. This PR adds failure-recovery behavior that only a failure scenario reaches:
The integration test covers only the successful default path. Suppose Add small unit tests in
As per path instructions: "Authentication, authorization, input validation, error propagation, or failure recovery changes covered only by a happy-path test." 🤖 Prompt for AI AgentsSource: Path instructions |
||
| def _broadcast_sample_state_loop(self): | ||
| logger.debug( | ||
| f"Starting broadcast sample state loop for pp_rank {self.dist.pp_rank}" | ||
|
|
@@ -3109,7 +3400,7 @@ def _broadcast_sample_state_loop(self): | |
| set_thread_local_mpi_comm(broadcast_mpi_comm) | ||
| try: | ||
| while True: | ||
| executed_batch = self.executed_batch_queue.get() | ||
| executed_batch = self._wait_for_executed_batch() | ||
| if executed_batch is None: | ||
| break | ||
| self._ring_broadcast_sample_state(executed_batch) | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🎯 Functional Correctness | 🟡 Minor | ⚡ Quick win
🔎 Supported by static analysis
🏁 Script executed:
Repository: NVIDIA/TensorRT-LLM
Length of output: 157
Add unit tests for the MPI progress poll interval rules.
_resolve_mpi_progress_poll_interval_msadds configuration behavior for default, malformed, non-finite, disabled, and clamped values. The existing tests do not exercise these rules, so regressions can pass unnoticed.Add a parameterized unit test that uses
monkeypatch.setenvand covers:5.0"abc","nan","inf"→5.0"0","-1"→ unchanged"0.1"→0.5"5000"→1000.0"7"→7.0Place the test under
tests/unittest/_torch/executor/.🤖 Prompt for AI Agents
Source: Path instructions