# Source code for tensorplay.distributed.elastic.agent.server.api Source: https://www.tensorplay.cn/docs/_modules/tensorplay/distributed/elastic/agent/server/api.html ``` """Elastic agent contract: worker specs, group state, and the run loop. The agent owns one homogeneous group of local workers. Its run loop rendezvous with peer agents, starts workers with per-rank environments, monitors them, restarts on failure (bounded by ``max_restarts``), and re-rendezvous when peer agents report queued nodes (scale-up). """ import abc import json import os import signal import socket import time import traceback import warnings from contextlib import contextmanager from dataclasses import dataclass, field from enum import Enum from typing import Any, Callable from tensorplay.distributed import Store from ...events import Event, EventSource, NodeState, record from ...metrics import prof, put_metric from ...multiprocessing.errors import ProcessFailure, SignalException from ...rendezvous import ( RendezvousGracefulExitError, RendezvousHandler, RendezvousInfo, ) from ...utils.logging import get_logger from ...utils.store import barrier __all__ = [ "WorkerSpec", "Worker", "WorkerState", "WorkerGroup", "RunResult", "ElasticAgent", "SimpleElasticAgent", ] DEFAULT_ROLE = "default" logger = get_logger(__name__) _TERMINAL_STATE_SYNC_ID = "tp_elastic/agent/terminal_state" [docs] @dataclass class WorkerSpec: """Blueprint of the worker group this agent manages. Every node runs the same spec: the same ``role`` name, the same ``local_world_size``, and the same entrypoint semantics, so global rank arithmetic across agents stays consistent. """ role: str local_world_size: int rdzv_handler: RendezvousHandler fn: Callable | None = None entrypoint: Callable | str | None = None args: tuple = () max_restarts: int = 3 monitor_interval: float = 0.1 master_port: int | None = None master_addr: str | None = None local_addr: str | None = None event_log_handler: str = "null" logs_specs: Any | None = None start_method: str = "spawn" redirects: Any = None tee: Any = None log_dir: str | None = None virtual_local_rank: bool = False numa_options: Any = None duplicate_stdout_filters: list[str] | None = None duplicate_stderr_filters: list[str] | None = None def __post_init__(self) -> None: if self.local_world_size <= 0: raise AssertionError if self.monitor_interval <= 0: raise AssertionError if self.fn: warnings.warn( "WorkerSpec.fn is deprecated; use WorkerSpec.entrypoint instead", DeprecationWarning, stacklevel=2, ) self.entrypoint = self.fn if not self.entrypoint: raise AssertionError [docs] def get_entrypoint_name(self) -> str: """Human-readable name of the entrypoint (module path or command).""" if isinstance(self.entrypoint, str): return os.path.basename(self.entrypoint) if self.entrypoint is None: raise AssertionError return self.entrypoint.__qualname__ [docs] @dataclass(init=False) class Worker: """One logical worker slot with its rank assignments.""" id: Any = None local_rank: int = -1 role_rank: int = -1 global_rank: int = -1 role_world_size: int = -1 world_size: int = -1 def __init__( self, local_rank: int = -1, global_rank: int = -1, role_rank: int = -1, world_size: int = -1, role_world_size: int = -1, id: Any = None, ) -> None: self.id = id self.local_rank = local_rank self.role_rank = role_rank self.global_rank = global_rank self.role_world_size = role_world_size self.world_size = world_size def __str__(self) -> str: return ( f"Worker(id={self.id}, local_rank={self.local_rank}, " f"global_rank={self.global_rank}, role_rank={self.role_rank}, " f"world_size={self.world_size}, role_world_size={self.role_world_size})" ) def __repr__(self) -> str: return str(self) [docs] class WorkerState(str, Enum): """State of the worker group in the agent run loop.""" UNKNOWN = "UNKNOWN" INIT = "INIT" HEALTHY = "HEALTHY" UNHEALTHY = "UNHEALTHY" STOPPED = "STOPPED" SUCCEEDED = "SUCCEEDED" FAILED = "FAILED" @staticmethod def is_running(state: "WorkerState") -> bool: return state in {WorkerState.HEALTHY, WorkerState.UNHEALTHY} [docs] class WorkerGroup: """Mutable group state driven by the agent.""" def __init__(self, spec: WorkerSpec) -> None: self.spec = spec self.workers: list[Worker] = [ Worker(local_rank=i) for i in range(spec.local_world_size) ] self.store: Store | None = None self.group_rank = None self.group_world_size = None self.master_addr = None self.master_port = None self.state = WorkerState.INIT class _RoleInstanceInfo: """Agent role descriptor exchanged through the store for rank assignment.""" def __init__( self, role: str, rank: int | None = None, local_world_size: int | None = None ) -> None: self.role = role if local_world_size is None: local_world_size = rank rank = 0 self.rank = int(rank or 0) self.local_world_size = int(local_world_size or 0) def serialize(self) -> bytes: return json.dumps( { "role": self.role, "rank": self.rank, "local_world_size": self.local_world_size, } ).encode() @staticmethod def deserialize(data: bytes) -> "_RoleInstanceInfo": payload = json.loads(data.decode()) return _RoleInstanceInfo( payload["role"], payload.get("rank", 0), payload["local_world_size"] ) @staticmethod def compare(obj1: "_RoleInstanceInfo", obj2: "_RoleInstanceInfo") -> int: if obj1.role == obj2.role: return obj1.rank - obj2.rank return 1 if obj1.role > obj2.role else -1 @staticmethod def find_role_boundaries( roles_infos: list["_RoleInstanceInfo"], role: str ) -> tuple[int, int]: start_idx, end_idx = -1, -1 for idx, role_info in enumerate(roles_infos): if role_info.role == role: start_idx = idx if start_idx == -1 else start_idx end_idx = idx return start_idx, end_idx [docs] @dataclass class RunResult: """Terminal outcome of the agent run for one role.""" state: WorkerState return_values: dict[int, Any] = field(default_factory=dict) failures: dict[int, ProcessFailure] = field(default_factory=dict) stdouts: dict[int, str] = field(default_factory=dict) stderrs: dict[int, str] = field(default_factory=dict) def is_failed(self) -> bool: return self.state == WorkerState.FAILED def _get_fq_hostname() -> str: return socket.getfqdn(socket.gethostname()) [docs] class ElasticAgent(abc.ABC): """Agent interface for one worker-group role.""" @abc.abstractmethod def _start_workers(self, worker_group: WorkerGroup) -> dict[int, Any]: ... @abc.abstractmethod def _stop_workers(self, worker_group: WorkerGroup, is_restart: bool = False) -> None: ... @abc.abstractmethod def _monitor_workers(self, worker_group: WorkerGroup) -> RunResult: ... [docs] def run(self, role: str = DEFAULT_ROLE) -> RunResult: """Run the worker group to a terminal state.""" raise NotImplementedError [docs] def get_worker_group(self, role: str = DEFAULT_ROLE) -> WorkerGroup: """Return the managed group.""" raise NotImplementedError [docs] class SimpleElasticAgent(ElasticAgent): """Reusable agent run loop over one worker group.""" def __init__( self, spec: WorkerSpec | None = None, exit_barrier_timeout: float = 300, shutdown_timeout: int = 30, ) -> None: self._worker_group: WorkerGroup | None = WorkerGroup(spec) if spec else None self._store: Store | None = None self._exit_barrier_timeout = exit_barrier_timeout self._remaining_restarts = spec.max_restarts if spec else 0 self._shutdown_timeout = shutdown_timeout self._total_execution_time = 0 self._in_exit_barrier = False def get_worker_group(self, role: str = DEFAULT_ROLE) -> WorkerGroup: if self._worker_group is None: raise RuntimeError("The agent has no worker group yet") return self._worker_group @abc.abstractmethod def _start_workers(self, worker_group: WorkerGroup) -> dict[int, Any]: ... @abc.abstractmethod def _stop_workers(self, worker_group: WorkerGroup, is_restart: bool = False) -> None: ... @abc.abstractmethod def _monitor_workers(self, worker_group: WorkerGroup) -> RunResult: ... def _shutdown(self, death_sig: "signal.Signals | None" = None, timeout: int = 30) -> None: """Stop workers and release the rendezvous.""" if self._worker_group is None: return try: if self._worker_group.state != WorkerState.STOPPED: self._stop_workers(self._worker_group) finally: if self._worker_group.spec.rdzv_handler is not None: try: self._worker_group.spec.rdzv_handler.shutdown() except Exception: logger.warning("Rendezvous shutdown failed", exc_info=True) @prof def _rendezvous(self, worker_group: WorkerGroup) -> None: """Rendezvous and assign ranks; refreshes all group state.""" spec = worker_group.spec rdzv_info: RendezvousInfo = spec.rdzv_handler.next_rendezvous() store = rdzv_info.store group_rank = rdzv_info.rank group_world_size = rdzv_info.world_size bootstrap = rdzv_info.bootstrap_store_info master_addr = spec.master_addr or (bootstrap.master_addr if bootstrap else "") master_port = spec.master_port or (bootstrap.port if bootstrap else -1) self._store = store workers = self._assign_worker_ranks(store, group_rank, group_world_size, spec) worker_group.workers = workers worker_group.store = store worker_group.group_rank = group_rank worker_group.group_world_size = group_world_size worker_group.master_addr = master_addr worker_group.master_port = master_port restart_count = spec.max_restarts - self._remaining_restarts logger.info( "[%s] Rendezvous complete: restart_count=%s master_addr=%s master_port=%s " "group_rank=%s group_world_size=%s global_ranks=%s world_size=%s", spec.role, restart_count, master_addr, master_port, group_rank, group_world_size, [w.global_rank for w in workers], workers[0].world_size if workers else -1, ) @prof def _assign_worker_ranks( self, store: Store, group_rank: int, group_world_size: int, spec: WorkerSpec ) -> list[Worker]: """Determine global/role ranks for local workers. Each agent publishes its role descriptor; the rank-0 agent aggregates them and writes base ranks per agent; every agent then derives its workers' ranks from its own base. """ if os.environ.get("TP_ELASTIC_WORKER_IDENTICAL", "0") == "1": global_world_size = group_world_size * spec.local_world_size base_global_rank = group_rank * spec.local_world_size base_role_rank = base_global_rank role_world_size = global_world_size else: role_info_prefix = "tp_elastic/role_info/" assigned_prefix = "tp_elastic/assigned_ranks/" agent_role_info = _RoleInstanceInfo( spec.role, group_rank, spec.local_world_size ) store.set( f"{role_info_prefix}{group_rank}", agent_role_info.serialize().decode(), ) if group_rank == 0: role_infos = [] for i in range(group_world_size): raw = store.get(f"{role_info_prefix}{i}") role_infos.append(_RoleInstanceInfo.deserialize(raw)) role_sizes: dict[str, int] = {} global_size = 0 for info in role_infos: role_sizes[info.role] = ( role_sizes.get(info.role, 0) + info.local_world_size ) global_size += info.local_world_size base_global_rank = 0 role_ranks: dict[str, int] = {} for i, info in enumerate(role_infos): payload = json.dumps( [ base_global_rank, global_size, role_ranks.get(info.role, 0), role_sizes[info.role], ] ) store.set(f"{assigned_prefix}{i}", payload) base_global_rank += info.local_world_size role_ranks[info.role] = ( role_ranks.get(info.role, 0) + info.local_world_size ) assigned = json.loads(store.get(f"{assigned_prefix}{group_rank}").decode()) base_global_rank, global_world_size, base_role_rank, role_world_size = assigned workers = [] for local_rank in range(spec.local_world_size): workers.append( Worker( local_rank=local_rank, global_rank=base_global_rank + local_rank, role_rank=base_role_rank + local_rank, world_size=global_world_size, role_world_size=role_world_size, ) ) return workers @prof def _initialize_workers(self, worker_group: WorkerGroup) -> None: """Rendezvous then start a fresh set of workers.""" role = worker_group.spec.role logger.info("[%s] Rendezvous'ing worker group", role) self._rendezvous(worker_group) logger.info("[%s] Starting worker group", role) worker_ids = self._start_workers(worker_group) for local_rank, worker_id in worker_ids.items(): worker_group.workers[local_rank].id = worker_id record( self._construct_event( state=NodeState.STARTING.value, source=EventSource.WORKER, worker=worker_group.workers[local_rank], ), worker_group.spec.event_log_handler, ) worker_group.state = WorkerState.HEALTHY @prof def _restart_workers(self, worker_group: WorkerGroup) -> None: """Stop, re-rendezvous, and start the group.""" role = worker_group.spec.role logger.info("[%s] Stopping worker group for restart", role) self._stop_workers(worker_group, is_restart=True) worker_group.state = WorkerState.STOPPED self._initialize_workers(worker_group) def _record_worker_events(self, result: RunResult) -> None: group = self._worker_group if group is None: return for worker in group.workers: record( self._construct_event( state=self._get_worker_state(worker, result), source=EventSource.WORKER, worker=worker, raw_error=( json.dumps(result.failures[worker.global_rank].error_file_data) if worker.global_rank in result.failures else None ), exit_code=( result.failures[worker.global_rank].exitcode if worker.global_rank in result.failures else None ), worker_pid=( result.failures[worker.global_rank].pid if worker.global_rank in result.failures else None ), ), group.spec.event_log_handler, ) def get_event_failed(self) -> Event: return self._construct_event( state=NodeState.FAILED.value, source=EventSource.AGENT, raw_error=traceback.format_exc(), ) def get_event_succeeded(self) -> Event: return self._construct_event( state=NodeState.SUCCEEDED.value, source=EventSource.AGENT, ) @contextmanager def record_duration(self, state: str): start_time = time.perf_counter() try: yield finally: duration_ms = (time.perf_counter() - start_time) * 1000 group = self._worker_group if group is not None: record( self._construct_event( state=state, source=EventSource.AGENT, duration_ms=duration_ms, ), group.spec.event_log_handler, ) def _construct_event( self, state: str, source: EventSource, worker: Worker | None = None, raw_error: str | None = None, duration_ms: float | None = None, exit_code: int | None = None, worker_pid: int | None = None, ) -> Event: group = self._worker_group if group is None: raise RuntimeError("The agent has no worker group yet") spec = group.spec metadata: dict[str, Any] = { "group_world_size": group.group_world_size, "entry_point": spec.get_entrypoint_name(), "run_id": spec.rdzv_handler.get_run_id(), "group_rank": group.group_rank, "role": spec.role, "hostname": _get_fq_hostname(), "state": state, "total_run_time": self._total_execution_time, "rdzv_backend": spec.rdzv_handler.get_backend(), "raw_error": raw_error, "agent_restarts": spec.max_restarts - self._remaining_restarts, "duration_ms": duration_ms, } if worker is not None: metadata.update( { "global_rank": worker.global_rank, "worker_id": str(worker.id), } ) worker_metadata = { "group_world_size": group.group_world_size, "entry_point": spec.get_entrypoint_name(), } if worker is not None: worker_metadata.update( { "local_rank": (worker.local_rank,), "role_rank": (worker.role_rank,), "role_world_size": (worker.role_world_size,), "exit_code": (exit_code,), "worker_pid": (worker_pid,), } ) metadata["metadata"] = json.dumps(worker_metadata) return Event( source=source, event_type=f"tp_elastic.worker.status.{state}", metadata=metadata, ) def _record_metric_with_condition(self, metric_name: str, condition: bool) -> None: group = self._worker_group if group is None: return put_metric( f"workers.{group.spec.role}.{metric_name}", int(condition), ) def _record_flakiness_metric(self, is_failed: bool = False) -> None: group = self._worker_group if group is None: return spec = group.spec if is_failed: flakiness = 100.0 else: flakiness = 100.0 - 100.0 * (self._remaining_restarts + 1) / ( spec.max_restarts + 1 ) put_metric(f"workers.{spec.role}.flakiness", int(flakiness)) def _record_metrics(self, group_results: RunResult) -> None: group = self._worker_group if group is None: return failed = group_results.is_failed() spec = group.spec restarted = self._remaining_restarts != spec.max_restarts self._record_flakiness_metric(failed) put_metric(f"workers.{spec.role}.run_total", 1) self._record_metric_with_condition( "run_success_with_retries", not failed and restarted ) self._record_metric_with_condition( "run_success_no_retries", not failed and not restarted ) self._record_metric_with_condition( "run_failed_with_retries", failed and restarted ) self._record_metric_with_condition( "run_failed_no_retries", failed and not restarted ) def _get_worker_state(self, worker: Worker, result: RunResult) -> str: failure = result.failures.get(worker.global_rank) if result.state in {WorkerState.UNHEALTHY, WorkerState.FAILED} and not failure: return "TERMINATED" if failure or worker.global_rank in result.return_values: return result.state.value raise ValueError(f"Unknown worker: {worker.global_rank}") @prof def _exit_barrier(self) -> None: """Wait for all agents to reach the exit point before tearing down.""" if self._store is None or self._worker_group is None: return if self._worker_group.group_world_size <= 1: return self._in_exit_barrier = True try: with barrier( self._store, self._worker_group.group_world_size, key_prefix=f"{_TERMINAL_STATE_SYNC_ID}/{self._worker_group.spec.role}", timeout=self._exit_barrier_timeout, ): logger.info( "Exit barrier reached for all %s agents", self._worker_group.group_world_size, ) finally: self._in_exit_barrier = False [docs] def run(self, role: str = DEFAULT_ROLE) -> RunResult: """Run the agent to a terminal state, shutting down on any exit path.""" start_time = time.monotonic() shutdown_called = False try: result = self._invoke_run(role) self._total_execution_time = int(time.monotonic() - start_time) self._record_metrics(result) self._record_worker_events(result) return result except RendezvousGracefulExitError as e: logger.info("Rendezvous gracefully exited: %s", e) return RunResult(state=WorkerState.SUCCEEDED) except SignalException as e: logger.warning("Received %s death signal, shutting down workers", e.sigval) self._shutdown(e.sigval, timeout=self._shutdown_timeout) shutdown_called = True raise finally: if not shutdown_called: self._shutdown(timeout=self._shutdown_timeout) self._total_execution_time = int(time.monotonic() - start_time) @abc.abstractmethod def _invoke_run(self, role: str = DEFAULT_ROLE) -> RunResult: ... ```