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
15 changes: 12 additions & 3 deletions tensorrt_llm/llmapi/mgmn_leader_node.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
'''
This script is used to start the MPICommSession in the rank0 and wait for the
MPI Proxy process to connect and get the MPI task to run.
Expand Down Expand Up @@ -41,14 +43,21 @@ def stop_server_main():
is_server=False,
socket_type=zmq.PAIR)

# Only queue the stop message once a peer is connected, and bound both
# sending and context termination if that peer disappears during shutdown.
queue.socket.setsockopt(zmq.IMMEDIATE, 1)
queue.socket.setsockopt(zmq.SNDTIMEO, 5000)
queue.socket.setsockopt(zmq.LINGER, 1000)
try:
logger_debug(
f"RemoteMpiCommSessionClient [rank{global_mpi_rank()}] send shutdown signal to server\n",
"green")
queue.put(None) # ask RemoteMpiCommSessionServer to shutdown
except zmq.error.ZMQError as e:
logger_debug(f"Error during RemoteMpiCommSessionClient shutdown: {e}\n",
"red")
except zmq.error.ZMQError as error:
raise click.ClickException(
f"Failed to send MPI Comm server shutdown: {error}") from error
finally:
queue.close()


@click.command()
Expand Down
274 changes: 158 additions & 116 deletions tensorrt_llm/llmapi/mpi_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,9 @@
import threading
import time
import traceback
from collections import deque
from collections.abc import Callable
from concurrent.futures import Future, ThreadPoolExecutor, as_completed
from concurrent.futures import Future, ThreadPoolExecutor
from concurrent.futures import wait as futures_wait
from typing import Any, Dict, List, NamedTuple, Optional, Tuple, TypeVar

Expand Down Expand Up @@ -342,6 +343,20 @@ def _process_start_time(pid: int) -> Optional[bytes]:
_DEFAULT_IDENTITY_TIMEOUT = 300.0


def _mgmn_shutdown_grace_seconds() -> float:
"""Bound draining a failed batch and shutting down the server-owned world."""
raw = os.environ.get("TLLM_MGMN_SHUTDOWN_GRACE_SECONDS", "60")
try:
value = float(raw)
if math.isfinite(value) and value > 0:
return value
except ValueError:
pass
logger.warning(
f"Ignoring invalid TLLM_MGMN_SHUTDOWN_GRACE_SECONDS={raw!r}; using 60s")
return 60.0


def _identity_barrier_timeout() -> float:
"""Deadline for the ``wait_shutdown`` worker-identity barrier, in seconds.

Expand Down Expand Up @@ -872,7 +887,6 @@ def __init__(self,
socket_type=zmq.PAIR,
use_hmac_encryption=True)
self.comm = comm
self.results = [] # the results may arrive in any order

if self.comm is not None:
self.session = MpiCommSession(n_workers=self.comm.Get_size(),
Expand All @@ -891,7 +905,7 @@ def task_wrapper(task: Callable[..., T], *args, **kwargs) -> T:
f"MpiCommSession rank{mpi_rank()} start task [{task}] with args: {args} and kwargs: {kwargs}\n",
"green")

# wait for all ranks to start the task
# Pin one task to each rank before any worker can accept another task.
mpi_barrier()

try:
Expand All @@ -901,130 +915,158 @@ def task_wrapper(task: Callable[..., T], *args, **kwargs) -> T:
f"MpiCommSession rank{mpi_rank()} task [{task}] failed with exception: {e}\n",
"red")
traceback.print_exc()
raise e
raise
finally:
logger_debug(
f"MpiCommSession rank{mpi_rank()} task [{task}] finished\n",
"green")
mpi_barrier()
# Let exceptions reach their futures even when a peer is stuck.
# The server gates the next batch on completion of all futures.

def serve(self):
logger_debug(f"RemoteMpiCommSessionServer listening on {self.addr}\n",
"yellow")
pending_futures = []
while True:
# Wait for any pending futures from previous tasks to complete
# This ensures all ranks are ready before accepting the next task
if pending_futures:
logger_debug(
f"RemoteMpiCommSessionServer waiting for {len(pending_futures)} pending futures to complete\n",
"grey")
n_failed = 0
first_exc = None
# Use as_completed so that failures are logged as soon as
# they occur rather than blocking behind a stuck future.
for future in as_completed(pending_futures):
try:
future.result() # Wait for completion
except Exception as e:
n_failed += 1
if first_exc is None:
first_exc = e
print_colored(
f"RemoteMpiCommSessionServer: MPI worker future "
f"failed: {type(e).__name__}: {e}\n", "red")
if n_failed == len(pending_futures):
# All workers failed — no point waiting further.
break
if n_failed:
logger.error(
f"RemoteMpiCommSessionServer: {n_failed}/"
f"{len(pending_futures)} MPI worker(s) failed. "
f"First error: {first_exc}")
pending_futures.clear()
logger_debug(
"RemoteMpiCommSessionServer all pending futures completed\n",
"grey")
def _shutdown_session(self, grace: float):
"""Close the world owned by this server, or abort it within the deadline.

message: Optional[RemoteTask] = self.queue.get()
if message is None:
logger_debug(
f"RemoteMpiCommSessionServer [rank{global_mpi_rank()}] received shutdown signal\n",
"green")
self.session.shutdown_abort()
break
else:
logger_debug(
f"RemoteMpiCommSessionServer [rank{global_mpi_rank()}] received task [{message.task}] from {self.addr}\n",
"green")
futures = self.session.submit(
RemoteMpiCommSessionServer.task_wrapper, message.task,
*message.args, **message.kwargs)
self.num_results = self.session.n_workers
assert len(futures) == self.num_results == mpi_world_size()
# Store futures to wait for them before the next task
pending_futures = list(futures)
for future in futures:
if message.sync:
future.add_done_callback(self.mpi_future_callback)
else:
# Fire-and-forget tasks have no result channel, but a
# crashed worker must still reach the client (the
# client-side session has no futures to watch); see
# RemoteWorkerDeath.
future.add_done_callback(self.mpi_async_error_callback)

def mpi_async_error_callback(self, future):
"""Forward a worker exception to the client for async tasks.

Runs on the executor's callback thread, like the existing sync-path
mpi_future_callback (same pre-existing cross-thread ZMQ-put pattern).
Best-effort: the socket may already be closed at shutdown.
MpiCommSession.shutdown deliberately leaves the shared COMM_WORLD pool
alive for reuse. Only the server's final shutdown owns closing it.
"""
if future.cancelled():
return
exc = future.exception()
if exc is None:
if grace <= 0:
self.session.abort()
return
print_colored(
f"RemoteMpiCommSessionServer: async MPI worker failed, forwarding "
f"to client: {type(exc).__name__}: {exc}\n", "red")
try:
self.queue.put(RemoteWorkerDeath.from_exception(exc))
except Exception as e:
logger_debug(f"Failed to forward worker death to client: {e}\n",
"red")
executor = None
if (isinstance(self.session, MpiCommSession)
and self.session.mpi_pool is not None
and self.session.mpi_pool is MPINodeState._global_mpi_pool):
executor = MPINodeState._global_comm_executor
finished = threading.Event()
errors = []

def shutdown():
try:
self.session.shutdown()
if executor is not None:
executor.__exit__(None, None, None)
MPINodeState._global_comm_executor = None
MPINodeState._global_mpi_pool = None
except BaseException as error:
errors.append(error)
finally:
finished.set()

thread = threading.Thread(target=shutdown,
name="RemoteMpiSessionShutdown",
daemon=True)
thread.start()
if not finished.wait(grace):
logger.error(f"Remote MPI shutdown exceeded {grace}s; aborting")
self.session.abort()
elif errors:
logger.error(f"Remote MPI shutdown failed: {errors[0]!r}; aborting")
self.session.abort()
else:
thread.join()

def mpi_future_callback(self, future):
logger_debug(f"rank{global_mpi_rank()} got future: {future}\n", "red")
if future.exception() is not None:
logger_debug(
f"mpi_future got exception: {future.exception()}, quitting\n",
"red")
self.queue.put(future.exception())
return
def serve(self):
"""Handle completions and control messages without waiting on a rank.

result = future.result()
self.results.append(result)
logger_debug(
f"RemoteMpiCommSessionServer working status: {len(self.results)}/{self.num_results}\n",
"grey")
if len(self.results) == self.num_results:
logger_debug(
"RemoteMpiCommSessionServer received all results, sending to client\n",
"green")
All socket operations run here, never in executor callbacks. Normal
tasks are serialized by batch; a stop request can bypass that gate.
"""
grace = _mgmn_shutdown_grace_seconds()
# A disconnected client must not prevent world teardown. Allow a short
# delivery window, including for the first worker-death notification.
self.queue.socket.setsockopt(zmq.SNDTIMEO, 1000)
self.queue.socket.setsockopt(zmq.LINGER, 1000)
queued_tasks = deque()
pending = []
results = []
sync = False
first_error = None
drain_deadline = None
stop_deadline = None
try:
while True:
for future in pending[:]:
if not future.done():
continue
pending.remove(future)
try:
results.append(future.result())
except Exception as error:
logger.error(f"Remote MPI worker failed: {error!r}")
if first_error is None:
first_error = error
# Preserve the sync API's one exception response,
# but never mix partial results into the next batch.
response = (error if sync else
RemoteWorkerDeath.from_exception(error))
try:
self.queue.put(response)
except Exception as send_error:
logger.error(
f"Failed to send MPI worker error: {send_error!r}"
)
raise
if not sync:
raise RuntimeError(
"Remote MPI asynchronous task failed"
) from error
drain_deadline = time.monotonic() + grace

if drain_deadline is not None and pending:
if time.monotonic() >= drain_deadline:
raise RuntimeError(
"MPI workers did not drain after a task failure"
) from first_error

if not pending:
if sync and first_error is None:
self.queue.put(results)
# Results and error state belong to a single submitted batch.
results = []
sync = False
first_error = None
drain_deadline = None
if queued_tasks:
message = queued_tasks.popleft()
pending = list(
self.session.submit(self.task_wrapper, message.task,
*message.args,
**message.kwargs))
assert len(pending) == self.session.n_workers
sync = message.sync

if stop_deadline is not None:
if not pending and not queued_tasks:
break
if time.monotonic() >= stop_deadline:
raise RuntimeError(
"Remote MPI shutdown deadline expired")
# Preserve tasks submitted before stop, but bound the whole
# drain so a wedged rank cannot strand server teardown.
time.sleep(0.1)
continue

# Keep receiving stop requests even while workers are stuck.
# Other requests wait until every future in the batch is done.
if self.queue.poll(0.1):
message = self.queue.get()
if message is None:
stop_deadline = time.monotonic() + grace
else:
queued_tasks.append(message)
finally:
try:
self.queue.put_noblock(self.results, retry=2)
except zmq.ZMQError as e:
# The client could be shutdown first.
if e.errno == zmq.EAGAIN:
pass
else:
raise e

logger_debug("RemoteMpiCommSessionServer sent results to client\n",
"green")
self.results.clear()
deadlines = [
deadline for deadline in (drain_deadline, stop_deadline)
if deadline is not None
]
if deadlines:
grace = max(0.0,
min(grace,
min(deadlines) - time.monotonic()))
self._shutdown_session(grace)
finally:
self.queue.close()


def find_free_port() -> int:
Expand Down
Loading
Loading