[object Object]

← 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

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 →