← back to Exo
improved master error handling
b62f68474afea16a1ae8b0055d9662b1f53ca3c7 · 2025-11-11 18:04:40 +0000 · Evan Quiney
Co-authored-by: Ryuichi Leo Takashige <rl.takashige@gmail.com>
Files touched
M src/exo/engines/mlx/auto_parallel.pyM src/exo/engines/mlx/utils_mlx.pyM src/exo/main.pyM src/exo/master/main.pyM src/exo/shared/apply.pyM src/exo/shared/types/commands.pyM src/exo/shared/types/state.pyM src/exo/shared/types/worker/shards.pyM src/exo/utils/pydantic_ext.pyM src/exo/worker/main.pyM src/exo/worker/plan.pyM src/exo/worker/runner/bootstrap.pyM src/exo/worker/runner/generate.py
Diff
commit b62f68474afea16a1ae8b0055d9662b1f53ca3c7
Author: Evan Quiney <evanev7@gmail.com>
Date: Tue Nov 11 18:04:40 2025 +0000
improved master error handling
Co-authored-by: Ryuichi Leo Takashige <rl.takashige@gmail.com>
---
src/exo/engines/mlx/auto_parallel.py | 21 ++++++++++-----------
src/exo/engines/mlx/utils_mlx.py | 14 +++++++++++---
src/exo/main.py | 2 +-
src/exo/master/main.py | 11 ++++-------
src/exo/shared/apply.py | 2 +-
src/exo/shared/types/commands.py | 7 ++-----
src/exo/shared/types/state.py | 2 +-
src/exo/shared/types/worker/shards.py | 2 +-
src/exo/utils/pydantic_ext.py | 1 -
src/exo/worker/main.py | 35 +++++++++++++++++++++++++++++------
src/exo/worker/plan.py | 4 +++-
src/exo/worker/runner/bootstrap.py | 1 +
src/exo/worker/runner/generate.py | 22 ++++++++++++++--------
13 files changed, 78 insertions(+), 46 deletions(-)
diff --git a/src/exo/engines/mlx/auto_parallel.py b/src/exo/engines/mlx/auto_parallel.py
index 452b53c8..345454db 100644
--- a/src/exo/engines/mlx/auto_parallel.py
+++ b/src/exo/engines/mlx/auto_parallel.py
@@ -86,11 +86,10 @@ class PipelineLastLayer(CustomMlxLayer):
self.original_layer_signature = signature(self.original_layer.__call__)
@override
- def __call__(
- self, x: mx.array, *args: object, **kwargs: object
- ) -> mx.array:
-
- cache = self.original_layer_signature.bind_partial(x, *args, **kwargs).arguments.get("cache", None)
+ def __call__(self, x: mx.array, *args: object, **kwargs: object) -> mx.array:
+ cache = self.original_layer_signature.bind_partial(
+ x, *args, **kwargs
+ ).arguments.get("cache", None)
assert cache is None or isinstance(cache, (KVCache, RotatingKVCache))
@@ -101,12 +100,12 @@ class PipelineLastLayer(CustomMlxLayer):
output, (self.r + 1) % self.s, group=self.group
)
if (
- cache is not None
- and hasattr(cache, "keys")
- and getattr(cache, "keys", None) is not None
- ):
- # This change happened upstream - check out mlx github somewhere??
- cache.keys = mx.depends(cache.keys, output) # type: ignore[reportUnknownMemberType]
+ cache is not None
+ and hasattr(cache, "keys")
+ and getattr(cache, "keys", None) is not None
+ ):
+ # This change happened upstream - check out mlx github somewhere??
+ cache.keys = mx.depends(cache.keys, output) # type: ignore[reportUnknownMemberType]
output = mx.distributed.all_gather(output, group=self.group)[-output.shape[0] :]
return output
diff --git a/src/exo/engines/mlx/utils_mlx.py b/src/exo/engines/mlx/utils_mlx.py
index 9e92e723..c82fbee6 100644
--- a/src/exo/engines/mlx/utils_mlx.py
+++ b/src/exo/engines/mlx/utils_mlx.py
@@ -22,8 +22,8 @@ from exo.engines.mlx.auto_parallel import (
pipeline_auto_parallel,
tensor_auto_parallel,
)
-from exo.shared.types.memory import Memory
from exo.shared.types.common import Host
+from exo.shared.types.memory import Memory
from exo.shared.types.tasks import ChatCompletionTaskParams
from exo.shared.types.worker.instances import (
BoundInstance,
@@ -146,7 +146,12 @@ def initialize_mlx(
model_path = build_model_path(bound_instance.bound_shard().model_meta.model_id)
model, _ = load_model(model_path, strict=True)
# TODO: we should really make this opt-in, but Kimi requires trust_remote_code=True
- tokenizer = cast(TokenizerWrapper, load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}))
+ tokenizer = cast(
+ TokenizerWrapper,
+ load_tokenizer(
+ model_path, tokenizer_config_extra={"trust_remote_code": True}
+ ),
+ )
assert isinstance(tokenizer, TokenizerWrapper)
else:
@@ -170,7 +175,10 @@ def shard_and_load(
assert isinstance(model, nn.Module)
# TODO: we should really make this opt-in, but Kimi requires trust_remote_code=True
- tokenizer = cast(TokenizerWrapper, load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}))
+ tokenizer = cast(
+ TokenizerWrapper,
+ load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}),
+ )
logger.info(f"Group size: {group.size()}, group rank: {group.rank()}")
diff --git a/src/exo/main.py b/src/exo/main.py
index b3432135..110d44a6 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -16,8 +16,8 @@ from exo.routing.router import Router, get_node_id_keypair
from exo.shared.constants import EXO_LOG
from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
-from exo.shared.types.common import NodeId, SessionId
from exo.shared.types.commands import KillCommand
+from exo.shared.types.common import NodeId, SessionId
from exo.utils.channels import Receiver, channel
from exo.utils.pydantic_ext import CamelCaseModel
from exo.worker.download.impl_shard_downloader import exo_shard_downloader
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 7badeeca..7f481cb5 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -15,8 +15,8 @@ from exo.shared.types.commands import (
CreateInstance,
DeleteInstance,
ForwarderCommand,
+ KillCommand,
RequestEventLog,
- SpinUpInstance,
TaskFinished,
TestCommand,
)
@@ -104,7 +104,7 @@ class Master:
generated_events: list[Event] = []
command = forwarder_command.command
match command:
- case TestCommand():
+ case TestCommand() | KillCommand():
pass
case ChatCompletion():
instance_task_counts: dict[InstanceId, int] = {}
@@ -123,10 +123,9 @@ class Master:
)
if not instance_task_counts:
- logger.warning(
+ raise ValueError(
f"No instance found for model {command.request_params.model}"
)
- continue
available_instance_ids = sorted(
instance_task_counts.keys(),
@@ -181,8 +180,6 @@ class Master:
del self.command_task_mapping[
command.finished_command_id
]
- case SpinUpInstance():
- raise NotImplementedError
case RequestEventLog():
# We should just be able to send everything, since other buffers will ignore old messages
for i in range(command.since_idx, len(self._event_log)):
@@ -191,7 +188,7 @@ class Master:
)
for event in generated_events:
await self.event_sender.send(event)
- except Exception as e:
+ except ValueError as e:
logger.opt(exception=e).warning("Error in command processor")
async def _event_processor(self) -> None:
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index b30512af..16cc6adb 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -29,8 +29,8 @@ from exo.shared.types.profiling import NodePerformanceProfile, SystemPerformance
from exo.shared.types.state import State
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.topology import NodeInfo
-from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.downloads import DownloadProgress
+from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index 979d42bd..9ea2aa3f 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -16,9 +16,11 @@ class BaseCommand(TaggedModel):
class TestCommand(BaseCommand):
pass
+
class KillCommand(BaseCommand):
pass
+
class ChatCompletion(BaseCommand):
request_params: ChatCompletionTaskParams
@@ -29,10 +31,6 @@ class CreateInstance(BaseCommand):
instance_meta: InstanceMeta
-class SpinUpInstance(BaseCommand):
- instance_id: InstanceId
-
-
class DeleteInstance(BaseCommand):
instance_id: InstanceId
@@ -51,7 +49,6 @@ Command = (
| RequestEventLog
| ChatCompletion
| CreateInstance
- | SpinUpInstance
| DeleteInstance
| TaskFinished
)
diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py
index 3cd3a256..efdb5bcb 100644
--- a/src/exo/shared/types/state.py
+++ b/src/exo/shared/types/state.py
@@ -8,9 +8,9 @@ from exo.shared.topology import Topology, TopologySnapshot
from exo.shared.types.common import NodeId
from exo.shared.types.profiling import NodePerformanceProfile
from exo.shared.types.tasks import Task, TaskId
+from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
-from exo.shared.types.worker.downloads import DownloadProgress
from exo.utils.pydantic_ext import CamelCaseModel
diff --git a/src/exo/shared/types/worker/shards.py b/src/exo/shared/types/worker/shards.py
index 303adcc3..e8e86730 100644
--- a/src/exo/shared/types/worker/shards.py
+++ b/src/exo/shared/types/worker/shards.py
@@ -1,6 +1,6 @@
from enum import Enum
-from pydantic import Field, ConfigDict
+from pydantic import Field
from exo.shared.types.models import ModelMetadata
from exo.utils.pydantic_ext import TaggedModel
diff --git a/src/exo/utils/pydantic_ext.py b/src/exo/utils/pydantic_ext.py
index 5631723c..1c459b2d 100644
--- a/src/exo/utils/pydantic_ext.py
+++ b/src/exo/utils/pydantic_ext.py
@@ -40,4 +40,3 @@ class TaggedModel(CamelCaseModel):
def __str__(self) -> str:
return f"{self.__class__.__name__}({super().__str__()})"
-
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 6ccf6554..830bd7ce 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -17,14 +17,21 @@ from exo.shared.types.events import (
NodeDownloadProgress,
NodeMemoryMeasured,
NodePerformanceMeasured,
- TaskCreated, TaskStatusUpdated,
+ TaskCreated,
+ TaskStatusUpdated,
TopologyEdgeCreated,
TopologyEdgeDeleted,
)
from exo.shared.types.multiaddr import Multiaddr
from exo.shared.types.profiling import MemoryPerformanceProfile, NodePerformanceProfile
from exo.shared.types.state import State
-from exo.shared.types.tasks import CreateRunner, DownloadModel, Task, TaskStatus, Shutdown
+from exo.shared.types.tasks import (
+ CreateRunner,
+ DownloadModel,
+ Shutdown,
+ Task,
+ TaskStatus,
+)
from exo.shared.types.topology import Connection
from exo.shared.types.worker.downloads import (
DownloadCompleted,
@@ -180,7 +187,11 @@ class Worker:
match task:
case CreateRunner():
self._create_supervisor(task)
- await self.event_sender.send(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+ await self.event_sender.send(
+ TaskStatusUpdated(
+ task_id=task.task_id, task_status=TaskStatus.Complete
+ )
+ )
case DownloadModel(shard_metadata=shard):
if shard not in self.download_status:
progress = DownloadPending(
@@ -204,9 +215,17 @@ class Worker:
await self.event_sender.send(
NodeDownloadProgress(download_progress=progress)
)
- await self.event_sender.send(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+ await self.event_sender.send(
+ TaskStatusUpdated(
+ task_id=task.task_id, task_status=TaskStatus.Complete
+ )
+ )
else:
- self.event_sender.send_nowait(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Running))
+ self.event_sender.send_nowait(
+ TaskStatusUpdated(
+ task_id=task.task_id, task_status=TaskStatus.Running
+ )
+ )
await self._handle_shard_download_process(
task, initial_progress
)
@@ -326,7 +345,11 @@ class Worker:
self.event_sender.send_nowait(
NodeDownloadProgress(download_progress=status)
)
- self.event_sender.send_nowait(TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Complete))
+ self.event_sender.send_nowait(
+ TaskStatusUpdated(
+ task_id=task.task_id, task_status=TaskStatus.Complete
+ )
+ )
elif (
progress.status == "in_progress"
and current_time() - last_progress_time > throttle_interval_secs
diff --git a/src/exo/worker/plan.py b/src/exo/worker/plan.py
index af46a3ff..dfdda537 100644
--- a/src/exo/worker/plan.py
+++ b/src/exo/worker/plan.py
@@ -60,7 +60,9 @@ def _kill_runner(
) -> Shutdown | None:
for runner in runners.values():
if (instance_id := runner.bound_instance.instance.instance_id) not in instances:
- return Shutdown(instance_id=instance_id, runner_id = runner.bound_instance.bound_runner_id)
+ return Shutdown(
+ instance_id=instance_id, runner_id=runner.bound_instance.bound_runner_id
+ )
""" --- Potential code to kill a runner if any runners in its instance have failed ---
global_runners_in_instance = runner.bound_instance.instance.shard_assignments.node_to_runner.values()
diff --git a/src/exo/worker/runner/bootstrap.py b/src/exo/worker/runner/bootstrap.py
index 989b8723..e05b4789 100644
--- a/src/exo/worker/runner/bootstrap.py
+++ b/src/exo/worker/runner/bootstrap.py
@@ -40,6 +40,7 @@ def entrypoint(
faulthandler.enable(file=sys.stderr, all_threads=True)
"""
import os
+
os.environ["MLX_METAL_FAST_SYNCH"] = "1"
global logger
diff --git a/src/exo/worker/runner/generate.py b/src/exo/worker/runner/generate.py
index 1293184c..09d51b6c 100644
--- a/src/exo/worker/runner/generate.py
+++ b/src/exo/worker/runner/generate.py
@@ -1,14 +1,14 @@
-from typing import Any, Callable, Generator, get_args, cast
+from typing import Any, Callable, Generator, cast, get_args
import mlx.core as mx
-from mlx_lm.models.cache import KVCache
from mlx_lm import stream_generate
+from mlx_lm.models.cache import KVCache
+from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.engines.mlx import Model
-from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.engines.mlx.utils_mlx import (
- make_kv_cache,
apply_chat_template,
+ make_kv_cache,
mx_barrier,
)
from exo.shared.openai_compat import FinishReason
@@ -16,7 +16,6 @@ from exo.shared.types.api import ChatCompletionMessage
from exo.shared.types.tasks import ChatCompletionTaskParams
from exo.shared.types.worker.commands_runner import (
GenerationResponse,
- TokenizedResponse,
)
from exo.worker.runner.bootstrap import logger
@@ -37,6 +36,7 @@ def maybe_quantize_kv_cache(
):
prompt_cache[e] = c.to_quantized(group_size=kv_group_size, bits=kv_bits)
+
def warmup_inference(
model: Model,
tokenizer: TokenizerWrapper,
@@ -109,11 +109,17 @@ def mlx_generate(
prefill_step_size=65536,
):
logger.info(out.text)
- if out.finish_reason != None and out.finish_reason not in get_args(FinishReason):
+ if out.finish_reason is not None and out.finish_reason not in get_args(
+ FinishReason
+ ):
# We don't throw here as this failure case is really not all that bad
# Just log the error and move on
- logger.warning(f"Model generated unexpected finish_reason: {out.finish_reason}")
+ logger.warning(
+ f"Model generated unexpected finish_reason: {out.finish_reason}"
+ )
yield GenerationResponse(
- text=out.text, token=out.token, finish_reason=cast(FinishReason | None, out.finish_reason)
+ text=out.text,
+ token=out.token,
+ finish_reason=cast(FinishReason | None, out.finish_reason),
)
← 631cb810 kimi k2 thinking
·
back to Exo
·
fix kimi eos token ids d793f5f9 →