Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
modal_spawn.py391 linesDownload Raw Back to serverless
1"""ModalSpawnExecutor — production Modal-backed serverless executor.2 3This is the v0-finished sibling of `ModalExecutor` (which remains a4skeleton per Wave 18 contract). The skeleton class stays unchanged to5preserve `test_skeleton_executors.py`'s pinned NotImplementedError6contract; this class is the working alternative for users who want7real Modal execution.8 9Design choices vs the skeleton's docstring:10 111. **User-provided `modal.Function` instead of internal app construction.**12   The skeleton showed a pattern where ModalExecutor builds its own13   `modal.App` and registers `run_replica` internally. That couples the14   executor to image/GPU/Volume choices the user actually wants to own.15   Instead, ModalSpawnExecutor takes a *pre-decorated* `modal.Function`16   from the caller — the user defines:17 18       @app.function(gpu="H100:4", image=my_image, volumes={"/vol": vol},19                     secrets=[modal.Secret.from_name("hf-token")],20                     timeout=4*3600)21       def run_replica(rendezvous_uri: str, world_size: int,22                       rank: int, **entrypoint_args):23           import os24           os.environ["REPLICA_RANK"] = str(rank)25           from composer_replication.diloco.serverless import (26               MockManager, ObjectStoreAllReduce,27           )28           store = ObjectStoreAllReduce(rendezvous_uri, rank=rank,29                                        world_size=world_size)30           manager = MockManager(store)31           # ... user's training loop with this manager ...32 33   then constructs:34 35       executor = ModalSpawnExecutor(modal_function=run_replica)36       handles = executor.launch_replicas(37           n_replicas=4,38           entrypoint=run_replica,        # ignored — function is bound39           entrypoint_args={"rendezvous_uri": "/vol/diloco/run42",40                            "world_size": 4},41       )42 432. **Rank as explicit kwarg, not env-var indirection.** Modal Functions44   start with a clean env, so the rank-via-env pattern that45   LocalProcessExecutor uses is fragile here (Modal would need46   container-level env injection per call, which `modal.Secret.from_dict`47   does but adds a round-trip per spawn). We pass rank as a kwarg to48   `.spawn(rank=i)` so it's plumbed through Modal's call args directly.49 503. **Handle metadata = `call_id`, no in-process state.** Unlike51   LocalProcessExecutor (which holds Process refs), this executor is52   stateless after launch — handles are reconstructed via53   `modal.FunctionCall.from_id(call_id)` for poll/cancel/collect.54   Lets the executor survive process restart mid-run.55 56References:57- modal-client 1.4.x docs on FunctionCall: https://modal.com/docs/reference/modal.FunctionCall58- ADR-005 (executor protocol design)59"""60from __future__ import annotations61 62import time63from typing import Any, Callable, Mapping64 65from composer_replication.diloco.serverless.executor import (66    ReplicaHandle,67    ServerlessExecutor,68)69 70 71class ModalSpawnExecutor:72    """Run replicas as parallel Modal Function spawns.73 74    Implements the `ServerlessExecutor` Protocol against Modal's75    `Function.spawn()` API. The user must provide a pre-decorated76    `modal.Function` (with `@app.function(...)` already applied) — see77    module docstring for the expected signature.78 79    Args:80        modal_function: a `modal.Function` registered against a `modal.App`.81            Must accept at minimum `rank: int` plus the kwargs in82            `entrypoint_args`. Image / GPU / Volume / Secret / timeout83            are pinned on the decorator and the executor won't override84            them.85        deploy: if True, calls `modal_function.app.deploy()` before86            spawning. Required when running outside a `modal run` context87            (e.g. from a regular Python script). Default False — assumes88            the user is inside a `modal run` block where the app is89            already live.90 91    Raises:92        RuntimeError: if `modal` client is not installed.93        TypeError: if `modal_function` is not a `modal.Function`.94    """95    backend_name = "modal_spawn"96    supports_inter_replica_network = False  # Modal containers are isolated by default97 98    def __init__(99        self,100        modal_function: Any,101        *,102        deploy: bool = False,103    ) -> None:104        try:105            import modal  # noqa: F401106        except ImportError as e:107            raise RuntimeError(108                "ModalSpawnExecutor requires the modal client. Install with "109                "`pip install modal` and configure with `modal token new`. "110                f"Got: {e!r}"111            )112 113        # Duck-type check — modal.Function objects expose .spawn / .remote /114        # ._app, which the user-supplied function will have if they used the115        # @app.function(...) decorator. We avoid `isinstance(_, modal.Function)`116        # to stay tolerant of modal-client minor-version changes that may117        # restructure the class.118        if not (hasattr(modal_function, "spawn") and hasattr(modal_function, "remote")):119            raise TypeError(120                f"modal_function must be a modal.Function (decorated via "121                f"`@app.function(...)`). Got {type(modal_function)!r} which "122                f"has no `.spawn()` method. "123                f"See ModalSpawnExecutor docstring for expected signature."124            )125 126        self.modal_function = modal_function127        self._deploy_requested = deploy128        self._deployed = False129        self._handles: dict[int, dict[str, Any]] = {}130 131    # -----------------------------------------------------------------132    # Lifecycle133    # -----------------------------------------------------------------134 135    def _maybe_deploy(self) -> None:136        if self._deploy_requested and not self._deployed:137            # `modal_function.app` exposes the underlying App. Calling138            # `.deploy()` registers it with Modal so spawn() works from139            # outside `modal run`.140            app = getattr(self.modal_function, "app", None)141            if app is None:142                raise RuntimeError(143                    "modal_function.app is None — can't deploy. The function "144                    "must have been decorated against a real modal.App."145                )146            app.deploy()147            self._deployed = True148 149    # -----------------------------------------------------------------150    # ServerlessExecutor Protocol151    # -----------------------------------------------------------------152 153    def launch_replicas(154        self,155        n_replicas: int,156        entrypoint: str | Callable[..., Any],157        entrypoint_args: Mapping[str, Any],158        *,159        gpu: str | None = None,160        timeout: int = 3600,161    ) -> list[ReplicaHandle]:162        """Spawn N parallel Modal Function calls.163 164        Note: `entrypoint` is **ignored** — the actual entrypoint is the165        `modal_function` passed to `__init__`. This keeps the executor166        Protocol-compatible while preserving the user's image/GPU167        decoration. `gpu` and `timeout` are similarly ignored (pinned168        on the function decorator).169        """170        del entrypoint, gpu, timeout  # pinned on the decorated function171 172        if n_replicas < 1:173            raise ValueError(f"n_replicas must be >= 1, got {n_replicas}")174 175        self._maybe_deploy()176 177        # Strip rank_env if present — we use explicit `rank` kwarg.178        spawn_kwargs = {k: v for k, v in entrypoint_args.items()179                        if k != "rank_env"}180 181        handles: list[ReplicaHandle] = []182        for rank in range(n_replicas):183            try:184                fcall = self.modal_function.spawn(rank=rank, **spawn_kwargs)185            except Exception as e:186                # Best-effort cancel any already-launched siblings187                for prior in handles:188                    try:189                        self.cancel(prior)190                    except Exception:191                        pass192                raise RuntimeError(193                    f"ModalSpawnExecutor.launch_replicas failed at rank={rank} "194                    f"of {n_replicas} (already-launched siblings cancelled). "195                    f"Underlying error: {e!r}"196                ) from e197 198            handle = ReplicaHandle(199                rank=rank,200                backend_name=self.backend_name,201                metadata={202                    "call_id": fcall.object_id,203                    "spawn_ts": time.time(),204                },205            )206            self._handles[rank] = {207                "fcall": fcall,208                "result": None,209            }210            handles.append(handle)211 212        return handles213 214    def poll(self, handle: ReplicaHandle) -> str:215        """Poll a Modal call's status.216 217        Modal's FunctionCall doesn't expose a non-blocking status getter218        directly (the API is `.get(timeout=...)`), so we poll by trying219        `.get(timeout=0)` and treating Timeout/Pending as "running".220 221        Returns one of: "pending" | "running" | "succeeded" | "failed" |222        "cancelled".223        """224        meta = self._handles.get(handle.rank)225        if meta is None:226            return "cancelled"227 228        # If we already collected this one, return cached result229        if meta["result"] is not None:230            return meta["result"]["status"]231 232        import modal233        from modal.exception import OutputExpiredError234 235        fcall = meta["fcall"]236        # Re-hydrate to get fresh state237        try:238            # `.get(timeout=0)` returns immediately if done; raises TimeoutError otherwise.239            result_value = fcall.get(timeout=0)240            meta["result"] = {241                "rank": handle.rank,242                "status": "succeeded",243                "exit_code": 0,244                "error": None,245                "result": result_value,246                "call_id": handle.metadata.get("call_id"),247            }248            return "succeeded"249        except TimeoutError:250            return "running"251        except OutputExpiredError as e:252            meta["result"] = {253                "rank": handle.rank,254                "status": "failed",255                "exit_code": 1,256                "error": f"OutputExpiredError: {e!r}",257                "result": None,258                "call_id": handle.metadata.get("call_id"),259            }260            return "failed"261        except Exception as e:262            # User-code exception bubbles up here as the original exception class263            meta["result"] = {264                "rank": handle.rank,265                "status": "failed",266                "exit_code": 1,267                "error": f"{type(e).__name__}: {e!r}",268                "result": None,269                "call_id": handle.metadata.get("call_id"),270            }271            return "failed"272 273    def stream_logs(self, handle: ReplicaHandle, *, n_lines: int = 200) -> str:274        """Read recent Modal logs for this call.275 276        Modal exposes per-FunctionCall logs via the dashboard URL. The277        client API doesn't expose log-streaming directly in 1.4.x, so we278        return a pointer to the dashboard URL plus any captured error279        from poll().280        """281        meta = self._handles.get(handle.rank)282        if meta is None:283            return f"<replica {handle.rank}: no metadata>"284 285        call_id = handle.metadata.get("call_id", "<unknown>")286        try:287            dashboard_url = meta["fcall"].get_dashboard_url()288        except Exception:289            dashboard_url = (290                f"https://modal.com/apps/<workspace>/<env>/calls/{call_id}"291            )292 293        if meta.get("result"):294            err = meta["result"].get("error") or "<no error>"295            return (296                f"[rank {handle.rank}] call_id={call_id}\n"297                f"  Dashboard: {dashboard_url}\n"298                f"  Result: {meta['result']['status']}\n"299                f"  Error: {err[-2000:] if err else '<none>'}"300            )301 302        return (303            f"[rank {handle.rank}] call_id={call_id} (still running)\n"304            f"  Dashboard: {dashboard_url}\n"305            f"  Logs not streamable via client API in modal-client 1.4.x; "306            f"use the dashboard URL or `modal app logs <app-id>`."307        )308 309    def cancel(self, handle: ReplicaHandle) -> None:310        """Best-effort cancel of a Modal call."""311        meta = self._handles.get(handle.rank)312        if meta is None:313            return314        try:315            meta["fcall"].cancel()316        except Exception:317            # Already terminated, network blip, etc. — best-effort.318            pass319 320    def collect(321        self,322        handles: list[ReplicaHandle],323        *,324        timeout: int | None = None,325    ) -> list[dict[str, Any]]:326        """Block until all replicas finish; return per-replica result dicts.327 328        Modal's `.get(timeout=...)` blocks until the call completes or329        the timeout elapses. We iterate handles, calling `.get()` with330        the remaining time budget, so the cumulative wall-clock is331        bounded by `timeout`.332        """333        deadline = time.time() + (timeout if timeout is not None else 86400)334        results: list[dict[str, Any]] = []335 336        for h in handles:337            meta = self._handles.get(h.rank)338            if meta is None:339                results.append({340                    "rank": h.rank,341                    "status": "cancelled",342                    "exit_code": None,343                    "error": "handle has no metadata (cancelled or unknown)",344                    "result": None,345                    "call_id": h.metadata.get("call_id"),346                })347                continue348 349            # Already collected by an earlier poll()350            if meta["result"] is not None:351                results.append(meta["result"])352                continue353 354            remaining = max(0.0, deadline - time.time())355            try:356                result_value = meta["fcall"].get(timeout=remaining)357                result_dict = {358                    "rank": h.rank,359                    "status": "succeeded",360                    "exit_code": 0,361                    "error": None,362                    "result": result_value,363                    "call_id": h.metadata.get("call_id"),364                }365            except TimeoutError as e:366                result_dict = {367                    "rank": h.rank,368                    "status": "running",369                    "exit_code": None,370                    "error": f"TimeoutError after deadline: {e!r}",371                    "result": None,372                    "call_id": h.metadata.get("call_id"),373                }374            except Exception as e:375                result_dict = {376                    "rank": h.rank,377                    "status": "failed",378                    "exit_code": 1,379                    "error": f"{type(e).__name__}: {e!r}",380                    "result": None,381                    "call_id": h.metadata.get("call_id"),382                }383 384            meta["result"] = result_dict385            results.append(result_dict)386 387        return results388 389 390__all__ = ["ModalSpawnExecutor"]391