Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
executor.py312 linesDownload Raw Back to serverless
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