[object Object]

← back to Exo

Handle model timeouts (#1177)

5c8a2379403064f3659c8bd5f705c7efac6001d3 · 2026-01-16 20:25:12 +0000 · rltakashige

- Add eval with a timeout.
- Add fast synch flag

## Motivation

Because of the experimental FAST SYNCH flag, some models may not work.
This PR catches when this occurs and allows users to specify a run
without fast synch

## Changes

- Adds a flag to enable or disable fast synch (--fast-synch and
--no-fast-synch)
- Adds a heuristic timeout
- Reduces exo_bench default timeout to 10 minutes.

## Why It Works

Heuristic timeout assumes normal loading times on Mac devices (60 +
model size in gb / 5: e.g. DeepSeek takes up to 120 seconds to load on
tensor parallel, and timeout is set to 60 + 120 = 180s.

We could raise this value if necessary.

## Test Plan

### Manual Testing
Catches that GPT OSS fails to load in Tensor RDMA
Can launch with --no-fast-synch flag to launch GPT OSS.

**GPT OSS 20B**
TP with fast synch
<img width="3064" height="456" alt="image"
src="https://github.com/user-attachments/assets/f6e25cd8-8621-4e99-99fe-292ee05c4035"
/>

TP without fast synch
<img width="3098" height="496" alt="image"
src="https://github.com/user-attachments/assets/d36453d9-6686-4cfe-aa7c-a7d458369d4d"
/>
[Note: the performance is really not great as fast synch is off]

(As a sanity check)
PP with fast synch
<img width="3124" height="496" alt="image"
src="https://github.com/user-attachments/assets/e97d4547-c6fa-483d-badb-4b371b900b4c"
/>

PP without fast synch
<img width="3078" height="508" alt="image"
src="https://github.com/user-attachments/assets/b2e20dfd-4b0e-4295-8a92-417dfe745c28"
/>

PP without RDMA
<img width="3070" height="498" alt="image"
src="https://github.com/user-attachments/assets/a8509d68-0aef-4cda-bca5-a67d39a0801e"
/>

TP without RDMA
<img width="3068" height="496" alt="image"
src="https://github.com/user-attachments/assets/b5691429-89f4-4369-bcf2-8fde2ad7154a"
/>

Files touched

Diff

commit 5c8a2379403064f3659c8bd5f705c7efac6001d3
Author: rltakashige <rl.takashige@gmail.com>
Date:   Fri Jan 16 20:25:12 2026 +0000

    Handle model timeouts (#1177)
    
    - Add eval with a timeout.
    - Add fast synch flag
    
    ## Motivation
    
    Because of the experimental FAST SYNCH flag, some models may not work.
    This PR catches when this occurs and allows users to specify a run
    without fast synch
    
    ## Changes
    
    - Adds a flag to enable or disable fast synch (--fast-synch and
    --no-fast-synch)
    - Adds a heuristic timeout
    - Reduces exo_bench default timeout to 10 minutes.
    
    ## Why It Works
    
    Heuristic timeout assumes normal loading times on Mac devices (60 +
    model size in gb / 5: e.g. DeepSeek takes up to 120 seconds to load on
    tensor parallel, and timeout is set to 60 + 120 = 180s.
    
    We could raise this value if necessary.
    
    ## Test Plan
    
    ### Manual Testing
    Catches that GPT OSS fails to load in Tensor RDMA
    Can launch with --no-fast-synch flag to launch GPT OSS.
    
    **GPT OSS 20B**
    TP with fast synch
    <img width="3064" height="456" alt="image"
    src="https://github.com/user-attachments/assets/f6e25cd8-8621-4e99-99fe-292ee05c4035"
    />
    
    TP without fast synch
    <img width="3098" height="496" alt="image"
    src="https://github.com/user-attachments/assets/d36453d9-6686-4cfe-aa7c-a7d458369d4d"
    />
    [Note: the performance is really not great as fast synch is off]
    
    (As a sanity check)
    PP with fast synch
    <img width="3124" height="496" alt="image"
    src="https://github.com/user-attachments/assets/e97d4547-c6fa-483d-badb-4b371b900b4c"
    />
    
    PP without fast synch
    <img width="3078" height="508" alt="image"
    src="https://github.com/user-attachments/assets/b2e20dfd-4b0e-4295-8a92-417dfe745c28"
    />
    
    PP without RDMA
    <img width="3070" height="498" alt="image"
    src="https://github.com/user-attachments/assets/a8509d68-0aef-4cda-bca5-a67d39a0801e"
    />
    
    TP without RDMA
    <img width="3068" height="496" alt="image"
    src="https://github.com/user-attachments/assets/b5691429-89f4-4369-bcf2-8fde2ad7154a"
    />
---
 bench/exo_bench.py                      | 12 +++++--
 src/exo/main.py                         | 23 +++++++++++++
 src/exo/worker/engines/mlx/utils_mlx.py | 60 +++++++++++++++++++++++++++++++--
 src/exo/worker/runner/bootstrap.py      | 14 ++++++--
 src/exo/worker/runner/runner.py         | 15 ++++++++-
 5 files changed, 114 insertions(+), 10 deletions(-)

diff --git a/bench/exo_bench.py b/bench/exo_bench.py
index 6fbed737..7366dbc0 100644
--- a/bench/exo_bench.py
+++ b/bench/exo_bench.py
@@ -27,7 +27,7 @@ class ExoHttpError(RuntimeError):
 
 
 class ExoClient:
-    def __init__(self, host: str, port: int, timeout_s: float = 2400.0):
+    def __init__(self, host: str, port: int, timeout_s: float = 600.0):
         self.host = host
         self.port = port
         self.timeout_s = timeout_s
@@ -324,6 +324,12 @@ def main() -> int:
         default=4,
         help="Only consider placements using <= this many nodes.",
     )
+    ap.add_argument(
+        "--min-nodes",
+        type=int,
+        default=1,
+        help="Only consider placements using >= this many nodes.",
+    )
     ap.add_argument(
         "--instance-meta", choices=["ring", "jaccl", "both"], default="both"
     )
@@ -345,7 +351,7 @@ def main() -> int:
         help="Warmup runs per placement (uses first pp/tg).",
     )
     ap.add_argument(
-        "--timeout", type=float, default=2400.0, help="HTTP timeout (seconds)."
+        "--timeout", type=float, default=600.0, help="HTTP timeout (seconds)."
     )
     ap.add_argument(
         "--json-out",
@@ -424,7 +430,7 @@ def main() -> int:
         ):
             continue
 
-        if 0 < n <= args.max_nodes:
+        if args.min_nodes <= n <= args.max_nodes:
             selected.append(p)
 
     if not selected:
diff --git a/src/exo/main.py b/src/exo/main.py
index 85bc095b..6425595d 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -205,6 +205,14 @@ def main():
     logger.info("Starting EXO")
     logger.info(f"EXO_LIBP2P_NAMESPACE: {os.getenv('EXO_LIBP2P_NAMESPACE')}")
 
+    # Set FAST_SYNCH override env var for runner subprocesses
+    if args.fast_synch is True:
+        os.environ["EXO_FAST_SYNCH"] = "on"
+        logger.info("FAST_SYNCH forced ON")
+    elif args.fast_synch is False:
+        os.environ["EXO_FAST_SYNCH"] = "off"
+        logger.info("FAST_SYNCH forced OFF")
+
     node = anyio.run(Node.create, args)
     anyio.run(node.run)
     logger.info("EXO Shutdown complete")
@@ -218,6 +226,7 @@ class Args(CamelCaseModel):
     api_port: PositiveInt = 52415
     tb_only: bool = False
     no_worker: bool = False
+    fast_synch: bool | None = None  # None = auto, True = force on, False = force off
 
     @classmethod
     def parse(cls) -> Self:
@@ -259,6 +268,20 @@ class Args(CamelCaseModel):
             "--no-worker",
             action="store_true",
         )
+        fast_synch_group = parser.add_mutually_exclusive_group()
+        fast_synch_group.add_argument(
+            "--fast-synch",
+            action="store_true",
+            dest="fast_synch",
+            default=None,
+            help="Force MLX FAST_SYNCH on (for JACCL backend)",
+        )
+        fast_synch_group.add_argument(
+            "--no-fast-synch",
+            action="store_false",
+            dest="fast_synch",
+            help="Force MLX FAST_SYNCH off",
+        )
 
         args = parser.parse_args()
         return cls(**vars(args))  # pyright: ignore[reportAny] - We are intentionally validating here, we can't do it statically
diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py
index 81fc68aa..f069ed01 100644
--- a/src/exo/worker/engines/mlx/utils_mlx.py
+++ b/src/exo/worker/engines/mlx/utils_mlx.py
@@ -2,7 +2,9 @@ import json
 import os
 import resource
 import sys
+import threading
 import time
+from collections.abc import Callable
 from pathlib import Path
 from typing import Any, cast
 
@@ -82,6 +84,45 @@ def get_weights_size(model_shard_meta: ShardMetadata) -> Memory:
     )
 
 
