Codeseys/composer-replication-framework
0
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 