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
214 changes: 153 additions & 61 deletions cvs/runners/aorta.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,11 @@
import shlex
import subprocess
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from functools import partial
from pathlib import Path
from threading import Lock
from typing import Any, Dict, List, Optional, Tuple, TYPE_CHECKING
from threading import Event, Lock, Thread
from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING

if TYPE_CHECKING:
import docker
Expand Down Expand Up @@ -201,7 +200,11 @@ def __init__(self, config: AortaConfig):
# Thread-safe storage for parallel deployment
self._docker_clients: Dict[str, docker.DockerClient] = {}
self._containers: Dict[str, Container] = {}
self._lock = Lock() # Protects _docker_clients and _containers
self._lock = Lock() # Protects _docker_clients, _containers, and _teardown_started
# Set under _lock once teardown() has taken its snapshot of _containers, so a
# straggling setup thread that registers after that point notices and cleans up
# after itself instead of leaking a container teardown() will never see again.
self._teardown_started = False

def validate_config(self) -> List[str]:
"""Validate Aorta-specific configuration."""
Expand Down Expand Up @@ -502,12 +505,62 @@ def _exec_in_container_streaming(

return exit_code, '\n'.join(output_lines)

def _setup_single_node(self, node: str) -> Tuple[str, bool, Optional[str]]:
@staticmethod
def _run_bounded_parallel(
tasks: Dict[Any, Callable[[], Any]], timeout_seconds: float
) -> Tuple[Dict[Any, Any], Dict[Any, Exception], List[Any]]:
"""
Run each zero-arg callable in ``tasks`` in its own thread, bounded by a
single deadline shared across all of them.

``ThreadPoolExecutor`` worker threads are non-daemon; CPython's atexit
hook joins every such thread at interpreter shutdown regardless of
``executor.shutdown(wait=False)``, so a node genuinely stuck in a
blocking Docker/SSH call would still hang the whole process at exit,
not just this method. Daemon threads are exempt from that join, so a
stuck task here can never block CVS from exiting.

Returns ``(results, errors, timed_out_keys)`` keyed by ``tasks``' keys.
A key lands in exactly one of ``results``, ``errors``, or
``timed_out_keys``.
"""
results: Dict[Any, Any] = {}
errors: Dict[Any, Exception] = {}

def _worker(key: Any, fn: Callable[[], Any]) -> None:
try:
results[key] = fn()
except Exception as e: # noqa: BLE001 - reported back per-key, not swallowed
errors[key] = e

threads = {key: Thread(target=_worker, args=(key, fn), daemon=True) for key, fn in tasks.items()}
for t in threads.values():
t.start()

deadline = time.monotonic() + timeout_seconds
timed_out = []
for key, t in threads.items():
t.join(timeout=max(0.0, deadline - time.monotonic()))
if t.is_alive():
timed_out.append(key)

return results, errors, timed_out

def _setup_single_node(self, node: str, cancel_event: Event) -> Tuple[str, bool, Optional[str]]:
"""
Set up a single node (thread-safe helper for parallel deployment).

Args:
node: Hostname or IP of the node
cancel_event: set by ``setup()`` once its overall deadline has
passed. A node that finishes launching its container after
that point tears the container down itself instead of
registering it into ``self._containers`` — ``teardown()``
may already have taken its snapshot of that dict and won't
see it otherwise, orphaning the container. The check and the
registration happen atomically under ``self._lock`` (see
``self._teardown_started``) so there is no window between
them for ``teardown()`` to race past.

Returns:
Tuple of (node, success, error_message)
Expand All @@ -529,9 +582,23 @@ def _setup_single_node(self, node: str) -> Tuple[str, bool, Optional[str]]:
# Launch container
container = self._launch_container(client, node)

# Thread-safe update of shared state
# Check-and-register must be atomic: if this happened as two separate
# steps, teardown() could take its snapshot of self._containers in the
# gap between the check and the registration and never see this
# container, orphaning it.
with self._lock:
self._containers[node] = container
too_late = cancel_event.is_set() or self._teardown_started
if not too_late:
self._containers[node] = container

if too_late:
log.warning(f"Setup on {node} finished after the deadline; cleaning up its container")
try:
container.stop(timeout=30)
container.remove(force=True)
except Exception as e:
log.warning(f"Error removing late container on {node}: {e}")
return (node, False, f"Setup timed out after {self.config.timeout_seconds}s")

# Build RCCL if not skipping
if not self.config.skip_rccl_build:
Expand Down Expand Up @@ -586,6 +653,9 @@ def setup(self) -> bool:

This significantly reduces setup time for multi-node clusters.
"""
with self._lock:
self._teardown_started = False

if not self.config.aorta_path.exists() and not self._ensure_aorta_repo():
log.error("Aorta path does not exist and auto-clone failed or is disabled")
return False
Expand All @@ -599,26 +669,31 @@ def setup(self) -> bool:

log.info(f"Setting up {num_nodes} node(s) in parallel...")

# Use ThreadPoolExecutor for parallel deployment
# Max workers = number of nodes (each node gets its own thread)
with ThreadPoolExecutor(max_workers=num_nodes) as executor:
# Submit all setup tasks
futures = {executor.submit(self._setup_single_node, node): node for node in nodes}

# Collect results as they complete
failed_nodes = []
for future in as_completed(futures):
node = futures[future]
try:
node_name, success, error_msg = future.result()
if not success:
log.error(f"Setup failed on {node_name}: {error_msg}")
failed_nodes.append((node_name, error_msg))
else:
log.info(f"Setup completed on {node_name}")
except Exception as e:
log.exception(f"Unexpected error setting up {node}: {e}")
failed_nodes.append((node, str(e)))
# Bounded by timeout_seconds so a stuck image pull or RCCL build cannot
# hang setup() forever; a still-stuck node's thread is abandoned (daemon)
# rather than joined, so it can't hang CVS itself either.
cancel_event = Event()
tasks = {node: partial(self._setup_single_node, node, cancel_event) for node in nodes}
results, errors, timed_out = self._run_bounded_parallel(tasks, self.config.timeout_seconds)
# Deadline has now passed for anyone still running; tell them to clean
# up their own container instead of registering it after the fact.
cancel_event.set()

failed_nodes = []
for node in nodes:
if node in timed_out:
log.error(f"Setup timed out on {node} after {self.config.timeout_seconds}s")
failed_nodes.append((node, f"Timed out after {self.config.timeout_seconds}s"))
elif node in errors:
log.exception(f"Unexpected error setting up {node}: {errors[node]}")
failed_nodes.append((node, str(errors[node])))
else:
node_name, success, error_msg = results[node]
if not success:
log.error(f"Setup failed on {node_name}: {error_msg}")
failed_nodes.append((node_name, error_msg))
else:
log.info(f"Setup completed on {node_name}")

if failed_nodes:
log.error(f"Setup failed on {len(failed_nodes)}/{num_nodes} nodes:")
Expand Down Expand Up @@ -906,43 +981,50 @@ def run(self, **kwargs) -> RunResult:
f"master={master_addr}:{master_port}"
)