+class ModelLoadingTimeoutError(Exception):
+    pass
+
+
+TimeoutCallback = Callable[[], None]
+
+
+def eval_with_timeout(
+    mlx_item: Any,  # pyright: ignore[reportAny]
+    timeout_seconds: float = 60.0,
+    on_timeout: TimeoutCallback | None = None,
+) -> None:
+    """Evaluate MLX item with a hard timeout.
+
+    If on_timeout callback is provided, it will be called before terminating
+    the process. This allows the runner to send a failure event before exit.
+    """
+    completed = threading.Event()
+
+    def watchdog() -> None:
+        if not completed.wait(timeout=timeout_seconds):
+            logger.error(
+                f"mlx_item evaluation timed out after {timeout_seconds:.0f}s. "
+                "This may indicate an issue with FAST_SYNCH and tensor parallel sharding. "
+                "Terminating process."
+            )
+            if on_timeout is not None:
+                on_timeout()
+            os._exit(1)
+
+    watchdog_thread = threading.Thread(target=watchdog, daemon=True)
+    watchdog_thread.start()
+
+    try:
+        mx.eval(mlx_item)  # pyright: ignore[reportAny]
+    finally:
+        completed.set()
+
+
 def mx_barrier(group: Group | None = None):
     mx.eval(
         mx.distributed.all_sum(
@@ -188,7 +229,9 @@ def initialize_mlx(
 
 
 def load_mlx_items(
-    bound_instance: BoundInstance, group: Group | None
+    bound_instance: BoundInstance,
+    group: Group | None,
+    on_timeout: TimeoutCallback | None = None,
 ) -> tuple[Model, TokenizerWrapper]:
     if group is None:
         logger.info(f"Single device used for {bound_instance.instance}")
@@ -202,7 +245,9 @@ def load_mlx_items(
     else:
         logger.info("Starting distributed init")
         start_time = time.perf_counter()
-        model, tokenizer = shard_and_load(bound_instance.bound_shard, group=group)
+        model, tokenizer = shard_and_load(
+            bound_instance.bound_shard, group=group, on_timeout=on_timeout
+        )
         end_time = time.perf_counter()
         logger.info(
             f"Time taken to shard and load model: {(end_time - start_time):.2f}s"
@@ -216,6 +261,7 @@ def load_mlx_items(
 def shard_and_load(
     shard_metadata: ShardMetadata,
     group: Group,
+    on_timeout: TimeoutCallback | None = None,
 ) -> tuple[nn.Module, TokenizerWrapper]:
     model_path = build_model_path(shard_metadata.model_meta.model_id)
 
@@ -252,7 +298,15 @@ def shard_and_load(
             logger.info(f"loading model from {model_path} with pipeline parallelism")
             model = pipeline_auto_parallel(model, group, shard_metadata)
 
-    mx.eval(model.parameters())
+    # Estimate timeout based on model size
+    base_timeout = float(os.environ.get("EXO_MODEL_LOAD_TIMEOUT", "60"))
+    model_size_gb = get_weights_size(shard_metadata).in_bytes / (1024**3)
+    timeout_seconds = base_timeout + model_size_gb / 5
+    logger.info(
+        f"Evaluating model parameters with timeout of {timeout_seconds:.0f}s "
+        f"(model size: {model_size_gb:.1f}GB)"
+    )
+    eval_with_timeout(model.parameters(), timeout_seconds, on_timeout)
 
     # TODO: Do we need this?
     mx.eval(model)
diff --git a/src/exo/worker/runner/bootstrap.py b/src/exo/worker/runner/bootstrap.py
index 27c9993e..28e298cb 100644
--- a/src/exo/worker/runner/bootstrap.py
+++ b/src/exo/worker/runner/bootstrap.py
@@ -17,15 +17,23 @@ def entrypoint(
     task_receiver: MpReceiver[Task],
     _logger: "loguru.Logger",
 ) -> None:
-    if (
-        isinstance(bound_instance.instance, MlxJacclInstance)
-        and len(bound_instance.instance.ibv_devices) >= 2
+    fast_synch_override = os.environ.get("EXO_FAST_SYNCH")
+    if fast_synch_override == "on" or (
+        fast_synch_override != "off"
+        and (
+            isinstance(bound_instance.instance, MlxJacclInstance)
+            and len(bound_instance.instance.ibv_devices) >= 2
+        )
     ):
         os.environ["MLX_METAL_FAST_SYNCH"] = "1"
+    else:
+        os.environ["MLX_METAL_FAST_SYNCH"] = "0"
 
     global logger
     logger = _logger
 
+    logger.info(f"Fast synch flag: {os.environ['MLX_METAL_FAST_SYNCH']}")
+
     # Import main after setting global logger - this lets us just import logger from this module
     try:
         from exo.worker.runner.runner import main
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index e195b9b3..28f89d54 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -150,7 +150,20 @@ def main(
                         )
                     )
 
-                    model, tokenizer = load_mlx_items(bound_instance, group)
+                    def on_model_load_timeout() -> None:
+                        event_sender.send(
+                            RunnerStatusUpdated(
+                                runner_id=runner_id,
+                                runner_status=RunnerFailed(
+                                    error_message="Model loading timed out"
+                                ),
+                            )
+                        )
+                        time.sleep(0.5)
+
+                    model, tokenizer = load_mlx_items(
+                        bound_instance, group, on_timeout=on_model_load_timeout
+                    )
 
                     current_status = RunnerLoaded()
                     logger.info("runner loaded")

← 745343c7 Return error responses for Chat Completions (#1173)  ·  back to Exo  ·  Add pre-commit checks documentation to AGENTS.md (#1184) c5158bee →