← back to Exo
add a server for distributed testing in /tests until we work out a stable solution. (#1098)
56af61fac97d4e044b19bd6469524eb221fc343a · 2026-01-08 12:50:04 +0000 · Evan Quiney
## Motivation
Testing multiple devices simultaneously requires coordination, and we
don't necessarily want to run a full EXO to test single components. We
need a mid-scale integration testing framework for distributed tests.
## Changes
Add a simple python server + bash query that runs Jaccl and Ring tests
without constructing a worker/master/networking. The query relies on all
devices being accessible over tailscale, currently.
## Test Plan
Manually tested RDMA + Ring inference on 2 nodes.
Files touched
M src/exo/worker/runner/bootstrap.pyM src/exo/worker/runner/runner.pyA tests/headless_runner.pyA tests/start_distributed_test.sh
Diff
commit 56af61fac97d4e044b19bd6469524eb221fc343a
Author: Evan Quiney <evanev7@gmail.com>
Date: Thu Jan 8 12:50:04 2026 +0000
add a server for distributed testing in /tests until we work out a stable solution. (#1098)
## Motivation
Testing multiple devices simultaneously requires coordination, and we
don't necessarily want to run a full EXO to test single components. We
need a mid-scale integration testing framework for distributed tests.
## Changes
Add a simple python server + bash query that runs Jaccl and Ring tests
without constructing a worker/master/networking. The query relies on all
devices being accessible over tailscale, currently.
## Test Plan
Manually tested RDMA + Ring inference on 2 nodes.
---
src/exo/worker/runner/bootstrap.py | 16 +-
src/exo/worker/runner/runner.py | 311 +++++++++++++++++--------------------
tests/headless_runner.py | 246 +++++++++++++++++++++++++++++
tests/start_distributed_test.sh | 52 +++++++
4 files changed, 452 insertions(+), 173 deletions(-)
diff --git a/src/exo/worker/runner/bootstrap.py b/src/exo/worker/runner/bootstrap.py
index 4cd834cb..27c9993e 100644
--- a/src/exo/worker/runner/bootstrap.py
+++ b/src/exo/worker/runner/bootstrap.py
@@ -6,7 +6,7 @@ from exo.shared.types.events import Event, RunnerStatusUpdated
from exo.shared.types.tasks import Task
from exo.shared.types.worker.instances import BoundInstance, MlxJacclInstance
from exo.shared.types.worker.runners import RunnerFailed
-from exo.utils.channels import MpReceiver, MpSender
+from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
logger: "loguru.Logger" = loguru.logger
@@ -31,6 +31,8 @@ def entrypoint(
from exo.worker.runner.runner import main
main(bound_instance, event_sender, task_receiver)
+ except ClosedResourceError:
+ logger.warning("Runner communication closed unexpectedly")
except Exception as e:
logger.opt(exception=e).warning(
f"Runner {bound_instance.bound_runner_id} crashed with critical exception {e}"
@@ -42,8 +44,10 @@ def entrypoint(
)
)
finally:
- event_sender.close()
- task_receiver.close()
- event_sender.join()
- task_receiver.join()
- logger.info("bye from the runner")
+ try:
+ event_sender.close()
+ task_receiver.close()
+ finally:
+ event_sender.join()
+ task_receiver.join()
+ logger.info("bye from the runner")
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index a3732065..f39e7c10 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -1,5 +1,7 @@
import time
+import mlx.core as mx
+
from exo.shared.types.api import ChatCompletionMessageText
from exo.shared.types.chunks import TokenChunk
from exo.shared.types.events import (
@@ -36,12 +38,11 @@ from exo.shared.types.worker.runners import (
RunnerStatus,
RunnerWarmingUp,
)
-from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
+from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.mlx.generator.generate import mlx_generate, warmup_inference
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
load_mlx_items,
- mlx_cleanup,
mlx_force_oom,
)
from exo.worker.runner.bootstrap import logger
@@ -57,182 +58,158 @@ def main(
bound_instance.bound_runner_id,
bound_instance.bound_shard,
)
- try:
- logger.info("hello from the runner")
- if getattr(shard_metadata, "immediate_exception", False):
- raise Exception("Fake exception - runner failed to spin up.")
- if timeout := getattr(shard_metadata, "should_timeout", 0):
- time.sleep(timeout)
-
- setup_start_time = time.time()
-
- model = None
- tokenizer = None
- sampler = None
- group = None
-
- current_status: RunnerStatus = RunnerIdle()
- logger.info("runner created")
- event_sender.send(
- RunnerStatusUpdated(runner_id=runner_id, runner_status=current_status)
- )
- with task_receiver as tasks:
- for task in tasks:
- event_sender.send(
- TaskStatusUpdated(
- task_id=task.task_id, task_status=TaskStatus.Running
- )
- )
- event_sender.send(TaskAcknowledged(task_id=task.task_id))
- match task:
- case ConnectToGroup() if isinstance(
- current_status, (RunnerIdle, RunnerFailed)
- ):
- logger.info("runner connecting")
- current_status = RunnerConnecting()
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
- )
- )
- group = initialize_mlx(bound_instance)
-
- logger.info("runner connected")
- current_status = RunnerConnected()
-
- # we load the model if it's connected with a group, or idle without a group. we should never tell a model to connect if it doesn't need to
- case LoadModel() if (
- isinstance(current_status, RunnerConnected)
- and group is not None
- ) or (isinstance(current_status, RunnerIdle) and group is None):
- current_status = RunnerLoading()
- logger.info("runner loading")
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
- )
+ logger.info("hello from the runner")
+ if getattr(shard_metadata, "immediate_exception", False):
+ raise Exception("Fake exception - runner failed to spin up.")
+ if timeout := getattr(shard_metadata, "should_timeout", 0):
+ time.sleep(timeout)
+
+ setup_start_time = time.time()
+
+ model = None
+ tokenizer = None
+ sampler = None
+ group = None
+
+ current_status: RunnerStatus = RunnerIdle()
+ logger.info("runner created")
+ event_sender.send(
+ RunnerStatusUpdated(runner_id=runner_id, runner_status=current_status)
+ )
+ with task_receiver as tasks:
+ for task in tasks:
+ event_sender.send(
+ TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Running)
+ )
+ event_sender.send(TaskAcknowledged(task_id=task.task_id))
+ match task:
+ case ConnectToGroup() if isinstance(
+ current_status, (RunnerIdle, RunnerFailed)
+ ):
+ logger.info("runner connecting")
+ current_status = RunnerConnecting()
+ event_sender.send(
+ RunnerStatusUpdated(
+ runner_id=runner_id, runner_status=current_status
)
-
- model, tokenizer, sampler = load_mlx_items(
- bound_instance, group
+ )
+ group = initialize_mlx(bound_instance)
+
+ logger.info("runner connected")
+ current_status = RunnerConnected()
+
+ # we load the model if it's connected with a group, or idle without a group. we should never tell a model to connect if it doesn't need to
+ case LoadModel() if (
+ isinstance(current_status, RunnerConnected) and group is not None
+ ) or (isinstance(current_status, RunnerIdle) and group is None):
+ current_status = RunnerLoading()
+ logger.info("runner loading")
+ event_sender.send(
+ RunnerStatusUpdated(
+ runner_id=runner_id, runner_status=current_status
)
+ )
- current_status = RunnerLoaded()
- logger.info("runner loaded")
- case StartWarmup() if isinstance(current_status, RunnerLoaded):
- assert model
- assert tokenizer
- assert sampler
- current_status = RunnerWarmingUp()
- logger.info("runner warming up")
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
- )
+ model, tokenizer, sampler = load_mlx_items(bound_instance, group)
+
+ current_status = RunnerLoaded()
+ logger.info("runner loaded")
+ case StartWarmup() if isinstance(current_status, RunnerLoaded):
+ assert model
+ assert tokenizer
+ assert sampler
+ current_status = RunnerWarmingUp()
+ logger.info("runner warming up")
+ event_sender.send(
+ RunnerStatusUpdated(
+ runner_id=runner_id, runner_status=current_status
)
+ )
- logger.info(f"warming up inference for instance: {instance}")
- toks = warmup_inference(
- model=model,
- tokenizer=tokenizer,
- sampler=sampler,
- # kv_prefix_cache=kv_prefix_cache, # supply for warmup-time prefix caching
- )
- logger.info(f"warmed up by generating {toks} tokens")
- logger.info(
- f"runner initialized in {time.time() - setup_start_time} seconds"
- )
- current_status = RunnerReady()
- logger.info("runner ready")
- case ChatCompletion(
- task_params=task_params, command_id=command_id
- ) if isinstance(current_status, RunnerReady):
- assert model
- assert tokenizer
- assert sampler
- logger.info(f"received chat request: {str(task)[:500]}")
- current_status = RunnerRunning()
- logger.info("runner running")
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
- )
+ logger.info(f"warming up inference for instance: {instance}")
+ toks = warmup_inference(
+ model=model,
+ tokenizer=tokenizer,
+ sampler=sampler,
+ # kv_prefix_cache=kv_prefix_cache, # supply for warmup-time prefix caching
+ )
+ logger.info(f"warmed up by generating {toks} tokens")
+ logger.info(
+ f"runner initialized in {time.time() - setup_start_time} seconds"
+ )
+ current_status = RunnerReady()
+ logger.info("runner ready")
+ case ChatCompletion(task_params=task_params, command_id=command_id) if (
+ isinstance(current_status, RunnerReady)
+ ):
+ assert model
+ assert tokenizer
+ assert sampler
+ logger.info(f"received chat request: {str(task)[:500]}")
+ current_status = RunnerRunning()
+ logger.info("runner running")
+ event_sender.send(
+ RunnerStatusUpdated(
+ runner_id=runner_id, runner_status=current_status
)
- assert task_params.messages[0].content is not None
- _check_for_debug_prompts(task_params.messages[0].content)
-
- # Generate responses using the actual MLX generation
- for response in mlx_generate(
- model=model,
- tokenizer=tokenizer,
- sampler=sampler,
- task=task_params,
- ):
- match response:
- case GenerationResponse():
- if shard_metadata.device_rank == 0:
- event_sender.send(
- ChunkGenerated(
- command_id=command_id,
- chunk=TokenChunk(
- idx=response.token,
- model=shard_metadata.model_meta.model_id,
- text=response.text,
- token_id=response.token,
- finish_reason=response.finish_reason,
- stats=response.stats,
- ),
- )
+ )
+ assert task_params.messages[0].content is not None
+ _check_for_debug_prompts(task_params.messages[0].content)
+
+ # Generate responses using the actual MLX generation
+ for response in mlx_generate(
+ model=model,
+ tokenizer=tokenizer,
+ sampler=sampler,
+ task=task_params,
+ ):
+ match response:
+ case GenerationResponse():
+ if shard_metadata.device_rank == 0:
+ event_sender.send(
+ ChunkGenerated(
+ command_id=command_id,
+ chunk=TokenChunk(
+ idx=response.token,
+ model=shard_metadata.model_meta.model_id,
+ text=response.text,
+ token_id=response.token,
+ finish_reason=response.finish_reason,
+ stats=response.stats,
+ ),
)
- # case TokenizedResponse():
- # TODO: something here ig
-
- current_status = RunnerReady()
- logger.info("runner ready")
- case Shutdown():
- current_status = RunnerShuttingDown()
- logger.info("runner shutting down")
- mlx_cleanup(model, tokenizer, group)
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
- )
+ )
+ # case TokenizedResponse():
+ # TODO: something here ig
+
+ current_status = RunnerReady()
+ logger.info("runner ready")
+ case Shutdown():
+ current_status = RunnerShuttingDown()
+ logger.info("runner shutting down")
+ event_sender.send(
+ RunnerStatusUpdated(
+ runner_id=runner_id, runner_status=current_status
)
- current_status = RunnerShutdown()
- case _:
- raise ValueError(
- f"Received {task.__class__.__name__} outside of state machine in {current_status=}"
- )
- event_sender.send(
- TaskStatusUpdated(
- task_id=task.task_id, task_status=TaskStatus.Complete
)
- )
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id, runner_status=current_status
+ current_status = RunnerShutdown()
+ case _:
+ raise ValueError(
+ f"Received {task.__class__.__name__} outside of state machine in {current_status=}"
)
- )
- if isinstance(current_status, RunnerShutdown):
- break
- except ClosedResourceError:
- logger.warning("runner communication closed unexpectedly")
- except Exception as e:
- logger.opt(exception=e).warning(
- f"Runner {runner_id} crashed with critical exception {e}"
- )
- event_sender.send(
- RunnerStatusUpdated(
- runner_id=runner_id,
- runner_status=RunnerFailed(error_message=str(e)),
+ event_sender.send(
+ TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete)
+ )
+ event_sender.send(
+ RunnerStatusUpdated(runner_id=runner_id, runner_status=current_status)
)
- )
- finally:
- event_sender.close()
- task_receiver.close()
- event_sender.join()
- task_receiver.join()
- logger.info("bye from the runner")
+ if isinstance(current_status, RunnerShutdown):
+ del model, tokenizer, group, sampler
+ mx.clear_cache()
+ import gc
+
+ gc.collect()
+ break
EXO_RUNNER_MUST_FAIL = "EXO RUNNER MUST FAIL"
diff --git a/tests/headless_runner.py b/tests/headless_runner.py
new file mode 100644
index 00000000..6728d143
--- /dev/null
+++ b/tests/headless_runner.py
@@ -0,0 +1,246 @@
+import multiprocessing as mp
+import socket
+import time
+import typing
+
+import anyio
+from fastapi import FastAPI
+from fastapi.responses import StreamingResponse
+from hypercorn import Config
+from hypercorn.asyncio import serve # pyright: ignore[reportUnknownVariableType]
+from loguru import logger
+from pydantic import BaseModel
+
+from exo.shared.logging import InterceptLogger, logger_setup
+from exo.shared.models.model_cards import MODEL_CARDS, ModelId
+from exo.shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
+from exo.shared.types.commands import CommandId
+from exo.shared.types.common import Host, NodeId
+from exo.shared.types.events import Event
+from exo.shared.types.tasks import (
+ ChatCompletion,
+ ConnectToGroup,
+ LoadModel,
+ Shutdown,
+ StartWarmup,
+ Task,
+)
+from exo.shared.types.worker.instances import (
+ BoundInstance,
+ Instance,
+ InstanceId,
+ MlxJacclInstance,
+ MlxRingInstance,
+)
+from exo.shared.types.worker.runners import RunnerId, ShardAssignments
+from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata
+from exo.utils.channels import MpReceiver, MpSender, mp_channel
+from exo.worker.download.impl_shard_downloader import (
+ build_full_shard,
+ exo_shard_downloader,
+)
+from exo.worker.runner.bootstrap import entrypoint
+
+
+class Tests(BaseModel):
+ # list[hostname, ip addr]
+ devs: list[list[str]]
+ model_id: str
+ kind: typing.Literal["init", "warmup", "inference"]
+
+
+hn = socket.gethostname()
+mp.set_start_method("spawn", force=True)
+logger_setup(None)
+
+
+async def main():
+ logger.info("starting cool server majig")
+ logger.info(hn)
+ await assert_downloads()
+ cfg = Config()
+ cfg.bind = "0.0.0.0:52415"
+ # nb: shared.logging needs updating if any of this changes
+ cfg.accesslog = "-"
+ cfg.errorlog = "-"
+ cfg.logger_class = InterceptLogger
+ app = FastAPI()
+ app.post("/ring")(ring_backend)
+ app.post("/jaccl")(jaccl_backend)
+ shutdown = anyio.Event()
+ await serve(
+ app, # type: ignore
+ cfg,
+ shutdown_trigger=lambda: shutdown.wait(),
+ )
+ await anyio.sleep_forever()
+ # gracefully shutdown the api
+ shutdown.set()
+
+
+async def assert_downloads():
+ sd = exo_shard_downloader()
+ # await sd.ensure_shard(await build_full_shard(MODEL_CARDS["qwen3-0.6b"].model_id))
+ await sd.ensure_shard(await build_full_shard(MODEL_CARDS["llama-3.2-1b"].model_id))
+
+
+async def ring_backend(test: Tests):
+ iid = InstanceId(str(hash(str(test.devs))))
+ return await execute_test(test, ring_instance(test, iid))
+
+
+def ring_instance(test: Tests, iid: InstanceId) -> Instance:
+ global hn
+ hbn = [Host(ip="i dont care", port=52416) for _ in test.devs]
+ world_size = len(test.devs)
+ for i in range(world_size):
+ if hn.startswith(test.devs[i][0]):
+ hn = test.devs[i][0]
+ if i - 1 >= 0:
+ hbn[i - 1] = Host(ip=test.devs[i - 1][1], port=52416)
+ if i + 1 < len(test.devs):
+ hbn[i + 1] = Host(ip=test.devs[i + 1][1], port=52416)
+ hbn[i] = Host(ip="0.0.0.0", port=52416)
+ break
+
+ meta = MODEL_CARDS[test.model_id].metadata
+ instance = MlxRingInstance(
+ instance_id=iid,
+ ephemeral_port=52416,
+ hosts_by_node={NodeId(hn): hbn},
+ shard_assignments=ShardAssignments(
+ model_id=ModelId(test.model_id),
+ node_to_runner={NodeId(host[0]): RunnerId(host[0]) for host in test.devs},
+ runner_to_shard={
+ RunnerId(test.devs[i][0]): PipelineShardMetadata(
+ model_meta=meta,
+ device_rank=i,
+ world_size=world_size,
+ start_layer=(meta.n_layers // world_size) * i,
+ end_layer=min(
+ meta.n_layers, (meta.n_layers // world_size) * (i + 1)
+ ),
+ n_layers=min(meta.n_layers, (meta.n_layers // world_size) * (i + 1))
+ - (meta.n_layers // world_size) * i,
+ )
+ for i in range(world_size)
+ },
+ ),
+ )
+
+ return instance
+
+
+async def execute_test(test: Tests, instance: Instance):
+ world_size = len(test.devs)
+ iid = InstanceId(str(hash(str(test.devs))))
+ _handle, recv, send = new_runner(instance)
+ if world_size > 1:
+ send.send(ConnectToGroup(instance_id=iid))
+ send.send(LoadModel(instance_id=iid))
+
+ match test.kind:
+ case "init":
+ pass
+ case "warmup":
+ send.send(StartWarmup(instance_id=iid))
+ case "inference":
+ send.send(StartWarmup(instance_id=iid))
+ send.send(
+ ChatCompletion(
+ task_params=ChatCompletionTaskParams(
+ model=test.model_id,
+ messages=[
+ ChatCompletionMessage(
+ role="system", content="You are a helpful assistant"
+ ),
+ ChatCompletionMessage(
+ role="user", content="What is the capital of France?"
+ ),
+ ],
+ ),
+ command_id=CommandId("yo"),
+ instance_id=iid,
+ )
+ )
+
+ send.send(Shutdown(runner_id=RunnerId(hn), instance_id=iid))
+
+ async def map_recv():
+ with recv:
+ try:
+ async for item in recv:
+ yield item.model_dump_json() + "\n"
+ except anyio.ClosedResourceError:
+ pass
+
+ ret = StreamingResponse(map_recv())
+ ret._pls_dont_gc = _handle # type: ignore
+ return ret
+
+
+async def jaccl_backend(test: Tests):
+ iid = InstanceId(str(hash(str(test.devs))))
+ return await execute_test(test, jaccl_instance(test, iid))
+
+
+def jaccl_instance(test: Tests, iid: InstanceId):
+ global hn
+ meta = MODEL_CARDS[test.model_id].metadata
+ world_size = len(test.devs)
+ for name, _ in test.devs:
+ if hn.startswith(name):
+ hn = name
+ break
+
+ return MlxJacclInstance(
+ instance_id=iid,
+ ibv_devices=[[None, "rdma_en3"], ["rdma_en3", None]],
+ # rank 0 is always coordinator
+ jaccl_coordinators={
+ NodeId(host[0]): test.devs[0][1] + ":52416" for host in test.devs
+ },
+ shard_assignments=ShardAssignments(
+ model_id=ModelId(test.model_id),
+ node_to_runner={NodeId(host[0]): RunnerId(host[0]) for host in test.devs},
+ runner_to_shard={
+ RunnerId(test.devs[i][0]): TensorShardMetadata(
+ model_meta=meta,
+ device_rank=i,
+ world_size=world_size,
+ start_layer=meta.n_layers,
+ end_layer=meta.n_layers,
+ n_layers=meta.n_layers,
+ )
+ for i in range(world_size)
+ },
+ ),
+ )
+
+
+def new_runner(
+ instance: Instance,
+) -> tuple[mp.Process, MpReceiver[Event], MpSender[Task]]:
+ bound_instance = BoundInstance(
+ instance=instance, bound_runner_id=RunnerId(hn), bound_node_id=NodeId(hn)
+ )
+ ev_send, ev_recv = mp_channel[Event]()
+ task_send, task_recv = mp_channel[Task]()
+
+ runner_process = mp.Process(
+ target=entrypoint,
+ args=(
+ bound_instance,
+ ev_send,
+ task_recv,
+ logger,
+ ),
+ )
+ runner_process._pls_dont_gc = (ev_send, task_recv) # type: ignore
+ runner_process.start()
+ time.sleep(0.1)
+ return (runner_process, ev_recv, task_send)
+
+
+if __name__ == "__main__":
+ anyio.run(main)
diff --git a/tests/start_distributed_test.sh b/tests/start_distributed_test.sh
new file mode 100755
index 00000000..33781e04
--- /dev/null
+++ b/tests/start_distributed_test.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+
+set -euo pipefail
+
+query() {
+ tailscale status | awk -v find="$1" '$2 == find { print $1 }'
+}
+
+if [[ $# -lt 2 ]]; then
+ echo "USAGE: $0 <test kind> [host1] [host2] ..."
+ exit 1
+fi
+
+
+kind=$1
+shift
+
+test_kinds="ring jaccl"
+
+if ! echo "$test_kinds" | grep -q "$kind"; then
+ printf "%s is not a known test kind.\nCurrent test kinds are %s" "$kind" "$test_kinds"
+ exit 1
+fi
+
+hostnames=("$@")
+weaved=()
+ips=()
+for name in "${hostnames[@]}"; do
+ ip=$(query "$name")
+ ips+=("$ip")
+ weaved+=("$name" "$ip")
+done
+
+devs_raw=$(printf "[\"%s\", \"%s\"], " "${weaved[@]}")
+devs="[${devs_raw%, }]"
+
+for i in "${!ips[@]}"; do
+ {
+ req="{
+ \"model_id\": \"llama-3.2-1b\",
+ \"devs\": ${devs},
+ \"kind\": \"inference\"
+ }"
+ echo "req $req"
+ curl -sN \
+ -X POST "http://${ips[$i]}:52415/${kind}" \
+ -H "Content-Type: application/json" -d "$req" \
+ 2>&1 | sed "s/^/\n${hostnames[$i]}@${ips[$i]}: /" || echo "curl to ${hostnames[$i]} failed"
+ } &
+done
+
+wait
← f76d543d We shouldn't fail on an HTTPException in the tier-2 discover
·
back to Exo
·
fmt: add swift formatting 55463a98 →