From a9c7b1c68a02e6565e3b8aab045cfe890e466f42 Mon Sep 17 00:00:00 2001 From: Evan Date: Mon, 23 Mar 2026 19:14:28 +0000 Subject: [PATCH] urgk --- .typings/mlx/nn/layers/activations.pyi | 2 +- .typings/mlx/nn/layers/containers.pyi | 2 +- .typings/mlx/nn/layers/convolution.pyi | 2 +- .../mlx/nn/layers/convolution_transpose.pyi | 2 +- .typings/mlx/nn/layers/distributed.pyi | 2 +- .typings/mlx/nn/layers/dropout.pyi | 2 +- .typings/mlx/nn/layers/embedding.pyi | 2 +- .typings/mlx/nn/layers/linear.pyi | 2 +- .typings/mlx/nn/layers/normalization.pyi | 2 +- .typings/mlx/nn/layers/pooling.pyi | 2 +- .../mlx/nn/layers/positional_encoding.pyi | 2 +- .typings/mlx/nn/layers/quantized.pyi | 2 +- .typings/mlx/nn/layers/recurrent.pyi | 2 +- .typings/mlx/nn/layers/transformer.pyi | 2 +- .typings/mlx/nn/layers/upsample.pyi | 2 +- python/exo_core/pyproject.toml | 6 ++- python/exo_core/src/exo_core/engine.py | 19 ++++---- .../src/exo_core/tokenizers/__init__.py | 0 .../tokenizers}/model_output_parsers.py | 17 ++++--- .../{utils => tokenizers}/tool_parsers.py | 0 python/exo_core/src/exo_core/types/chunks.py | 2 +- .../exo_core/src/exo_core/types/instances.py | 3 +- python/exo_core/src/exo_core/types/runners.py | 3 +- .../src/mlx_engine}/batch_generator.py | 48 ++++++++++--------- python/mlx_engine/src/mlx_engine/builder.py | 38 +++++++++------ python/mlx_engine/src/mlx_engine/cache.py | 7 ++- .../src/mlx_engine/dsml_encoding.py | 3 +- .../mlx_engine/generator/batch_generate.py | 18 +++---- .../src/mlx_engine/generator/generate.py | 14 +++--- python/parts.nix | 8 +--- python/vllm_engine/src/vllm_engine/builder.py | 29 ++++++----- .../src/vllm_engine/growable_cache.py | 3 +- .../src/vllm_engine/vllm_generator.py | 29 ++++++----- src/exo/api/main.py | 18 +++---- src/exo/download/coordinator.py | 8 ++-- src/exo/download/impl_shard_downloader.py | 2 +- src/exo/download/shard_downloader.py | 2 +- src/exo/download/tests/test_re_download.py | 7 ++- src/exo/main.py | 2 +- src/exo/master/main.py | 2 +- src/exo/master/placement.py | 2 +- src/exo/master/tests/test_master.py | 2 +- src/exo/master/tests/test_placement.py | 2 +- src/exo/master/tests/test_placement_utils.py | 2 +- src/exo/routing/event_router.py | 2 +- src/exo/routing/router.py | 2 +- src/exo/shared/election.py | 2 +- src/exo/shared/tests/conftest.py | 2 +- src/exo/shared/tests/test_election.py | 2 +- src/exo/shared/tracing.py | 1 - src/exo/shared/types/commands.py | 9 ++-- src/exo/utils/info_gatherer/info_gatherer.py | 2 +- src/exo/utils/info_gatherer/net_profile.py | 2 +- src/exo/utils/tests/test_mp_channel.py | 3 +- .../worker/engines/image/distributed_model.py | 6 +-- src/exo/worker/engines/image/generate.py | 16 +++---- src/exo/worker/main.py | 6 +-- src/exo/worker/runner/bootstrap.py | 7 +-- src/exo/worker/runner/image_models/runner.py | 12 ++--- src/exo/worker/runner/llm_inference/runner.py | 18 +++---- src/exo/worker/runner/runner_supervisor.py | 2 +- .../unittests/test_mlx/test_auto_parallel.py | 4 +- .../test_prefix_cache_architectures.py | 11 ++--- .../unittests/test_mlx/test_tokenizers.py | 3 +- .../unittests/test_runner/test_dsml_e2e.py | 2 +- .../test_runner/test_event_ordering.py | 2 +- .../test_runner/test_parse_tool_calls.py | 2 +- .../test_runner/test_runner_supervisor.py | 2 +- tests/headless_runner.py | 2 +- uv.lock | 8 +++- 70 files changed, 235 insertions(+), 221 deletions(-) create mode 100644 python/exo_core/src/exo_core/tokenizers/__init__.py rename {src/exo/worker/runner/llm_inference => python/exo_core/src/exo_core/tokenizers}/model_output_parsers.py (98%) rename python/exo_core/src/exo_core/{utils => tokenizers}/tool_parsers.py (100%) rename {src/exo/worker/runner/llm_inference => python/mlx_engine/src/mlx_engine}/batch_generator.py (94%) diff --git a/.typings/mlx/nn/layers/activations.pyi b/.typings/mlx/nn/layers/activations.pyi index adacb4da..d00298f5 100644 --- a/.typings/mlx/nn/layers/activations.pyi +++ b/.typings/mlx/nn/layers/activations.pyi @@ -6,7 +6,7 @@ from functools import partial from typing import Any import mlx.core as mx -from base import Module +from .base import Module @partial(mx.compile, shapeless=True) def sigmoid(x: mx.array) -> mx.array: diff --git a/.typings/mlx/nn/layers/containers.pyi b/.typings/mlx/nn/layers/containers.pyi index 068ea179..bdf0b270 100644 --- a/.typings/mlx/nn/layers/containers.pyi +++ b/.typings/mlx/nn/layers/containers.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Callable import mlx.core as mx -from base import Module +from .base import Module class Sequential(Module): """A layer that calls the passed callables in order. diff --git a/.typings/mlx/nn/layers/convolution.pyi b/.typings/mlx/nn/layers/convolution.pyi index 28b4ffd3..4d849b2e 100644 --- a/.typings/mlx/nn/layers/convolution.pyi +++ b/.typings/mlx/nn/layers/convolution.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Union import mlx.core as mx -from base import Module +from .base import Module class Conv1d(Module): """Applies a 1-dimensional convolution over the multi-channel input sequence. diff --git a/.typings/mlx/nn/layers/convolution_transpose.pyi b/.typings/mlx/nn/layers/convolution_transpose.pyi index 8fe11b4a..f9d78ddc 100644 --- a/.typings/mlx/nn/layers/convolution_transpose.pyi +++ b/.typings/mlx/nn/layers/convolution_transpose.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Union import mlx.core as mx -from base import Module +from .base import Module class ConvTranspose1d(Module): """Applies a 1-dimensional transposed convolution over the multi-channel input sequence. diff --git a/.typings/mlx/nn/layers/distributed.pyi b/.typings/mlx/nn/layers/distributed.pyi index 5be9cc4b..fda35c4a 100644 --- a/.typings/mlx/nn/layers/distributed.pyi +++ b/.typings/mlx/nn/layers/distributed.pyi @@ -6,7 +6,7 @@ from functools import lru_cache from typing import Callable, Optional, Union import mlx.core as mx -from base import Module +from .base import Module from mlx.nn.layers.linear import Linear @lru_cache diff --git a/.typings/mlx/nn/layers/dropout.pyi b/.typings/mlx/nn/layers/dropout.pyi index 00ef6f01..b4506128 100644 --- a/.typings/mlx/nn/layers/dropout.pyi +++ b/.typings/mlx/nn/layers/dropout.pyi @@ -3,7 +3,7 @@ This type stub file was generated by pyright. """ import mlx.core as mx -from base import Module +from .base import Module class Dropout(Module): r"""Randomly zero a portion of the elements during training. diff --git a/.typings/mlx/nn/layers/embedding.pyi b/.typings/mlx/nn/layers/embedding.pyi index e273c801..14fd15a0 100644 --- a/.typings/mlx/nn/layers/embedding.pyi +++ b/.typings/mlx/nn/layers/embedding.pyi @@ -3,7 +3,7 @@ This type stub file was generated by pyright. """ import mlx.core as mx -from base import Module +from .base import Module from .quantized import QuantizedEmbedding diff --git a/.typings/mlx/nn/layers/linear.pyi b/.typings/mlx/nn/layers/linear.pyi index 07e93a43..afc8de9c 100644 --- a/.typings/mlx/nn/layers/linear.pyi +++ b/.typings/mlx/nn/layers/linear.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Any import mlx.core as mx -from base import Module +from .base import Module from .quantized import QuantizedLinear diff --git a/.typings/mlx/nn/layers/normalization.pyi b/.typings/mlx/nn/layers/normalization.pyi index 216ccfff..fa71fbd4 100644 --- a/.typings/mlx/nn/layers/normalization.pyi +++ b/.typings/mlx/nn/layers/normalization.pyi @@ -3,7 +3,7 @@ This type stub file was generated by pyright. """ import mlx.core as mx -from base import Module +from .base import Module class InstanceNorm(Module): r"""Applies instance normalization [1] on the inputs. diff --git a/.typings/mlx/nn/layers/pooling.pyi b/.typings/mlx/nn/layers/pooling.pyi index 36b0ca24..dbb51a1a 100644 --- a/.typings/mlx/nn/layers/pooling.pyi +++ b/.typings/mlx/nn/layers/pooling.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Optional, Tuple, Union import mlx.core as mx -from base import Module +from .base import Module class _Pool(Module): def __init__( diff --git a/.typings/mlx/nn/layers/positional_encoding.pyi b/.typings/mlx/nn/layers/positional_encoding.pyi index 14e07e14..3019e98d 100644 --- a/.typings/mlx/nn/layers/positional_encoding.pyi +++ b/.typings/mlx/nn/layers/positional_encoding.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Optional import mlx.core as mx -from base import Module +from .base import Module class RoPE(Module): """Implements the rotary positional encoding. diff --git a/.typings/mlx/nn/layers/quantized.pyi b/.typings/mlx/nn/layers/quantized.pyi index 137a4c8e..6b064f90 100644 --- a/.typings/mlx/nn/layers/quantized.pyi +++ b/.typings/mlx/nn/layers/quantized.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Callable, Optional, Union import mlx.core as mx -from base import Module +from .base import Module def quantize( model: Module, diff --git a/.typings/mlx/nn/layers/recurrent.pyi b/.typings/mlx/nn/layers/recurrent.pyi index d31d9382..868dbf72 100644 --- a/.typings/mlx/nn/layers/recurrent.pyi +++ b/.typings/mlx/nn/layers/recurrent.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Callable, Optional import mlx.core as mx -from base import Module +from .base import Module class RNN(Module): r"""An Elman recurrent layer. diff --git a/.typings/mlx/nn/layers/transformer.pyi b/.typings/mlx/nn/layers/transformer.pyi index 9274a823..c47d28c8 100644 --- a/.typings/mlx/nn/layers/transformer.pyi +++ b/.typings/mlx/nn/layers/transformer.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Any, Callable, Optional import mlx.core as mx -from base import Module +from .base import Module class MultiHeadAttention(Module): """Implements the scaled dot product attention with multiple heads. diff --git a/.typings/mlx/nn/layers/upsample.pyi b/.typings/mlx/nn/layers/upsample.pyi index 1ef3298c..6e9a0806 100644 --- a/.typings/mlx/nn/layers/upsample.pyi +++ b/.typings/mlx/nn/layers/upsample.pyi @@ -5,7 +5,7 @@ This type stub file was generated by pyright. from typing import Literal, Tuple, Union import mlx.core as mx -from base import Module +from .base import Module def upsample_nearest(x: mx.array, scale_factor: Tuple) -> mx.array: ... def upsample_linear( diff --git a/python/exo_core/pyproject.toml b/python/exo_core/pyproject.toml index bafe5d41..7386e537 100644 --- a/python/exo_core/pyproject.toml +++ b/python/exo_core/pyproject.toml @@ -5,7 +5,11 @@ description = "Add your description here" readme = "README.md" authors = [{ name = "Evan", email = "evanev7@gmail.com" }] requires-python = ">=3.13" -dependencies = ["pydantic>=2.13.0b2"] +dependencies = [ + "mlx-lm", # TODO: depend on transformers or other + "openai-harmony", # inherit from workspace + "pydantic", # inherit from workspace +] [build-system] requires = ["uv_build>=0.9.24,<0.10.0"] diff --git a/python/exo_core/src/exo_core/engine.py b/python/exo_core/src/exo_core/engine.py index 01c018f1..5b8439f4 100644 --- a/python/exo_core/src/exo_core/engine.py +++ b/python/exo_core/src/exo_core/engine.py @@ -1,14 +1,15 @@ from abc import ABC, abstractmethod from collections.abc import Callable, Iterable -from typing import Self from exo_core.types.tasks import TaskId -class Cancelled: pass +class Cancelled: + pass -class Finished: pass +class Finished: + pass CANCEL_ALL_TASKS = TaskId("CANCEL_TALL_TASKS") @@ -45,12 +46,12 @@ class Engine[TaskType, ResponseType](ABC): class EngineBuilder[SetupType, TaskType, ResponseType](ABC): - @classmethod - @abstractmethod - def create( - cls, - bound_instance: SetupType, - ) -> Self: ... + # @classmethod + # @abstractmethod + # def create( + # cls, + # bound_instance: SetupType, + # ) -> Self: ... @abstractmethod def connect(self) -> None: ... diff --git a/python/exo_core/src/exo_core/tokenizers/__init__.py b/python/exo_core/src/exo_core/tokenizers/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/exo/worker/runner/llm_inference/model_output_parsers.py b/python/exo_core/src/exo_core/tokenizers/model_output_parsers.py similarity index 98% rename from src/exo/worker/runner/llm_inference/model_output_parsers.py rename to python/exo_core/src/exo_core/tokenizers/model_output_parsers.py index 05c1f69c..6bd60710 100644 --- a/src/exo/worker/runner/llm_inference/model_output_parsers.py +++ b/python/exo_core/src/exo_core/tokenizers/model_output_parsers.py @@ -2,8 +2,10 @@ from collections.abc import Generator from functools import cache from typing import Any -from exo_core.types.common import ModelId -from exo_core.types.runner_response import GenerationResponse, ToolCallResponse +from loguru import logger +from mlx_engine.utils_mlx import ( + detect_thinking_prompt_suffix, +) from mlx_lm.tokenizer_utils import TokenizerWrapper from openai_harmony import ( HarmonyEncodingName, @@ -13,12 +15,13 @@ from openai_harmony import ( load_harmony_encoding, ) -from exo.api.types import ToolCallItem -from mlx_engine.utils_mlx import ( - detect_thinking_prompt_suffix, +from exo_core.tokenizers.tool_parsers import ToolParser +from exo_core.types.common import ModelId +from exo_core.types.runner_response import ( + GenerationResponse, + ToolCallItem, + ToolCallResponse, ) -from loguru import logger -from exo_core.utils.tool_parsers import ToolParser @cache diff --git a/python/exo_core/src/exo_core/utils/tool_parsers.py b/python/exo_core/src/exo_core/tokenizers/tool_parsers.py similarity index 100% rename from python/exo_core/src/exo_core/utils/tool_parsers.py rename to python/exo_core/src/exo_core/tokenizers/tool_parsers.py diff --git a/python/exo_core/src/exo_core/types/chunks.py b/python/exo_core/src/exo_core/types/chunks.py index 3694057a..c6cd9941 100644 --- a/python/exo_core/src/exo_core/types/chunks.py +++ b/python/exo_core/src/exo_core/types/chunks.py @@ -1,8 +1,8 @@ from collections.abc import Generator from typing import Any, Literal -from exo_core.model_cards import ModelId from exo_core.models import TaggedModel +from exo_core.types.common import ModelId from exo_core.types.runner_response import ( FinishReason, GenerationStats, diff --git a/python/exo_core/src/exo_core/types/instances.py b/python/exo_core/src/exo_core/types/instances.py index addb8f72..92c1baac 100644 --- a/python/exo_core/src/exo_core/types/instances.py +++ b/python/exo_core/src/exo_core/types/instances.py @@ -5,7 +5,8 @@ from pydantic import model_validator from exo_core.model_cards import ModelTask from exo_core.models import CamelCaseModel, TaggedModel from exo_core.types.common import Host, Id, NodeId -from exo_core.types.runners import RunnerId, ShardAssignments, ShardMetadata +from exo_core.types.runners import RunnerId, ShardAssignments +from exo_core.types.shards import ShardMetadata class InstanceId(Id): diff --git a/python/exo_core/src/exo_core/types/runners.py b/python/exo_core/src/exo_core/types/runners.py index d6f89f6d..1bc059c9 100644 --- a/python/exo_core/src/exo_core/types/runners.py +++ b/python/exo_core/src/exo_core/types/runners.py @@ -2,9 +2,8 @@ from collections.abc import Mapping from pydantic import model_validator -from exo_core.model_cards import ModelId from exo_core.models import CamelCaseModel, TaggedModel -from exo_core.types.common import Id, NodeId +from exo_core.types.common import Id, ModelId, NodeId from exo_core.types.shards import ShardMetadata diff --git a/src/exo/worker/runner/llm_inference/batch_generator.py b/python/mlx_engine/src/mlx_engine/batch_generator.py similarity index 94% rename from src/exo/worker/runner/llm_inference/batch_generator.py rename to python/mlx_engine/src/mlx_engine/batch_generator.py index 733f6502..6d903773 100644 --- a/src/exo/worker/runner/llm_inference/batch_generator.py +++ b/python/mlx_engine/src/mlx_engine/batch_generator.py @@ -8,19 +8,23 @@ from typing import TYPE_CHECKING import mlx.core as mx from exo_core.constants import EXO_MAX_CONCURRENT_REQUESTS from exo_core.types.chunks import ErrorChunk, PrefillProgressChunk -from exo_core.types.common import ModelId +from exo_core.types.common import CommandId, ModelId from exo_core.types.runner_response import GenerationResponse, ToolCallResponse from exo_core.types.tasks import CANCEL_ALL_TASKS, TaskId, TextGeneration from exo_core.types.text_generation import TextGenerationTaskParams +from exo_core.utils.channels import MpReceiver, MpSender from mlx_lm.tokenizer_utils import TokenizerWrapper -from exo.shared.types.events import ChunkGenerated, Event -from exo.utils.channels import MpReceiver, MpSender from mlx_engine.cache import KVPrefixCache from mlx_engine.generator.batch_generate import ExoBatchGenerator if TYPE_CHECKING: from vllm_engine.vllm_generator import VllmBatchEngine +from exo_core.engine import Cancelled, Engine, Finished +from exo_core.tokenizers.model_output_parsers import apply_all_parsers +from exo_core.tokenizers.tool_parsers import ToolParser +from loguru import logger + from mlx_engine.generator.generate import ( PrefillCancelled, ) @@ -29,11 +33,6 @@ from mlx_engine.utils_mlx import ( mx_all_gather_tasks, mx_any, ) -from loguru import logger -from exo_core.engine import Engine, Cancelled, Finished - -from .model_output_parsers import apply_all_parsers -from exo_core.utils.tool_parsers import ToolParser class GeneratorQueue[T]: @@ -74,7 +73,9 @@ def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None: @dataclass(eq=False) -class SequentialGenerator(Engine[TextGeneration, GenerationResponse | ToolCallResponse]): +class SequentialGenerator( + Engine[TextGeneration, GenerationResponse | ToolCallResponse] +): tokenizer: TokenizerWrapper group: mx.distributed.Group | None kv_prefix_cache: KVPrefixCache | None @@ -82,7 +83,7 @@ class SequentialGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRe model_id: ModelId device_rank: int cancel_receiver: MpReceiver[TaskId] - event_sender: MpSender[Event] + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]] _generate_fn: Callable[..., Generator[GenerationResponse]] _warmup_fn: Callable[[], int] check_for_cancel_every: int = 50 @@ -197,10 +198,11 @@ class SequentialGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRe def _send_error(self, task: TextGeneration, e: Exception) -> None: if self.device_rank == 0: + # TODO: sync channels? self.event_sender.send( - ChunkGenerated( - command_id=task.command_id, - chunk=ErrorChunk( + ( + task.command_id, + ErrorChunk( model=self.model_id, finish_reason="error", error_message=str(e), @@ -215,9 +217,9 @@ class SequentialGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRe def on_prefill_progress(processed: int, total: int) -> None: if self.device_rank == 0: self.event_sender.send( - ChunkGenerated( - command_id=task.command_id, - chunk=PrefillProgressChunk( + ( + task.command_id, + PrefillProgressChunk( model=self.model_id, processed_tokens=processed, total_tokens=total, @@ -268,7 +270,7 @@ class BatchGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRespons model_id: ModelId device_rank: int cancel_receiver: MpReceiver[TaskId] - event_sender: MpSender[Event] + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]] _gen: "ExoBatchGenerator | VllmBatchEngine" max_concurrent_requests: int = EXO_MAX_CONCURRENT_REQUESTS check_for_cancel_every: int = 50 @@ -412,9 +414,9 @@ class BatchGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRespons def _send_error(self, task: TextGeneration, e: Exception) -> None: if self.device_rank == 0: self.event_sender.send( - ChunkGenerated( - command_id=task.command_id, - chunk=ErrorChunk( + ( + task.command_id, + ErrorChunk( model=self.model_id, finish_reason="error", error_message=str(e), @@ -429,9 +431,9 @@ class BatchGenerator(Engine[TextGeneration, GenerationResponse | ToolCallRespons def on_prefill_progress(processed: int, total: int) -> None: if self.device_rank == 0: self.event_sender.send( - ChunkGenerated( - command_id=task.command_id, - chunk=PrefillProgressChunk( + ( + task.command_id, + PrefillProgressChunk( model=self.model_id, processed_tokens=processed, total_tokens=total, diff --git a/python/mlx_engine/src/mlx_engine/builder.py b/python/mlx_engine/src/mlx_engine/builder.py index 47b4063a..b628db8f 100644 --- a/python/mlx_engine/src/mlx_engine/builder.py +++ b/python/mlx_engine/src/mlx_engine/builder.py @@ -1,27 +1,36 @@ +import contextlib +import os from dataclasses import dataclass -from typing import Self, Callable -from exo_core.engine import EngineBuilder, Engine -from exo_core.types.common import ModelId +from typing import Callable, Self + +import mlx.core as mx +from exo_core.engine import EngineBuilder +from exo_core.tokenizers.tool_parsers import make_mlx_parser +from exo_core.types.chunks import ErrorChunk, PrefillProgressChunk +from exo_core.types.common import CommandId, ModelId from exo_core.types.instances import BoundInstance -from exo_core.types.tasks import TextGeneration -from exo_core.types.runner_response import GenerationResponse -from mlx_engine.utils_mlx import initialize_mlx, load_mlx_items -from mlx_engine.types import Model +from exo_core.types.runner_response import GenerationResponse, ToolCallResponse +from exo_core.types.tasks import TaskId, TextGeneration +from exo_core.utils.channels import MpReceiver, MpSender +from loguru import logger +from mlx_lm.tokenizer_utils import TokenizerWrapper + +from mlx_engine.batch_generator import BatchGenerator, SequentialGenerator +from mlx_engine.cache import KVPrefixCache +from mlx_engine.generator.batch_generate import ExoBatchGenerator from mlx_engine.generator.generate import ( mlx_generate, warmup_inference, ) -from exo_core.utils.tool_parsers import make_mlx_parser +from mlx_engine.types import Model +from mlx_engine.utils_mlx import initialize_mlx, load_mlx_items @dataclass -class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse]): - import mlx.core as mx - from mlx_lm.tokenizer_utils import TokenizerWrapper - +class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse]): model_id: ModelId bound_instance: BoundInstance - event_sender: MpSender[Event] + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]] cancel_receiver: MpReceiver[TaskId] inference_model: Model | None = None tokenizer: TokenizerWrapper | None = None @@ -31,7 +40,7 @@ class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse def create( cls, bound_instance: BoundInstance, - event_sender: MpSender[Event], + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]], cancel_receiver: MpReceiver[TaskId], ) -> Self: return cls( @@ -105,7 +114,6 @@ class MlxBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse _generate_fn=generate_fn, _warmup_fn=warmup_fn, ) - from exo.worker.runner.llm_inference.batch_generator import ExoBatchGenerator logger.info("using BatchGenerator") gen = ExoBatchGenerator( diff --git a/python/mlx_engine/src/mlx_engine/cache.py b/python/mlx_engine/src/mlx_engine/cache.py index 7b67b26c..0f74ebdd 100644 --- a/python/mlx_engine/src/mlx_engine/cache.py +++ b/python/mlx_engine/src/mlx_engine/cache.py @@ -25,7 +25,7 @@ if TYPE_CHECKING: # Fraction of device memory above which LRU eviction kicks in. # Smaller machines need more aggressive eviction. def _default_memory_threshold() -> float: - total_gb = Memory.from_bytes(psutil.virtual_memory().total).in_gb + total_gb = Memory.from_bytes(psutil.virtual_memory().total).in_gb # pyright: ignore[reportAny] if total_gb >= 128: return 0.85 if total_gb >= 64: @@ -220,7 +220,6 @@ class KVPrefixCache: def lookup( self, prompt_token_ids: list[int] ) -> tuple["TorchKVCache | None", int, int | None]: - from exo.worker.engines.vllm.kv_cache import TorchKVCache prompt_mx = mx.array(prompt_token_ids) max_length = len(prompt_token_ids) @@ -352,14 +351,14 @@ def get_prefix_length(prompt: mx.array, cached_prompt: mx.array) -> int: def get_available_memory() -> Memory: - mem: int = psutil.virtual_memory().available + mem: int = psutil.virtual_memory().available # pyright: ignore[reportAny] return Memory.from_bytes(mem) def get_memory_used_percentage() -> float: mem = psutil.virtual_memory() # percent is 0-100 - return float(mem.percent / 100) + return float(mem.percent / 100) # pyright: ignore[reportAny] def make_kv_cache( diff --git a/python/mlx_engine/src/mlx_engine/dsml_encoding.py b/python/mlx_engine/src/mlx_engine/dsml_encoding.py index 397d94ea..3cc8f244 100644 --- a/python/mlx_engine/src/mlx_engine/dsml_encoding.py +++ b/python/mlx_engine/src/mlx_engine/dsml_encoding.py @@ -2,9 +2,8 @@ import json import re from typing import Any -from mlx_lm.chat_templates import deepseek_v32 - from exo_core.types.runner_response import ToolCallItem +from mlx_lm.chat_templates import deepseek_v32 BOS_TOKEN: str = deepseek_v32.bos_token EOS_TOKEN: str = deepseek_v32.eos_token diff --git a/python/mlx_engine/src/mlx_engine/generator/batch_generate.py b/python/mlx_engine/src/mlx_engine/generator/batch_generate.py index 5e0485bd..2aec5b22 100644 --- a/python/mlx_engine/src/mlx_engine/generator/batch_generate.py +++ b/python/mlx_engine/src/mlx_engine/generator/batch_generate.py @@ -4,7 +4,15 @@ from typing import Callable, cast import mlx.core as mx from exo_core.types.common import ModelId -from exo_core.types.runner_response import GenerationResponse +from exo_core.types.runner_response import ( + CompletionTokensDetails, + FinishReason, + GenerationResponse, + GenerationStats, + PromptTokensDetails, + TopLogprobItem, + Usage, +) from exo_core.types.tasks import TaskId from exo_core.types.text_generation import TextGenerationTaskParams from exo_core.utils.memory import Memory @@ -16,14 +24,6 @@ from mlx_lm.models.cache import RotatingKVCache from mlx_lm.sample_utils import make_logits_processors, make_sampler from mlx_lm.tokenizer_utils import StreamingDetokenizer, TokenizerWrapper -from exo.api.types import ( - CompletionTokensDetails, - FinishReason, - GenerationStats, - PromptTokensDetails, - TopLogprobItem, - Usage, -) from mlx_engine.cache import ( CacheSnapshot, KVPrefixCache, diff --git a/python/mlx_engine/src/mlx_engine/generator/generate.py b/python/mlx_engine/src/mlx_engine/generator/generate.py index 559e67f1..e1a9b17c 100644 --- a/python/mlx_engine/src/mlx_engine/generator/generate.py +++ b/python/mlx_engine/src/mlx_engine/generator/generate.py @@ -7,7 +7,13 @@ from typing import Callable, Generator, cast, get_args import mlx.core as mx from exo_core.types.common import ModelId from exo_core.types.runner_response import ( + CompletionTokensDetails, + FinishReason, GenerationResponse, + GenerationStats, + PromptTokensDetails, + TopLogprobItem, + Usage, ) from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams from exo_core.utils.memory import Memory @@ -20,14 +26,6 @@ from mlx_lm.models.cache import ArraysCache, RotatingKVCache from mlx_lm.sample_utils import make_logits_processors, make_sampler from mlx_lm.tokenizer_utils import TokenizerWrapper -from exo.api.types import ( - CompletionTokensDetails, - FinishReason, - GenerationStats, - PromptTokensDetails, - TopLogprobItem, - Usage, -) from mlx_engine.auto_parallel import ( PipelineFirstLayer, PipelineLastLayer, diff --git a/python/parts.nix b/python/parts.nix index c6f209e2..231bb750 100644 --- a/python/parts.nix +++ b/python/parts.nix @@ -188,12 +188,8 @@ autoPatchelfIgnoreMissingDeps = (old.autoPatchelfIgnoreMissingDeps or [ ]) ++ [ "libcuda.so.1" ]; }); xgrammar = prev.xgrammar.overrideAttrs (old: { - nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [ final.setuptools final.scikit-build-core final.packaging final.pathspec pkgs.cmake final.nanobind ]; - - prePatch = '' - cat cpp/nanobind/CMakeLists.txt - ''; - patches = (old.patches or [ ]) ++ [ ../nix/nanobind_cmake.patch ]; + nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [ pkgs.cmake ]; + patches = (old.patches or [ ]) ++ [ ../nix/xgrammar_cmake.patch ]; }); vllm = prev.vllm.overrideAttrs (old: { patches = (old.patches or [ ]) ++ [ ../nix/vllm_uv2nix_cmake.patch ]; diff --git a/python/vllm_engine/src/vllm_engine/builder.py b/python/vllm_engine/src/vllm_engine/builder.py index b8560a09..0b1eb8b6 100644 --- a/python/vllm_engine/src/vllm_engine/builder.py +++ b/python/vllm_engine/src/vllm_engine/builder.py @@ -1,30 +1,37 @@ +import contextlib +import os from dataclasses import dataclass -from typing import Self, Callable +from typing import Callable, Self + from exo_core.constants import EXO_MODELS_DIR -from exo_core.engine import EngineBuilder, Engine -from exo_core.types.common import ModelId +from exo_core.engine import Engine, EngineBuilder +from exo_core.types.chunks import ErrorChunk, PrefillProgressChunk +from exo_core.types.common import CommandId, ModelId from exo_core.types.instances import BoundInstance -from exo_core.types.tasks import TextGeneration -from exo_core.types.runner_response import GenerationResponse -from vllm_engine.vllm_generator import VllmBatchEngine -from vllm_engine.vllm_generator import load_vllm_engine +from exo_core.types.runner_response import GenerationResponse, ToolCallResponse +from exo_core.types.tasks import TaskId, TextGeneration +from exo_core.utils.channels import MpReceiver, MpSender +from loguru import logger +from mlx_engine.batch_generator import BatchGenerator + +from vllm_engine.vllm_generator import VllmBatchEngine, load_vllm_engine @dataclass -class VllmBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse]): +class VllmBuilder(EngineBuilder[BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse]): model_id: ModelId model_path: str trust_remote_code: bool cancel_receiver: MpReceiver[TaskId] - event_sender: MpSender[Event] + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]] bound_instance: BoundInstance @classmethod def create( cls, bound_instance: BoundInstance, - event_sender: MpSender[Event], - cancel_receiver: MpReceiver[TaskId], + cancel_receiver: MpReceiver[TaskId], + event_sender: MpSender[tuple[CommandId, ErrorChunk | PrefillProgressChunk]], ) -> Self: mid = bound_instance.instance.shard_assignments.model_id return cls( diff --git a/python/vllm_engine/src/vllm_engine/growable_cache.py b/python/vllm_engine/src/vllm_engine/growable_cache.py index af623c26..828d5762 100644 --- a/python/vllm_engine/src/vllm_engine/growable_cache.py +++ b/python/vllm_engine/src/vllm_engine/growable_cache.py @@ -1,9 +1,8 @@ import torch +from loguru import logger from mlx_engine.cache import KVPrefixCache from vllm.v1.worker.gpu_model_runner import GPUModelRunner -from loguru import logger - INITIAL_FRACTION = 0.05 GROWTH_HEADROOM_BYTES = 512 * 1024 * 1024 MIN_GROWTH_BLOCKS = 16 diff --git a/python/vllm_engine/src/vllm_engine/vllm_generator.py b/python/vllm_engine/src/vllm_engine/vllm_generator.py index 19a99d73..c60dbc71 100644 --- a/python/vllm_engine/src/vllm_engine/vllm_generator.py +++ b/python/vllm_engine/src/vllm_engine/vllm_generator.py @@ -8,28 +8,27 @@ from collections.abc import Callable, Generator from dataclasses import dataclass, field import torch -from mlx_engine.cache import KVPrefixCache -from mlx_engine.utils_mlx import get_eos_token_ids_for_model +from exo_core.tokenizers.tool_parsers import ToolParser, infer_tool_parser from exo_core.types.common import ModelId -from exo_core.types.runner_response import GenerationResponse -from exo_core.types.tasks import TaskId -from exo_core.types.text_generation import TextGenerationTaskParams -from exo_core.utils.memory import Memory -from exo_core.engine import Engine -from loguru import logger -from vllm.engine.arg_utils import EngineArgs -from vllm.v1.attention.backends.registry import AttentionBackendEnum -from vllm.sampling_params import SamplingParams -from vllm.v1.engine.llm_engine import LLMEngine -from vllm.v1.kv_cache_interface import KVCacheConfig - from exo_core.types.runner_response import ( CompletionTokensDetails, + GenerationResponse, GenerationStats, PromptTokensDetails, Usage, ) -from exo.worker.runner.llm_inference.tool_parsers import ToolParser, infer_tool_parser +from exo_core.types.tasks import TaskId +from exo_core.types.text_generation import TextGenerationTaskParams +from exo_core.utils.memory import Memory +from loguru import logger +from mlx_engine.cache import KVPrefixCache +from mlx_engine.utils_mlx import get_eos_token_ids_for_model +from vllm.engine.arg_utils import EngineArgs +from vllm.sampling_params import SamplingParams +from vllm.v1.attention.backends.registry import AttentionBackendEnum +from vllm.v1.engine.llm_engine import LLMEngine +from vllm.v1.kv_cache_interface import KVCacheConfig + from vllm_engine.growable_cache import ( get_model_runner, patch_vllm, diff --git a/src/exo/api/main.py b/src/exo/api/main.py index aa71caf5..39a5409b 100644 --- a/src/exo/api/main.py +++ b/src/exo/api/main.py @@ -36,8 +36,17 @@ from exo_core.types.chunks import ( ) from exo_core.types.common import CommandId, Id, ModelId, NodeId, SystemId from exo_core.types.downloads import DownloadCompleted +from exo_core.types.image_generation import ( + AdvancedImageParams, + BenchImageGenerationTaskParams, + ImageEditsTaskParams, + ImageGenerationTaskParams, + ImageSize, + normalize_image_size, +) from exo_core.types.instances import Instance, InstanceId, InstanceMeta from exo_core.types.shards import Sharding +from exo_core.utils.channels import Receiver, Sender, channel from exo_core.utils.memory import Memory from fastapi import FastAPI, File, Form, HTTPException, Query, Request, UploadFile from fastapi.middleware.cors import CORSMiddleware @@ -71,14 +80,6 @@ from exo.api.adapters.responses import ( generate_responses_stream, responses_request_to_text_generation, ) -from exo_core.types.image_generation import ( - normalize_image_size, - AdvancedImageParams, - BenchImageGenerationTaskParams, - ImageEditsTaskParams, - ImageGenerationTaskParams, - ImageSize, -) from exo.api.types import ( AddCustomModelParams, BenchChatCompletionRequest, @@ -173,7 +174,6 @@ from exo.shared.types.events import ( ) from exo.shared.types.state import State from exo.utils.banner import print_startup_banner -from exo.utils.channels import Receiver, Sender, channel from exo.utils.disk_event_log import DiskEventLog from exo.utils.power_sampler import PowerSampler from exo.utils.task_group import TaskGroup diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py index 902aa308..9412cacb 100644 --- a/src/exo/download/coordinator.py +++ b/src/exo/download/coordinator.py @@ -4,18 +4,17 @@ import anyio from anyio import current_time from exo_core.constants import EXO_MODELS_DIR, EXO_MODELS_PATH from exo_core.model_cards import get_model_cards -from exo_core.types.common import NodeId, ModelId +from exo_core.types.common import ModelId, NodeId from exo_core.types.downloads import ( DownloadCompleted, DownloadFailed, DownloadOngoing, DownloadPending, DownloadProgress, -) -from exo_core.types.shards import PipelineShardMetadata, ShardMetadata -from exo_core.types.downloads import ( RepoDownloadProgress, ) +from exo_core.types.shards import PipelineShardMetadata, ShardMetadata +from exo_core.utils.channels import Receiver, Sender from exo_core.utils.downloads import ( delete_model, map_repo_download_progress_to_download_progress_data, @@ -34,7 +33,6 @@ from exo.shared.types.events import ( Event, NodeDownloadProgress, ) -from exo.utils.channels import Receiver, Sender from exo.utils.task_group import TaskGroup diff --git a/src/exo/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py index 4b9f0af8..dc95f9c5 100644 --- a/src/exo/download/impl_shard_downloader.py +++ b/src/exo/download/impl_shard_downloader.py @@ -6,11 +6,11 @@ from typing import AsyncIterator, Callable from exo_core.model_cards import ModelCard, get_model_cards from exo_core.types.common import ModelId +from exo_core.types.downloads import RepoDownloadProgress from exo_core.types.shards import ( PipelineShardMetadata, ShardMetadata, ) -from exo_core.types.downloads import RepoDownloadProgress from exo_core.utils.downloads import download_shard from loguru import logger diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py index 91ab893f..c309c2f5 100644 --- a/src/exo/download/shard_downloader.py +++ b/src/exo/download/shard_downloader.py @@ -7,11 +7,11 @@ from typing import AsyncIterator, Callable from exo_core.model_cards import ModelCard, ModelTask from exo_core.types.common import ModelId +from exo_core.types.downloads import RepoDownloadProgress from exo_core.types.shards import ( PipelineShardMetadata, ShardMetadata, ) -from exo_core.types.downloads import RepoDownloadProgress from exo_core.utils.memory import Memory diff --git a/src/exo/download/tests/test_re_download.py b/src/exo/download/tests/test_re_download.py index 59e85c55..5e8d35ff 100644 --- a/src/exo/download/tests/test_re_download.py +++ b/src/exo/download/tests/test_re_download.py @@ -9,10 +9,10 @@ from typing import Callable from unittest.mock import AsyncMock, patch from exo_core.model_cards import ModelCard, ModelTask -from exo_core.types.common import NodeId, SystemId, ModelId -from exo_core.types.downloads import DownloadCompleted +from exo_core.types.common import ModelId, NodeId, SystemId +from exo_core.types.downloads import DownloadCompleted, RepoDownloadProgress from exo_core.types.shards import PipelineShardMetadata, ShardMetadata -from exo_core.types.downloads import RepoDownloadProgress +from exo_core.utils.channels import Receiver, Sender, channel from exo_core.utils.memory import Memory from exo.download.coordinator import DownloadCoordinator @@ -24,7 +24,6 @@ from exo.shared.types.commands import ( StartDownload, ) from exo.shared.types.events import Event, NodeDownloadProgress -from exo.utils.channels import Receiver, Sender, channel NODE_ID = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa") MODEL_ID = ModelId("test-org/test-model") diff --git a/src/exo/main.py b/src/exo/main.py index e91f31a2..cc1d9c8c 100644 --- a/src/exo/main.py +++ b/src/exo/main.py @@ -10,6 +10,7 @@ import anyio from exo_core.constants import EXO_LOG from exo_core.models import CamelCaseModel from exo_core.types.common import NodeId, SessionId +from exo_core.utils.channels import Receiver, channel from loguru import logger from pydantic import PositiveInt @@ -22,7 +23,6 @@ from exo.routing.event_router import EventRouter from exo.routing.router import Router, get_node_id_keypair from exo.shared.election import Election, ElectionResult from exo.shared.logging import logger_cleanup, logger_setup -from exo.utils.channels import Receiver, channel from exo.utils.task_group import TaskGroup from exo.worker.main import Worker diff --git a/src/exo/master/main.py b/src/exo/master/main.py index 7b7a1622..e8fade0e 100644 --- a/src/exo/master/main.py +++ b/src/exo/master/main.py @@ -17,6 +17,7 @@ from exo_core.types.tasks import ( from exo_core.types.tasks import ( TextGeneration as TextGenerationTask, ) +from exo_core.utils.channels import Receiver, Sender from loguru import logger from exo.master.placement import ( @@ -59,7 +60,6 @@ from exo.shared.types.events import ( TracesMerged, ) from exo.shared.types.state import State -from exo.utils.channels import Receiver, Sender from exo.utils.disk_event_log import DiskEventLog from exo.utils.event_buffer import MultiSourceBuffer from exo.utils.task_group import TaskGroup diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py index a1d417c1..e3c105dc 100644 --- a/src/exo/master/placement.py +++ b/src/exo/master/placement.py @@ -3,7 +3,7 @@ from collections.abc import Mapping from copy import deepcopy from typing import Sequence -from exo_core.types.common import NodeId, ModelId +from exo_core.types.common import ModelId, NodeId from exo_core.types.downloads import ( DownloadOngoing, DownloadProgress, diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py index 04827857..5c40f96d 100644 --- a/src/exo/master/tests/test_master.py +++ b/src/exo/master/tests/test_master.py @@ -14,6 +14,7 @@ from exo_core.types.shards import PipelineShardMetadata, Sharding from exo_core.types.tasks import TaskStatus from exo_core.types.tasks import TextGeneration as TextGenerationTask from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams +from exo_core.utils.channels import channel from exo_core.utils.memory import Memory from loguru import logger @@ -38,7 +39,6 @@ from exo.shared.types.events import ( from exo.shared.types.profiling import ( MemoryUsage, ) -from exo.utils.channels import channel @pytest.mark.asyncio diff --git a/src/exo/master/tests/test_placement.py b/src/exo/master/tests/test_placement.py index 59d0cae8..f34d03ec 100644 --- a/src/exo/master/tests/test_placement.py +++ b/src/exo/master/tests/test_placement.py @@ -1,6 +1,6 @@ import pytest from exo_core.model_cards import ModelCard, ModelTask -from exo_core.types.common import CommandId, NodeId, ModelId +from exo_core.types.common import CommandId, ModelId, NodeId from exo_core.types.instances import ( Instance, InstanceId, diff --git a/src/exo/master/tests/test_placement_utils.py b/src/exo/master/tests/test_placement_utils.py index 4bfeb6d3..ae0daaae 100644 --- a/src/exo/master/tests/test_placement_utils.py +++ b/src/exo/master/tests/test_placement_utils.py @@ -1,6 +1,6 @@ import pytest from exo_core.model_cards import ModelCard, ModelTask -from exo_core.types.common import NodeId, ModelId +from exo_core.types.common import ModelId, NodeId from exo_core.types.shards import ( CfgShardMetadata, PipelineShardMetadata, diff --git a/src/exo/routing/event_router.py b/src/exo/routing/event_router.py index 40612d7f..8cff1d6c 100644 --- a/src/exo/routing/event_router.py +++ b/src/exo/routing/event_router.py @@ -5,6 +5,7 @@ import anyio from anyio import BrokenResourceError, ClosedResourceError from anyio.abc import CancelScope from exo_core.types.common import SessionId, SystemId +from exo_core.utils.channels import Receiver, Sender, channel from loguru import logger from exo.shared.types.commands import ForwarderCommand, RequestEventLog @@ -15,7 +16,6 @@ from exo.shared.types.events import ( IndexedEvent, LocalForwarderEvent, ) -from exo.utils.channels import Receiver, Sender, channel from exo.utils.event_buffer import OrderedBuffer from exo.utils.task_group import TaskGroup diff --git a/src/exo/routing/router.py b/src/exo/routing/router.py index 98071e33..52c58931 100644 --- a/src/exo/routing/router.py +++ b/src/exo/routing/router.py @@ -13,6 +13,7 @@ from anyio import ( ) from exo_core.constants import EXO_NODE_ID_KEYPAIR from exo_core.models import CamelCaseModel +from exo_core.utils.channels import Receiver, Sender, channel from exo_pyo3_bindings import ( AllQueuesFullError, Keypair, @@ -24,7 +25,6 @@ from exo_pyo3_bindings import ( from filelock import FileLock from loguru import logger -from exo.utils.channels import Receiver, Sender, channel from exo.utils.task_group import TaskGroup from .connection_message import ConnectionMessage diff --git a/src/exo/shared/election.py b/src/exo/shared/election.py index 2be700e1..93ce3f17 100644 --- a/src/exo/shared/election.py +++ b/src/exo/shared/election.py @@ -8,11 +8,11 @@ from anyio import ( ) from exo_core.models import CamelCaseModel from exo_core.types.common import NodeId, SessionId +from exo_core.utils.channels import Receiver, Sender from loguru import logger from exo.routing.connection_message import ConnectionMessage from exo.shared.types.commands import ForwarderCommand -from exo.utils.channels import Receiver, Sender from exo.utils.task_group import TaskGroup DEFAULT_ELECTION_TIMEOUT = 3.0 diff --git a/src/exo/shared/tests/conftest.py b/src/exo/shared/tests/conftest.py index 6fc9b549..fcf50376 100644 --- a/src/exo/shared/tests/conftest.py +++ b/src/exo/shared/tests/conftest.py @@ -6,8 +6,8 @@ from typing import Generator import pytest from _pytest.logging import LogCaptureFixture from exo_core.model_cards import ModelCard, ModelTask -from exo_core.types.shards import PipelineShardMetadata, ShardMetadata from exo_core.types.common import ModelId +from exo_core.types.shards import PipelineShardMetadata, ShardMetadata from exo_core.utils.memory import Memory from loguru import logger diff --git a/src/exo/shared/tests/test_election.py b/src/exo/shared/tests/test_election.py index 6c600066..64e92bee 100644 --- a/src/exo/shared/tests/test_election.py +++ b/src/exo/shared/tests/test_election.py @@ -1,11 +1,11 @@ import pytest from anyio import create_task_group, fail_after, move_on_after from exo_core.types.common import NodeId, SessionId, SystemId +from exo_core.utils.channels import channel from exo.routing.connection_message import ConnectionMessage from exo.shared.election import Election, ElectionMessage, ElectionResult from exo.shared.types.commands import ForwarderCommand, TestCommand -from exo.utils.channels import channel # ======= # # Helpers # diff --git a/src/exo/shared/tracing.py b/src/exo/shared/tracing.py index bd2bdb4c..b3912f45 100644 --- a/src/exo/shared/tracing.py +++ b/src/exo/shared/tracing.py @@ -9,7 +9,6 @@ from pathlib import Path from typing import cast, final from exo_core.constants import EXO_TRACING_ENABLED - from loguru import logger # Context variable to track the current trace category for hierarchical nesting diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py index af970662..5ff2b993 100644 --- a/src/exo/shared/types/commands.py +++ b/src/exo/shared/types/commands.py @@ -2,15 +2,14 @@ from exo_core.model_cards import ModelCard from exo_core.models import CamelCaseModel, TaggedModel from exo_core.types.chunks import InputImageChunk from exo_core.types.common import CommandId, ModelId, NodeId, SystemId -from exo_core.types.instances import Instance, InstanceId, InstanceMeta -from exo_core.types.shards import Sharding, ShardMetadata -from exo_core.types.text_generation import TextGenerationTaskParams -from pydantic import Field - from exo_core.types.image_generation import ( ImageEditsTaskParams, ImageGenerationTaskParams, ) +from exo_core.types.instances import Instance, InstanceId, InstanceMeta +from exo_core.types.shards import Sharding, ShardMetadata +from exo_core.types.text_generation import TextGenerationTaskParams +from pydantic import Field class BaseCommand(TaggedModel): diff --git a/src/exo/utils/info_gatherer/info_gatherer.py b/src/exo/utils/info_gatherer/info_gatherer.py index 40db497b..be246f43 100644 --- a/src/exo/utils/info_gatherer/info_gatherer.py +++ b/src/exo/utils/info_gatherer/info_gatherer.py @@ -13,6 +13,7 @@ from anyio.streams.buffered import BufferedByteReceiveStream from anyio.streams.text import TextReceiveStream from exo_core.constants import EXO_CONFIG_FILE, EXO_MODELS_DIR from exo_core.models import TaggedModel +from exo_core.utils.channels import Sender from exo_core.utils.memory import Memory from loguru import logger from pydantic import ValidationError @@ -28,7 +29,6 @@ from exo.shared.types.thunderbolt import ( ThunderboltConnectivity, ThunderboltIdentifier, ) -from exo.utils.channels import Sender from exo.utils.task_group import TaskGroup from .macmon import MacmonMetrics diff --git a/src/exo/utils/info_gatherer/net_profile.py b/src/exo/utils/info_gatherer/net_profile.py index 0e3b70eb..bd99d699 100644 --- a/src/exo/utils/info_gatherer/net_profile.py +++ b/src/exo/utils/info_gatherer/net_profile.py @@ -5,11 +5,11 @@ import anyio import httpx from anyio import create_task_group from exo_core.types.common import NodeId +from exo_core.utils.channels import Sender, channel from loguru import logger from exo.shared.topology import Topology from exo.shared.types.profiling import NodeNetworkInfo -from exo.utils.channels import Sender, channel REACHABILITY_ATTEMPTS = 3 diff --git a/src/exo/utils/tests/test_mp_channel.py b/src/exo/utils/tests/test_mp_channel.py index 73a26adb..0ed93e4a 100644 --- a/src/exo/utils/tests/test_mp_channel.py +++ b/src/exo/utils/tests/test_mp_channel.py @@ -3,10 +3,9 @@ import time import pytest from anyio import fail_after +from exo_core.utils.channels import MpReceiver, MpSender, mp_channel from loguru import logger -from exo.utils.channels import MpReceiver, MpSender, mp_channel - def foo(recv: MpReceiver[str]): expected = ["hi", "hi 2", "bye"] diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index 843e4d80..b25279ee 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -3,13 +3,15 @@ from pathlib import Path from typing import Any, Literal, Optional import mlx.core as mx +from exo_core.types.image_generation import AdvancedImageParams from exo_core.types.instances import BoundInstance from exo_core.types.shards import CfgShardMetadata, PipelineShardMetadata from exo_core.utils.downloads import build_model_path +from loguru import logger from mflux.models.common.config.config import Config +from mlx_engine.utils_mlx import mlx_distributed_init, mx_barrier from PIL import Image -from exo_core.types.image_generation import AdvancedImageParams from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models import ( create_adapter_for_model, @@ -17,8 +19,6 @@ from exo.worker.engines.image.models import ( ) from exo.worker.engines.image.models.base import ModelAdapter from exo.worker.engines.image.pipeline import DiffusionRunner -from mlx_engine.utils_mlx import mlx_distributed_init, mx_barrier -from loguru import logger class DistributedImageModel: diff --git a/src/exo/worker/engines/image/generate.py b/src/exo/worker/engines/image/generate.py index 1a078f1a..ac50421d 100644 --- a/src/exo/worker/engines/image/generate.py +++ b/src/exo/worker/engines/image/generate.py @@ -7,20 +7,20 @@ from pathlib import Path from typing import Generator, Literal import mlx.core as mx -from exo_core.types.runner_response import ( - ImageGenerationResponse, - PartialImageResponse, -) -from exo_core.utils.memory import Memory -from PIL import Image - from exo_core.types.image_generation import ( AdvancedImageParams, ImageEditsTaskParams, ImageGenerationTaskParams, ImageSize, ) -from exo_core.types.runner_response import ImageGenerationStats +from exo_core.types.runner_response import ( + ImageGenerationResponse, + ImageGenerationStats, + PartialImageResponse, +) +from exo_core.utils.memory import Memory +from PIL import Image + from exo.worker.engines.image.distributed_model import DistributedImageModel diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py index f92ba490..cb5a8b20 100644 --- a/src/exo/worker/main.py +++ b/src/exo/worker/main.py @@ -3,8 +3,9 @@ from datetime import datetime, timezone import anyio from anyio import fail_after -from exo_core.types.common import CommandId, NodeId, SystemId, ModelId +from exo_core.types.common import CommandId, ModelId, NodeId, SystemId from exo_core.types.downloads import DownloadCompleted +from exo_core.types.image_generation import ImageEditsTaskParams from exo_core.types.runners import RunnerId from exo_core.types.tasks import ( CancelTask, @@ -15,10 +16,10 @@ from exo_core.types.tasks import ( Task, TaskStatus, ) +from exo_core.utils.channels import Receiver, Sender, channel from exo_core.utils.downloads import resolve_model_in_path from loguru import logger -from exo_core.types.image_generation import ImageEditsTaskParams from exo.shared.apply import apply from exo.shared.types.commands import ( ForwarderCommand, @@ -39,7 +40,6 @@ from exo.shared.types.events import ( from exo.shared.types.multiaddr import Multiaddr from exo.shared.types.state import State from exo.shared.types.topology import Connection, SocketConnection -from exo.utils.channels import Receiver, Sender, channel from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer from exo.utils.info_gatherer.net_profile import check_reachable from exo.utils.keyed_backoff import KeyedBackoff diff --git a/src/exo/worker/runner/bootstrap.py b/src/exo/worker/runner/bootstrap.py index 6e9e235d..372ed7a7 100644 --- a/src/exo/worker/runner/bootstrap.py +++ b/src/exo/worker/runner/bootstrap.py @@ -5,13 +5,12 @@ import sys from pathlib import Path import loguru -from exo_core.constants import EXO_MODELS_DIR from exo_core.types.instances import BoundInstance, VllmInstance from exo_core.types.runners import RunnerFailed from exo_core.types.tasks import Task, TaskId +from exo_core.utils.channels import ClosedResourceError, MpReceiver, MpSender from exo.shared.types.events import Event, RunnerStatusUpdated -from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender _CUDA_HOST_LIBS = ["libcuda.so.1", "libnvidia-ml.so.1", "libnvidia-ptxjitcompiler.so.1"] _CUDA_HOST_SEARCH_DIRS = [ @@ -73,6 +72,7 @@ def entrypoint( _ensure_cuda_libs() from vllm_engine.builder import VllmBuilder + from .llm_inference.runner import Runner builder = VllmBuilder.create( @@ -92,9 +92,10 @@ def entrypoint( ) runner.main() else: - from .llm_inference.runner import Runner from mlx_engine.builder import MlxBuilder + from .llm_inference.runner import Runner + builder = MlxBuilder.create( bound_instance, event_sender=event_sender, diff --git a/src/exo/worker/runner/image_models/runner.py b/src/exo/worker/runner/image_models/runner.py index 4cf0e9d5..e8f73b05 100644 --- a/src/exo/worker/runner/image_models/runner.py +++ b/src/exo/worker/runner/image_models/runner.py @@ -10,6 +10,7 @@ from exo_core.types.common import CommandId, ModelId from exo_core.types.instances import BoundInstance from exo_core.types.runner_response import ( ImageGenerationResponse, + ImageGenerationStats, PartialImageResponse, ) from exo_core.types.runners import ( @@ -43,8 +44,12 @@ from exo_core.types.tasks import ( TaskId, TaskStatus, ) +from exo_core.utils.channels import MpReceiver, MpSender +from loguru import logger +from mlx_engine.utils_mlx import ( + initialize_mlx, +) -from exo_core.types.runner_response import ImageGenerationStats from exo.shared.tracing import clear_trace_buffer, get_trace_buffer from exo.shared.types.events import ( ChunkGenerated, @@ -55,17 +60,12 @@ from exo.shared.types.events import ( TraceEventData, TracesCollected, ) -from exo.utils.channels import MpReceiver, MpSender from exo.worker.engines.image import ( DistributedImageModel, generate_image, initialize_image_model, warmup_image_generator, ) -from mlx_engine.utils_mlx import ( - initialize_mlx, -) -from loguru import logger def _is_primary_output_node(shard_metadata: ShardMetadata) -> bool: diff --git a/src/exo/worker/runner/llm_inference/runner.py b/src/exo/worker/runner/llm_inference/runner.py index 79ff9259..237c6508 100644 --- a/src/exo/worker/runner/llm_inference/runner.py +++ b/src/exo/worker/runner/llm_inference/runner.py @@ -4,6 +4,7 @@ from enum import Enum from typing import TYPE_CHECKING from anyio import WouldBlock +from exo_core.engine import Cancelled, Engine, EngineBuilder, Finished from exo_core.model_cards import ModelTask from exo_core.types.chunks import ( ErrorChunk, @@ -40,6 +41,8 @@ from exo_core.types.tasks import ( TaskStatus, TextGeneration, ) +from exo_core.utils.channels import MpReceiver, MpSender +from loguru import logger from exo.shared.types.events import ( ChunkGenerated, @@ -48,12 +51,6 @@ from exo.shared.types.events import ( TaskAcknowledged, TaskStatusUpdated, ) -from exo.utils.channels import MpReceiver, MpSender -from loguru import logger - -from .batch_generator import Cancelled, Finished - -from exo_core.engine import EngineBuilder, Engine try: from vllm_engine.builder import VllmBuilder @@ -74,7 +71,9 @@ class ExitCode(str, Enum): Shutdown = "Shutdown" -BuilderType = EngineBuilder[BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse] +BuilderType = EngineBuilder[ + BoundInstance, TextGeneration, GenerationResponse | ToolCallResponse +] EngineType = Engine[TextGeneration, GenerationResponse | ToolCallResponse] @@ -220,8 +219,9 @@ class Runner: self.update_status(RunnerLoaded()) logger.info("runner loaded") - case StartWarmup() if isinstance(self.current_status, RunnerLoaded) and isinstance(self.generator, Engine): - assert isinstance(self.generator, InferenceGenerator) + case StartWarmup() if isinstance( + self.current_status, RunnerLoaded + ) and isinstance(self.generator, Engine): logger.info("runner warming up") self.update_status(RunnerWarmingUp()) diff --git a/src/exo/worker/runner/runner_supervisor.py b/src/exo/worker/runner/runner_supervisor.py index 7e83224e..59aafc9d 100644 --- a/src/exo/worker/runner/runner_supervisor.py +++ b/src/exo/worker/runner/runner_supervisor.py @@ -32,6 +32,7 @@ from exo_core.types.tasks import ( TaskStatus, TextGeneration, ) +from exo_core.utils.channels import MpReceiver, MpSender, Sender, mp_channel from loguru import logger from exo.shared.types.events import ( @@ -41,7 +42,6 @@ from exo.shared.types.events import ( TaskAcknowledged, TaskStatusUpdated, ) -from exo.utils.channels import MpReceiver, MpSender, Sender, mp_channel from exo.utils.task_group import TaskGroup from exo.worker.runner.bootstrap import entrypoint diff --git a/src/exo/worker/tests/unittests/test_mlx/test_auto_parallel.py b/src/exo/worker/tests/unittests/test_mlx/test_auto_parallel.py index 850db320..5eea67b9 100644 --- a/src/exo/worker/tests/unittests/test_mlx/test_auto_parallel.py +++ b/src/exo/worker/tests/unittests/test_mlx/test_auto_parallel.py @@ -7,13 +7,13 @@ from typing import Any import mlx.core as mx import mlx.nn as mlx_nn import pytest - -from exo.worker.engines.mlx.auto_parallel import ( +from mlx_engine.auto_parallel import ( CustomMlxLayer, PipelineFirstLayer, PipelineLastLayer, patch_pipeline_model, ) + from exo.worker.tests.unittests.test_mlx.conftest import MockLayer diff --git a/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py b/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py index d9946cc7..4fa23912 100644 --- a/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py +++ b/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py @@ -14,15 +14,14 @@ import pytest from exo_core.types.common import ModelId from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams from mlx.utils import tree_flatten, tree_unflatten -from mlx_lm.tokenizer_utils import TokenizerWrapper - -from exo.shared.types.mlx import Model -from exo.worker.engines.mlx.cache import KVPrefixCache -from exo.worker.engines.mlx.generator.generate import mlx_generate -from exo.worker.engines.mlx.utils_mlx import ( +from mlx_engine.cache import KVPrefixCache +from mlx_engine.generator.generate import mlx_generate +from mlx_engine.types import Model +from mlx_engine.utils_mlx import ( apply_chat_template, load_tokenizer_for_model_id, ) +from mlx_lm.tokenizer_utils import TokenizerWrapper HF_CACHE = Path.home() / ".cache" / "huggingface" / "hub" diff --git a/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py b/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py index c2c2b4d9..0fc023a7 100644 --- a/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py +++ b/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py @@ -17,8 +17,7 @@ from exo_core.utils.downloads import ( ensure_models_dir, fetch_file_list_with_cache, ) - -from exo.worker.engines.mlx.utils_mlx import ( +from mlx_engine.utils_mlx import ( get_eos_token_ids_for_model, load_tokenizer_for_model_id, ) diff --git a/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py b/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py index d9145154..b0e4d222 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py +++ b/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py @@ -6,7 +6,6 @@ from exo_core.types.runner_response import ( GenerationResponse, ToolCallResponse, ) - from mlx_engine.dsml_encoding import ( ASSISTANT_TOKEN, BOS_TOKEN, @@ -20,6 +19,7 @@ from mlx_engine.dsml_encoding import ( encode_messages, parse_dsml_output, ) + from exo.worker.runner.llm_inference.model_output_parsers import parse_deepseek_v32 # ── Shared fixtures ────────────────────────────────────────────── diff --git a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py index dbb072fa..43b0f4eb 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py +++ b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py @@ -30,6 +30,7 @@ from exo_core.types.tasks import ( TextGeneration, ) from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams +from exo_core.utils.channels import mp_channel import exo.worker.runner.llm_inference.batch_generator as mlx_batch_generator import exo.worker.runner.llm_inference.model_output_parsers as mlx_model_output_parsers @@ -41,7 +42,6 @@ from exo.shared.types.events import ( TaskAcknowledged, TaskStatusUpdated, ) -from exo.utils.channels import mp_channel from ...constants import ( CHAT_COMPLETION_TASK_ID, diff --git a/src/exo/worker/tests/unittests/test_runner/test_parse_tool_calls.py b/src/exo/worker/tests/unittests/test_runner/test_parse_tool_calls.py index cd9dc843..6c0d4730 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_parse_tool_calls.py +++ b/src/exo/worker/tests/unittests/test_runner/test_parse_tool_calls.py @@ -5,9 +5,9 @@ from collections.abc import Generator from typing import Any from exo_core.types.runner_response import GenerationResponse, ToolCallResponse +from exo_core.utils.tool_parsers import make_mlx_parser from exo.worker.runner.llm_inference.model_output_parsers import parse_tool_calls -from exo.worker.runner.llm_inference.tool_parsers import make_mlx_parser def _make_responses( diff --git a/src/exo/worker/tests/unittests/test_runner/test_runner_supervisor.py b/src/exo/worker/tests/unittests/test_runner/test_runner_supervisor.py index 6b747aa3..ef0dcffe 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_runner_supervisor.py +++ b/src/exo/worker/tests/unittests/test_runner/test_runner_supervisor.py @@ -9,9 +9,9 @@ from exo_core.types.instances import BoundInstance, InstanceId from exo_core.types.runners import RunnerFailed, RunnerId from exo_core.types.tasks import Task, TaskId, TextGeneration from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams +from exo_core.utils.channels import channel, mp_channel from exo.shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated -from exo.utils.channels import channel, mp_channel from exo.worker.runner.runner_supervisor import RunnerSupervisor from exo.worker.tests.unittests.conftest import get_bound_mlx_ring_instance diff --git a/tests/headless_runner.py b/tests/headless_runner.py index 20f73525..9498770c 100644 --- a/tests/headless_runner.py +++ b/tests/headless_runner.py @@ -29,6 +29,7 @@ from exo_core.types.tasks import ( TextGeneration, ) from exo_core.types.text_generation import InputMessage, TextGenerationTaskParams +from exo_core.utils.channels import channel, mp_channel from fastapi import FastAPI from fastapi.responses import Response, StreamingResponse from hypercorn import Config @@ -38,7 +39,6 @@ from pydantic import BaseModel from exo.shared.types.commands import CommandId from exo.shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated -from exo.utils.channels import channel, mp_channel from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer from exo.worker.runner.bootstrap import entrypoint diff --git a/uv.lock b/uv.lock index 4478b651..50941414 100644 --- a/uv.lock +++ b/uv.lock @@ -943,11 +943,17 @@ name = "exo-core" version = "0.1.0" source = { editable = "python/exo_core" } dependencies = [ + { name = "mlx-lm", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cuda' and extra == 'project-9-exo-bench')" }, + { name = "openai-harmony", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cuda' and extra == 'project-9-exo-bench')" }, { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cuda' and extra == 'project-9-exo-bench')" }, ] [package.metadata] -requires-dist = [{ name = "pydantic", specifier = ">=2.13.0b2" }] +requires-dist = [ + { name = "mlx-lm", git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Feval-left-padding-in-batched-rotation" }, + { name = "openai-harmony", specifier = ">=0.0.8" }, + { name = "pydantic", specifier = ">=2.13.0b2" }, +] [[package]] name = "exo-pyo3-bindings"