with ThreadPoolExecutor(max_workers=nnodes) as executor:
futures = {}
for rank, node in enumerate(nodes):
cmd = self._build_torchrun_command(
node_rank=rank,
nnodes=nnodes,
master_addr=master_addr,
master_port=master_port,
nproc_per_node=nproc_per_node,
)
future = executor.submit(
partial(
self._run_single_node,
node=node,
node_rank=rank,
launch_cmd=cmd,
env=base_env,
)
tasks = {}
for rank, node in enumerate(nodes):
cmd = self._build_torchrun_command(
node_rank=rank,
nnodes=nnodes,
master_addr=master_addr,
master_port=master_port,
nproc_per_node=nproc_per_node,
)
tasks[(rank, node)] = partial(
self._run_single_node,
node=node,
node_rank=rank,
launch_cmd=cmd,
env=base_env,
)

# Bounded by timeout_seconds so a stalled NCCL collective (or any
# other indefinite hang inside the container) cannot hang run()
# forever with no CVS-level report; a still-stuck node's thread is
# abandoned (daemon) rather than joined, so it can't hang CVS itself.
results, errors, not_done = self._run_bounded_parallel(tasks, self.config.timeout_seconds)

for rank, node in tasks:
if (rank, node) in not_done:
log.error(f"Node {node} (rank {rank}) timed out after {self.config.timeout_seconds}s")
stdout_dict[node] = (
f"Timed out after {self.config.timeout_seconds}s waiting for node to complete"
)
futures[future] = (rank, node)

for future in as_completed(futures):
rank, node = futures[future]
try:
n, ec, out = future.result()
stdout_dict[n] = out
exit_codes[n] = ec
except Exception as e:
log.exception(f"Node {node} (rank {rank}) raised: {e}")
stdout_dict[node] = str(e)
exit_codes[node] = -1
exit_codes[node] = -1
elif (rank, node) in errors:
log.exception(f"Node {node} (rank {rank}) raised: {errors[(rank, node)]}")
stdout_dict[node] = str(errors[(rank, node)])
exit_codes[node] = -1
else:
n, ec, out = results[(rank, node)]
stdout_dict[n] = out
exit_codes[n] = ec

failed = {n: c for n, c in exit_codes.items() if c != 0}
if failed:
log.error(f"Disaggregated run failed on {len(failed)}/{nnodes} nodes: {failed}")
return RunResult(
status=RunStatus.FAILED,
status=RunStatus.TIMEOUT if not_done else RunStatus.FAILED,
start_time=start_time,
end_time=time.time(),
stdout=stdout_dict,
Expand Down Expand Up @@ -1223,7 +1305,19 @@ def teardown(self) -> bool:

success = True

for node, container in self._containers.items():
# Snapshot-and-clear under the same lock _setup_single_node uses to register
# containers, so a straggling setup thread deterministically either lands in
# this snapshot (and gets torn down below) or observes _teardown_started and
# cleans up after itself -- never both, and never neither. Iterating
# self._containers directly here (unsynchronized) would also risk a
# "dictionary changed size during iteration" crash if a setup thread inserted
# into it concurrently.
with self._lock:
self._teardown_started = True
containers = dict(self._containers)
self._containers.clear()

for node, container in containers.items():
try:
log.info(f"Stopping container on {node}...")
container.stop(timeout=30)
Expand Down Expand Up @@ -1261,8 +1355,6 @@ def teardown(self) -> bool:
log.warning(f"Error removing container on {node}: {e}")
success = False

self._containers.clear()

# Close Docker clients - suppress BrokenPipeError during SSH cleanup
for node, client in self._docker_clients.items():
try:
Expand Down
Loading
Loading