← 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
M bench/exo_bench.pyM src/exo/main.pyM src/exo/worker/engines/mlx/utils_mlx.pyM src/exo/worker/runner/bootstrap.pyM src/exo/worker/runner/runner.py
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 →