Codeseys/composer-replication-framework
0
1"""ServerlessExecutor Protocol + LocalProcessExecutor.2 3Per ADR-005:4- `ServerlessExecutor` is a structural Protocol — backends implement it5 by writing a class with the right methods, no formal inheritance needed.6- `LocalProcessExecutor` is the reference implementation that uses Python's7 `multiprocessing` module. It's used for tests and for development; the8 cloud adapters (Modal, HF Jobs, …) implement the same Protocol against9 their respective backends.10"""11from __future__ import annotations12 13import multiprocessing as mp14import sys15import time16from dataclasses import dataclass, field17from typing import Any, Callable, Mapping, Protocol, runtime_checkable18 19 20@dataclass21class ReplicaHandle:22 """Opaque handle to a launched replica. Backend-specific contents.23 24 All executors return `list[ReplicaHandle]` from `launch_replicas`.25 Each handle's `metadata` dict is backend-specific; users shouldn't26 rely on its shape.27 """28 rank: int29 backend_name: str30 metadata: dict[str, Any] = field(default_factory=dict)31 """Backend-specific data (e.g. Modal call ID, HF Jobs job ID, local32 Process object). Not stable across backends."""33 34 35@runtime_checkable36class ServerlessExecutor(Protocol):37 """Uniform interface for launching N replicas on a serverless backend.38 39 Implementations: `LocalProcessExecutor` (test/dev), `ModalSpawnExecutor`40 (Modal, production), `EKSExecutor` (Amazon EKS / Kubernetes Indexed Job,41 production), `ModalExecutor` / `HFJobsExecutor` (v0 skeletons). Future42 adapters: `RunPodExecutor`, `SageMakerExecutor`.43 44 Note on rank assignment: the Protocol guarantees that handles are45 returned in rank order (`handles[i].rank == i`). The replica entrypoint46 learns its own rank either from an environment variable47 (`REPLICA_RANK`) or from a backend-provided mechanism (Modal's48 `Function.shard_rank`, etc.). The executor abstraction normalizes49 rank by setting the env var.50 """51 backend_name: str52 supports_inter_replica_network: bool53 54 def launch_replicas(55 self,56 n_replicas: int,57 entrypoint: str | Callable[..., Any],58 entrypoint_args: Mapping[str, Any],59 *,60 gpu: str | None = None,61 timeout: int = 3600,62 ) -> list[ReplicaHandle]:63 """Spin up N replicas in parallel.64 65 Args:66 n_replicas: number of replicas to launch67 entrypoint: either an importable Python path (e.g.68 "composer_replication.diloco.serverless.replica_entrypoint")69 or a Callable (Local executor only).70 entrypoint_args: kwargs passed to the entrypoint. The kwarg71 `rank_env` (default "REPLICA_RANK") names the environment72 variable in which the rank will be set on the replica.73 gpu: backend-specific GPU spec, e.g. "A100", "H100". `None`74 means CPU-only (smoke tests).75 timeout: per-replica wall-clock timeout in seconds.76 77 Returns:78 list[ReplicaHandle] of length n_replicas, in rank order.79 """80 ...81 82 def poll(self, handle: ReplicaHandle) -> str:83 """Poll a replica's status. Returns one of:84 "pending" | "running" | "succeeded" | "failed" | "cancelled".85 """86 ...87 88 def stream_logs(self, handle: ReplicaHandle, *, n_lines: int = 200) -> str:89 """Read up to n_lines of recent stdout/stderr from a replica."""90 ...91 92 def cancel(self, handle: ReplicaHandle) -> None:93 """Best-effort cancel. No exception if already terminated."""94 ...95 96 def collect(97 self,98 handles: list[ReplicaHandle],99 *,100 timeout: int | None = None,101 ) -> list[dict[str, Any]]:102 """Block until all replicas finish; return per-replica result dicts.103 104 Each result dict contains at least:105 {"rank": int, "status": str, "exit_code": int | None,106 "error": str | None}107 """108 ...109 110 111# ---------------------------------------------------------------------112# LocalProcessExecutor — reference implementation using multiprocessing113# ---------------------------------------------------------------------114 115 116def _local_replica_target(117 rank: int,118 rank_env: str,119 entrypoint: Any,120 entrypoint_args: Mapping[str, Any],121 result_queue: mp.Queue,122) -> None:123 """multiprocessing target — runs in the child process."""124 import os125 import traceback126 127 os.environ[rank_env] = str(rank)128 try:129 if callable(entrypoint):130 result = entrypoint(**entrypoint_args)131 elif isinstance(entrypoint, str):132 # importable path133 mod_path, _, fn_name = entrypoint.rpartition(".")134 if not mod_path:135 # Top-level script path; just import it and call its main()136 import importlib137 mod = importlib.import_module(entrypoint)138 fn = getattr(mod, "main", None)139 if fn is None:140 raise RuntimeError(141 f"entrypoint '{entrypoint}' has no main() function"142 )143 result = fn(**entrypoint_args)144 else:145 import importlib146 mod = importlib.import_module(mod_path)147 fn = getattr(mod, fn_name)148 result = fn(**entrypoint_args)149 else:150 raise TypeError(151 f"entrypoint must be Callable or importable str, got {type(entrypoint)!r}"152 )153 result_queue.put({"rank": rank, "status": "succeeded",154 "exit_code": 0, "error": None, "result": result})155 except Exception as e:156 tb = traceback.format_exc()157 result_queue.put({"rank": rank, "status": "failed",158 "exit_code": 1, "error": f"{e!r}\n{tb}", "result": None})159 160 161class LocalProcessExecutor:162 """Runs replicas as subprocesses on the local machine.163 164 Use cases:165 - Test the serverless layer end-to-end without cloud spend.166 - Develop the algorithm locally with N>1 replicas and `file://`167 rendezvous before deploying to Modal/HF Jobs.168 - CI smoke tests.169 """170 backend_name = "local_process"171 supports_inter_replica_network = True # localhost works172 173 def __init__(self) -> None:174 # use 'spawn' so the child has a fresh interpreter (avoid CUDA fork issues)175 try:176 self._ctx = mp.get_context("spawn")177 except ValueError:178 # Fallback for environments where 'spawn' isn't available179 self._ctx = mp.get_context()180 self._handles: dict[int, dict[str, Any]] = {}181 182 def launch_replicas(183 self,184 n_replicas: int,185 entrypoint: str | Callable[..., Any],186 entrypoint_args: Mapping[str, Any],187 *,188 gpu: str | None = None,189 timeout: int = 3600,190 ) -> list[ReplicaHandle]:191 if gpu is not None:192 # Local executor doesn't pin GPUs; emit a soft warning.193 sys.stderr.write(194 f"[LocalProcessExecutor] gpu={gpu!r} ignored — "195 f"local processes share whatever GPUs are visible.\n"196 )197 rank_env = entrypoint_args.get("rank_env", "REPLICA_RANK")198 199 handles: list[ReplicaHandle] = []200 result_queue: mp.Queue = self._ctx.Queue()201 for rank in range(n_replicas):202 args_for_rank = dict(entrypoint_args)203 args_for_rank.pop("rank_env", None)204 proc = self._ctx.Process(205 target=_local_replica_target,206 args=(rank, rank_env, entrypoint, args_for_rank, result_queue),207 name=f"composer-replica-{rank}",208 )209 proc.start()210 handle = ReplicaHandle(211 rank=rank, backend_name=self.backend_name,212 metadata={"pid": proc.pid, "start_ts": time.time()},213 )214 self._handles[rank] = {"proc": proc, "queue": result_queue,215 "deadline": time.time() + timeout,216 "result": None}217 handles.append(handle)218 return handles219 220 def poll(self, handle: ReplicaHandle) -> str:221 meta = self._handles.get(handle.rank)222 if meta is None:223 return "cancelled"224 proc: mp.Process = meta["proc"]225 if proc.is_alive():226 return "running"227 if meta.get("result") is not None:228 return meta["result"]["status"]229 # Process exited; read result if available230 try:231 queue: mp.Queue = meta["queue"]232 while not queue.empty():233 r = queue.get_nowait()234 self._handles[r["rank"]]["result"] = r235 except Exception:236 pass237 if meta.get("result") is not None:238 return meta["result"]["status"]239 return "failed" if proc.exitcode != 0 else "succeeded"240 241 def stream_logs(self, handle: ReplicaHandle, *, n_lines: int = 200) -> str:242 # multiprocessing.Process doesn't natively capture stdout; we'd243 # need a Pipe or file redirection. For the local reference impl,244 # we just point the user at the result dict's `error` field.245 meta = self._handles.get(handle.rank)246 if meta is None:247 return f"<replica {handle.rank}: no metadata>"248 if meta.get("result"):249 err = meta["result"].get("error") or ""250 return f"[rank {handle.rank}] {err[-2000:]}"251 return f"<replica {handle.rank}: still running, no captured logs>"252 253 def cancel(self, handle: ReplicaHandle) -> None:254 meta = self._handles.get(handle.rank)255 if meta is None:256 return257 proc: mp.Process = meta["proc"]258 if proc.is_alive():259 proc.terminate()260 proc.join(timeout=5)261 if proc.is_alive():262 proc.kill()263 264 def collect(265 self,266 handles: list[ReplicaHandle],267 *,268 timeout: int | None = None,269 ) -> list[dict[str, Any]]:270 deadline = time.time() + (timeout if timeout is not None else 3600)271 # Wait for all processes to finish272 for h in handles:273 meta = self._handles.get(h.rank)274 if meta is None:275 continue276 proc: mp.Process = meta["proc"]277 remaining = max(0.0, deadline - time.time())278 proc.join(timeout=remaining)279 if proc.is_alive():280 proc.terminate()281 proc.join(timeout=5)282 # Drain results283 results_by_rank: dict[int, dict[str, Any]] = {}284 for h in handles:285 meta = self._handles.get(h.rank)286 if meta is None:287 results_by_rank[h.rank] = {288 "rank": h.rank, "status": "cancelled",289 "exit_code": None, "error": "no metadata", "result": None,290 }291 continue292 queue: mp.Queue = meta["queue"]293 while not queue.empty():294 try:295 r = queue.get_nowait()296 results_by_rank[r["rank"]] = r297 except Exception:298 break299 if h.rank not in results_by_rank:300 proc: mp.Process = meta["proc"]301 results_by_rank[h.rank] = {302 "rank": h.rank,303 "status": "succeeded" if proc.exitcode == 0 else "failed",304 "exit_code": proc.exitcode,305 "error": None if proc.exitcode == 0 else f"exit code {proc.exitcode}",306 "result": None,307 }308 return [results_by_rank[h.rank] for h in handles]309 310 311__all__ = ["LocalProcessExecutor", "ReplicaHandle", "ServerlessExecutor"]312 