From acb97127bf1326d55f07d2301da6e6f19da19df5 Mon Sep 17 00:00:00 2001 From: Alex Cheema <41707476+AlexCheema@users.noreply.github.com> Date: Tue, 3 Feb 2026 06:01:56 -0800 Subject: [PATCH 1/9] Normalize TextGenerationTaskParams.input to list[InputMessage] (#1360) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Motivation With the addition of the Responses API, we introduced `str | list[InputMessage]` as the type for `TextGenerationTaskParams.input` since the Responses API supports sending input as a plain string. But there was no reason to leak that flexibility past the API adapter boundary — it just meant every downstream consumer had to do `if isinstance(messages, str):` checks, adding complexity for no benefit. ## Changes - Changed `TextGenerationTaskParams.input` from `str | list[InputMessage]` to `list[InputMessage]` - Each API adapter (Chat Completions, Claude Messages, Responses) now normalizes to `list[InputMessage]` at the boundary - Removed `isinstance(task_params.input, str)` branches in `utils_mlx.py` and `runner.py` - Wrapped string inputs in `[InputMessage(role="user", content=...)]` in the warmup path and all test files ## Why It Works The API adapters are the only place where we deal with raw user input formats. By normalizing there, all downstream code (worker, runner, MLX engine) can just assume `list[InputMessage]` and skip the type-checking branches. The type system (`basedpyright`) catches any missed call sites at compile time. ## Test Plan ### Automated Testing - `uv run basedpyright` — 0 errors - `uv run ruff check` — passes - `nix fmt` — applied - `uv run pytest` — 174 passed, 1 skipped Co-authored-by: Claude Opus 4.5 --- src/exo/master/adapters/chat_completions.py | 4 +++- src/exo/master/adapters/claude.py | 4 +++- src/exo/master/adapters/responses.py | 10 +++++++--- src/exo/master/tests/test_master.py | 8 +++++--- src/exo/shared/types/text_generation.py | 2 +- .../worker/engines/mlx/generator/generate.py | 4 ++-- src/exo/worker/engines/mlx/utils_mlx.py | 15 +++++---------- src/exo/worker/runner/runner.py | 13 ++++--------- .../tests/unittests/test_mlx/conftest.py | 6 +++--- .../test_plan/test_task_forwarding.py | 18 +++++++++++++----- .../test_runner/test_event_ordering.py | 4 ++-- tests/headless_runner.py | 8 ++++++-- 12 files changed, 54 insertions(+), 42 deletions(-) diff --git a/src/exo/master/adapters/chat_completions.py b/src/exo/master/adapters/chat_completions.py index 5a27664d..e144696b 100644 --- a/src/exo/master/adapters/chat_completions.py +++ b/src/exo/master/adapters/chat_completions.py @@ -66,7 +66,9 @@ def chat_request_to_text_generation( return TextGenerationTaskParams( model=request.model, - input=input_messages if input_messages else "", + input=input_messages + if input_messages + else [InputMessage(role="user", content="")], instructions=instructions, max_output_tokens=request.max_tokens, temperature=request.temperature, diff --git a/src/exo/master/adapters/claude.py b/src/exo/master/adapters/claude.py index 13398012..6c17b49c 100644 --- a/src/exo/master/adapters/claude.py +++ b/src/exo/master/adapters/claude.py @@ -141,7 +141,9 @@ def claude_request_to_text_generation( return TextGenerationTaskParams( model=request.model, - input=input_messages if input_messages else "", + input=input_messages + if input_messages + else [InputMessage(role="user", content="")], instructions=instructions, max_output_tokens=request.max_tokens, temperature=request.temperature, diff --git a/src/exo/master/adapters/responses.py b/src/exo/master/adapters/responses.py index 27d845cd..c2a416ac 100644 --- a/src/exo/master/adapters/responses.py +++ b/src/exo/master/adapters/responses.py @@ -43,10 +43,10 @@ def _extract_content(content: str | list[ResponseContentPart]) -> str: def responses_request_to_text_generation( request: ResponsesRequest, ) -> TextGenerationTaskParams: - input_value: str | list[InputMessage] + input_value: list[InputMessage] built_chat_template: list[dict[str, Any]] | None = None if isinstance(request.input, str): - input_value = request.input + input_value = [InputMessage(role="user", content=request.input)] else: input_messages: list[InputMessage] = [] chat_template_messages: list[dict[str, Any]] = [] @@ -95,7 +95,11 @@ def responses_request_to_text_generation( } ) - input_value = input_messages if input_messages else "" + input_value = ( + input_messages + if input_messages + else [InputMessage(role="user", content="")] + ) built_chat_template = chat_template_messages if chat_template_messages else None return TextGenerationTaskParams( diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py index d9987727..ddf9aec8 100644 --- a/src/exo/master/tests/test_master.py +++ b/src/exo/master/tests/test_master.py @@ -28,7 +28,7 @@ from exo.shared.types.profiling import ( ) from exo.shared.types.tasks import TaskStatus from exo.shared.types.tasks import TextGeneration as TextGenerationTask -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import ( InstanceMeta, MlxRingInstance, @@ -136,7 +136,9 @@ async def test_master(): command_id=CommandId(), task_params=TextGenerationTaskParams( model=ModelId("llama-3.2-1b"), - input="Hello, how are you?", + input=[ + InputMessage(role="user", content="Hello, how are you?") + ], ), ) ), @@ -189,7 +191,7 @@ async def test_master(): assert isinstance(events[2].event.task, TextGenerationTask) assert events[2].event.task.task_params == TextGenerationTaskParams( model=ModelId("llama-3.2-1b"), - input="Hello, how are you?", + input=[InputMessage(role="user", content="Hello, how are you?")], ) await master.shutdown() diff --git a/src/exo/shared/types/text_generation.py b/src/exo/shared/types/text_generation.py index b9c5565c..31f97a70 100644 --- a/src/exo/shared/types/text_generation.py +++ b/src/exo/shared/types/text_generation.py @@ -28,7 +28,7 @@ class TextGenerationTaskParams(BaseModel, frozen=True): """ model: ModelId - input: str | list[InputMessage] + input: list[InputMessage] instructions: str | None = None max_output_tokens: int | None = None temperature: float | None = None diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index c9585f57..a38d70c5 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -17,7 +17,7 @@ from exo.shared.types.api import ( from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory from exo.shared.types.mlx import KVCacheType -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.runner_response import ( GenerationResponse, ) @@ -100,7 +100,7 @@ def warmup_inference( tokenizer=tokenizer, task_params=TextGenerationTaskParams( model=ModelId(""), - input=content, + input=[InputMessage(role="user", content=content)], ), ) diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index de5cb190..d7fb9958 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -436,16 +436,11 @@ def apply_chat_template( ) # Convert input to messages - if isinstance(task_params.input, str): - # Simple string input becomes a single user message - formatted_messages.append({"role": "user", "content": task_params.input}) - else: - # List of InputMessage - for msg in task_params.input: - if not msg.content: - logger.warning("Received message with empty content, skipping") - continue - formatted_messages.append({"role": msg.role, "content": msg.content}) + for msg in task_params.input: + if not msg.content: + logger.warning("Received message with empty content, skipping") + continue + formatted_messages.append({"role": msg.role, "content": msg.content}) prompt: str = tokenizer.apply_chat_template( formatted_messages, diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 5439b72b..0b527318 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -918,15 +918,10 @@ def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None: Extracts the first user input text and checks for debug triggers. """ - prompt: str - if isinstance(task_params.input, str): - prompt = task_params.input - else: - # List of InputMessage - get first message content - if len(task_params.input) == 0: - logger.debug("Empty message list in debug prompt check") - return - prompt = task_params.input[0].content + if len(task_params.input) == 0: + logger.debug("Empty message list in debug prompt check") + return + prompt = task_params.input[0].content if not prompt: return diff --git a/src/exo/worker/tests/unittests/test_mlx/conftest.py b/src/exo/worker/tests/unittests/test_mlx/conftest.py index 7015d70b..9e897141 100644 --- a/src/exo/worker/tests/unittests/test_mlx/conftest.py +++ b/src/exo/worker/tests/unittests/test_mlx/conftest.py @@ -14,7 +14,7 @@ from exo.shared.constants import EXO_MODELS_DIR from exo.shared.models.model_cards import ModelCard, ModelTask from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.generator.generate import mlx_generate @@ -114,7 +114,7 @@ def run_gpt_oss_pipeline_device( task = TextGenerationTaskParams( model=DEFAULT_GPT_OSS_MODEL_ID, - input=prompt_text, + input=[InputMessage(role="user", content=prompt_text)], max_output_tokens=max_tokens, ) @@ -182,7 +182,7 @@ def run_gpt_oss_tensor_parallel_device( task = TextGenerationTaskParams( model=DEFAULT_GPT_OSS_MODEL_ID, - input=prompt_text, + input=[InputMessage(role="user", content=prompt_text)], max_output_tokens=max_tokens, ) diff --git a/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py b/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py index 07376787..64be74e6 100644 --- a/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py +++ b/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py @@ -2,7 +2,7 @@ from typing import cast import exo.worker.plan as plan_mod from exo.shared.types.tasks import Task, TaskId, TaskStatus, TextGeneration -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import BoundInstance, InstanceId from exo.shared.types.worker.runners import ( RunnerIdle, @@ -59,7 +59,9 @@ def test_plan_forwards_pending_chat_completion_when_runner_ready(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -106,7 +108,9 @@ def test_plan_does_not_forward_chat_completion_if_any_runner_not_ready(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -150,7 +154,9 @@ def test_plan_does_not_forward_tasks_for_other_instances(): instance_id=other_instance_id, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -198,7 +204,9 @@ def test_plan_ignores_non_pending_or_non_chat_tasks(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Complete, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) other_task_id = TaskId("other-task") 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 9d7703d8..16a43f2b 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 @@ -22,7 +22,7 @@ from exo.shared.types.tasks import ( TaskStatus, TextGeneration, ) -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.runner_response import GenerationResponse from exo.shared.types.worker.runners import ( RunnerConnected, @@ -86,7 +86,7 @@ SHUTDOWN_TASK = Shutdown( CHAT_PARAMS = TextGenerationTaskParams( model=MODEL_A_ID, - input="hello", + input=[InputMessage(role="user", content="hello")], stream=True, max_output_tokens=4, temperature=0.0, diff --git a/tests/headless_runner.py b/tests/headless_runner.py index ed57823b..56fb2632 100644 --- a/tests/headless_runner.py +++ b/tests/headless_runner.py @@ -23,7 +23,7 @@ from exo.shared.types.tasks import ( Task, TextGeneration, ) -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import ( BoundInstance, Instance, @@ -196,7 +196,11 @@ async def execute_test(test: Tests, instance: Instance, hn: str) -> list[Event]: task_params=TextGenerationTaskParams( model=test.model_id, instructions="You are a helpful assistant", - input="What is the capital of France?", + input=[ + InputMessage( + role="user", content="What is the capital of France?" + ) + ], ), command_id=CommandId("yo"), instance_id=iid, From a0f4f363555744f2a9660679437be064bb2bb712 Mon Sep 17 00:00:00 2001 From: rltakashige Date: Tue, 3 Feb 2026 20:03:29 +0000 Subject: [PATCH 2/9] Reduce reliance on internet (#1363) ## Motivation Offline users currently have to wait for every retry to fail before being able to launch a model. For users that restart clusters often or share API keys between devices, we also spam HuggingFace with downloads every 5 minutes. These issues are caused by _emit_existing_download_progress being inefficient. ## Changes - Only query HuggingFace once while EXO is running (assumption being that a change should only be reflected on a new EXO session) - Only query HuggingFace when there is an internet connection (polling connectivity every 10 seconds) - Request download progress if we switch from no connectivity -> connected to reduce the wait. - Reduce download progress sleep as it's no longer expensive (queries cache most of the time). - Reduce retries as 30 is way too many. ## Test Plan ### Manual Testing Manually tested the behaviour. ### Automated Testing None, should I add any? We do have some tests for this folder, but they are probably not too helpful. --- src/exo/download/coordinator.py | 34 +++++++++- src/exo/download/download_utils.py | 82 ++++++++++++++++++----- src/exo/download/impl_shard_downloader.py | 28 +++++++- src/exo/download/shard_downloader.py | 5 ++ src/exo/shared/models/model_cards.py | 8 +-- 5 files changed, 130 insertions(+), 27 deletions(-) diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py index c2f7b9e9..f5798ad3 100644 --- a/src/exo/download/coordinator.py +++ b/src/exo/download/coordinator.py @@ -1,4 +1,5 @@ import asyncio +import socket from dataclasses import dataclass, field from typing import Iterator @@ -60,10 +61,37 @@ class DownloadCoordinator: async def run(self) -> None: logger.info("Starting DownloadCoordinator") + self._test_internet_connection() async with self._tg as tg: tg.start_soon(self._command_processor) tg.start_soon(self._forward_events) tg.start_soon(self._emit_existing_download_progress) + tg.start_soon(self._check_internet_connection) + + def _test_internet_connection(self) -> None: + try: + socket.create_connection(("1.1.1.1", 443), timeout=3).close() + self.shard_downloader.set_internet_connection(True) + except OSError: + self.shard_downloader.set_internet_connection(False) + logger.debug( + f"Internet connectivity: {self.shard_downloader.internet_connection}" + ) + + async def _check_internet_connection(self) -> None: + first_connection = True + while True: + await asyncio.sleep(10) + + # Assume that internet connection is set to False on 443 errors. + if self.shard_downloader.internet_connection: + continue + + self._test_internet_connection() + + if first_connection and self.shard_downloader.internet_connection: + first_connection = False + self._tg.start_soon(self._emit_existing_download_progress) def shutdown(self) -> None: self._tg.cancel_scope.cancel() @@ -241,7 +269,7 @@ class DownloadCoordinator: async def _emit_existing_download_progress(self) -> None: try: while True: - logger.info( + logger.debug( "DownloadCoordinator: Fetching and emitting existing download progress..." ) async for ( @@ -274,10 +302,10 @@ class DownloadCoordinator: await self.event_sender.send( NodeDownloadProgress(download_progress=status) ) - logger.info( + logger.debug( "DownloadCoordinator: Done emitting existing download progress." ) - await anyio.sleep(5 * 60) # 5 minutes + await anyio.sleep(60) except Exception as e: logger.error( f"DownloadCoordinator: Error emitting existing download progress: {e}" diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py index 6dec4718..618e4f38 100644 --- a/src/exo/download/download_utils.py +++ b/src/exo/download/download_utils.py @@ -49,6 +49,10 @@ class HuggingFaceAuthenticationError(Exception): """Raised when HuggingFace returns 401/403 for a model download.""" +class HuggingFaceRateLimitError(Exception): + """429 Huggingface code""" + + async def _build_auth_error_message(status_code: int, model_id: ModelId) -> str: token = await get_hf_token() if status_code == 401 and token is None: @@ -154,49 +158,76 @@ async def seed_models(seed_dir: str | Path): logger.error(traceback.format_exc()) +_fetched_file_lists_this_session: set[str] = set() + + async def fetch_file_list_with_cache( - model_id: ModelId, revision: str = "main", recursive: bool = False + model_id: ModelId, + revision: str = "main", + recursive: bool = False, + skip_internet: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: target_dir = (await ensure_models_dir()) / "caches" / model_id.normalize() await aios.makedirs(target_dir, exist_ok=True) cache_file = target_dir / f"{model_id.normalize()}--{revision}--file_list.json" + cache_key = f"{model_id.normalize()}--{revision}" + + if cache_key in _fetched_file_lists_this_session and await aios.path.exists( + cache_file + ): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + + if skip_internet: + if await aios.path.exists(cache_file): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + raise FileNotFoundError( + f"No internet connection and no cached file list for {model_id}" + ) - # Always try fresh first try: file_list = await fetch_file_list_with_retry( - model_id, revision, recursive=recursive + model_id, + revision, + recursive=recursive, + on_connection_lost=on_connection_lost, ) - # Update cache with fresh data async with aiofiles.open(cache_file, "w") as f: await f.write( TypeAdapter(list[FileListEntry]).dump_json(file_list).decode() ) + _fetched_file_lists_this_session.add(cache_key) return file_list except Exception as e: - # Fetch failed - try cache fallback if await aios.path.exists(cache_file): logger.warning( f"Failed to fetch file list for {model_id}, using cached data: {e}" ) async with aiofiles.open(cache_file, "r") as f: return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) - # No cache available, propagate the error - raise + raise FileNotFoundError(f"Failed to fetch file list for {model_id}: {e}") from e async def fetch_file_list_with_retry( - model_id: ModelId, revision: str = "main", path: str = "", recursive: bool = False + model_id: ModelId, + revision: str = "main", + path: str = "", + recursive: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _fetch_file_list(model_id, revision, path, recursive) except HuggingFaceAuthenticationError: raise except Exception as e: + on_connection_lost() if attempt == n_attempts - 1: raise e - await asyncio.sleep(min(8, 0.1 * float(2.0 ** int(attempt)))) + await asyncio.sleep(2.0**attempt) raise Exception( f"Failed to fetch file list for {model_id=} {revision=} {path=} {recursive=}" ) @@ -216,7 +247,11 @@ async def _fetch_file_list( if response.status in [401, 403]: msg = await _build_auth_error_message(response.status, model_id) raise HuggingFaceAuthenticationError(msg) - if response.status == 200: + elif response.status == 429: + raise HuggingFaceRateLimitError( + f"Couldn't download {model_id} because of HuggingFace rate limit." + ) + elif response.status == 200: data_json = await response.text() data = TypeAdapter(list[FileListEntry]).validate_json(data_json) files: list[FileListEntry] = [] @@ -249,7 +284,7 @@ def create_http_session( else: total_timeout = 1800 connect_timeout = 60 - sock_read_timeout = 1800 + sock_read_timeout = 60 sock_connect_timeout = 60 ssl_context = ssl.create_default_context( @@ -324,8 +359,9 @@ async def download_file_with_retry( path: str, target_dir: Path, on_progress: Callable[[int, int, bool], None] = lambda _, __, ___: None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> Path: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _download_file( @@ -333,14 +369,19 @@ async def download_file_with_retry( ) except HuggingFaceAuthenticationError: raise - except Exception as e: - if isinstance(e, FileNotFoundError) or attempt == n_attempts - 1: + except HuggingFaceRateLimitError as e: + if attempt == n_attempts - 1: raise e logger.error( f"Download error on attempt {attempt}/{n_attempts} for {model_id=} {revision=} {path=} {target_dir=}" ) logger.error(traceback.format_exc()) - await asyncio.sleep(min(8, 0.1 * (2.0**attempt))) + await asyncio.sleep(2.0**attempt) + except Exception as e: + on_connection_lost() + if attempt == n_attempts - 1: + raise e + break raise Exception( f"Failed to download file {model_id=} {revision=} {path=} {target_dir=}" ) @@ -542,7 +583,9 @@ async def download_shard( on_progress: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], max_parallel_downloads: int = 8, skip_download: bool = False, + skip_internet: bool = False, allow_patterns: list[str] | None = None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> tuple[Path, RepoDownloadProgress]: if not skip_download: logger.debug(f"Downloading {shard.model_card.model_id=}") @@ -562,7 +605,11 @@ async def download_shard( all_start_time = time.time() file_list = await fetch_file_list_with_cache( - shard.model_card.model_id, revision, recursive=True + shard.model_card.model_id, + revision, + recursive=True, + skip_internet=skip_internet, + on_connection_lost=on_connection_lost, ) filtered_file_list = list( filter_repo_objects( @@ -672,6 +719,7 @@ async def download_shard( lambda curr_bytes, total_bytes, is_renamed: schedule_progress( file, curr_bytes, total_bytes, is_renamed ), + on_connection_lost=on_connection_lost, ) if not skip_download: diff --git a/src/exo/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py index 1b7f5eab..0e7aea1e 100644 --- a/src/exo/download/impl_shard_downloader.py +++ b/src/exo/download/impl_shard_downloader.py @@ -1,4 +1,5 @@ import asyncio +from asyncio import create_task from collections.abc import Awaitable from pathlib import Path from typing import AsyncIterator, Callable @@ -49,6 +50,10 @@ class SingletonShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.active_downloads: dict[ShardMetadata, asyncio.Task[Path]] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -85,6 +90,10 @@ class CachedShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.cache: dict[tuple[str, ShardMetadata], Path] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -142,6 +151,8 @@ class ResumableShardDownloader(ShardDownloader): self.on_progress_wrapper, max_parallel_downloads=self.max_parallel_downloads, allow_patterns=allow_patterns, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) return target_dir @@ -154,12 +165,23 @@ class ResumableShardDownloader(ShardDownloader): """Helper coroutine that builds the shard for a model and gets its download status.""" shard = await build_full_shard(model_id) return await download_shard( - shard, self.on_progress_wrapper, skip_download=True + shard, + self.on_progress_wrapper, + skip_download=True, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) - # Kick off download status coroutines concurrently + semaphore = asyncio.Semaphore(self.max_parallel_downloads) + + async def download_with_semaphore( + model_card: ModelCard, + ) -> tuple[Path, RepoDownloadProgress]: + async with semaphore: + return await _status_for_model(model_card.model_id) + tasks = [ - asyncio.create_task(_status_for_model(model_card.model_id)) + create_task(download_with_semaphore(model_card)) for model_card in await get_model_cards() ] diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py index 30c11d25..9dd8c324 100644 --- a/src/exo/download/shard_downloader.py +++ b/src/exo/download/shard_downloader.py @@ -16,6 +16,11 @@ from exo.shared.types.worker.shards import ( # TODO: the PipelineShardMetadata getting reinstantiated is a bit messy. Should this be a classmethod? class ShardDownloader(ABC): + internet_connection: bool = False + + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + @abstractmethod async def ensure_shard( self, shard: ShardMetadata, config_only: bool = False diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index bf2a892f..58d84849 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -108,9 +108,9 @@ class ModelCard(CamelCaseModel): async def fetch_from_hf(model_id: ModelId) -> "ModelCard": """Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta.""" # TODO: failure if files do not exist - config_data = await get_config_data(model_id) + config_data = await fetch_config_data(model_id) num_layers = config_data.layer_count - mem_size_bytes = await get_safetensors_size(model_id) + mem_size_bytes = await fetch_safetensors_size(model_id) mc = ModelCard( model_id=ModelId(model_id), @@ -258,7 +258,7 @@ class ConfigData(BaseModel): return data -async def get_config_data(model_id: ModelId) -> ConfigData: +async def fetch_config_data(model_id: ModelId) -> ConfigData: """Downloads and parses config.json for a model.""" from exo.download.download_utils import ( download_file_with_retry, @@ -280,7 +280,7 @@ async def get_config_data(model_id: ModelId) -> ConfigData: return ConfigData.model_validate_json(await f.read()) -async def get_safetensors_size(model_id: ModelId) -> Memory: +async def fetch_safetensors_size(model_id: ModelId) -> Memory: """Gets model size from safetensors index or falls back to HF API.""" from exo.download.download_utils import ( download_file_with_retry, From 2063278906d96e9017c48d53d54d9e1631b6cda9 Mon Sep 17 00:00:00 2001 From: Alex Cheema <41707476+AlexCheema@users.noreply.github.com> Date: Wed, 4 Feb 2026 05:06:15 -0800 Subject: [PATCH 3/9] feat: add custom HuggingFace model support (#1368) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Motivation Users should be able to run any HuggingFace model, not just the ones we ship TOML cards for. Continues the aim of #1191 with a minimal implementation on top of the current TOML model card system. Custom cards are saved to `~/.exo/custom_model_cards/` rather than the bundled `resources/inference_model_cards/` because `RESOURCES_DIR` is read-only in PyInstaller bundles (`sys._MEIPASS`). This also fixes `fetch_from_hf` which was saving cards to the wrong path (`resources/` root instead of `resources/inference_model_cards/`). ## Changes - Add `EXO_CUSTOM_MODEL_CARDS_DIR` constant (`~/.exo/custom_model_cards/`) - Update `model_cards.py`: add custom dir to search path, fix `save_to_custom_dir`, add `delete_custom_card`/`is_custom_card` - Add `POST /models/add` and `DELETE /models/custom/{model_id}` API endpoints - Add `is_custom` field to `ModelListModel` API response - Dashboard: add custom model input form in dropdown, delete button for custom models, show actual API errors, auto-select newly added model ## Why It Works Two separate directories for model cards: the bundled read-only `resources/inference_model_cards/` for built-in cards, and user-writable `~/.exo/custom_model_cards/` for custom cards. Both are scanned when listing models. This works in all environments including PyInstaller bundles where `RESOURCES_DIR` points to `sys._MEIPASS`. ## Test Plan ### Manual Testing - Add a custom model via the dropdown (e.g. `mlx-community/Llama-3.2-1B-Instruct-4bit`) - Verify it appears in the model list with the delete (x) button - Delete it and verify it disappears - Try adding an invalid model ID and verify the actual error is shown ### Automated Testing - `uv run basedpyright` — 0 errors - `uv run ruff check` — passes - `uv run pytest src/` — passes - `cd dashboard && npm run build` — builds --------- Co-authored-by: Claude Opus 4.5 --- dashboard/src/routes/+page.svelte | 126 +++++++++++++++++++++++++-- src/exo/master/api.py | 36 ++++++++ src/exo/shared/constants.py | 2 + src/exo/shared/models/model_cards.py | 36 ++++++-- src/exo/shared/types/api.py | 5 ++ 5 files changed, 191 insertions(+), 14 deletions(-) diff --git a/dashboard/src/routes/+page.svelte b/dashboard/src/routes/+page.svelte index 183b6547..6ecce7cc 100644 --- a/dashboard/src/routes/+page.svelte +++ b/dashboard/src/routes/+page.svelte @@ -100,6 +100,7 @@ storage_size_megabytes?: number; tasks?: string[]; hugging_face_id?: string; + is_custom?: boolean; }> >([]); @@ -215,6 +216,11 @@ let isModelDropdownOpen = $state(false); let modelDropdownSearch = $state(""); + // Custom model add state + let customModelInput = $state(""); + let isAddingCustomModel = $state(false); + let customModelError = $state(""); + // Slider dragging state let isDraggingSlider = $state(false); let sliderTrackElement: HTMLDivElement | null = $state(null); @@ -530,6 +536,57 @@ } } + async function addCustomModel() { + const modelId = customModelInput.trim(); + if (!modelId) return; + + isAddingCustomModel = true; + customModelError = ""; + + try { + const response = await fetch("/models/add", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ model_id: modelId }), + }); + + if (!response.ok) { + try { + const err = await response.json(); + customModelError = + err.detail || `Failed to add model (${response.status})`; + } catch { + customModelError = `Failed to add model (${response.status}: ${response.statusText})`; + } + return; + } + + const added = await response.json(); + customModelInput = ""; + await fetchModels(); + selectPreviewModel(added.id); + isModelDropdownOpen = false; + } catch { + customModelError = "Network error"; + } finally { + isAddingCustomModel = false; + } + } + + async function deleteCustomModel(modelId: string) { + try { + const response = await fetch( + `/models/custom/${encodeURIComponent(modelId)}`, + { method: "DELETE" }, + ); + if (response.ok) { + await fetchModels(); + } + } catch { + console.error("Failed to delete custom model"); + } + } + async function launchInstance( modelId: string, specificPreview?: PlacementPreview | null, @@ -2475,7 +2532,7 @@ >
+ +
{ + e.preventDefault(); + addCustomModel(); + }} + class="flex gap-1.5" + > + + +
+ {#if customModelError} +
+ {customModelError} +
+ {/if}
@@ -2557,14 +2642,37 @@ {/if} {model.name || model.id} - - {sizeGB >= 1 - ? sizeGB.toFixed(0) - : sizeGB.toFixed(1)}GB + + + {sizeGB >= 1 + ? sizeGB.toFixed(0) + : sizeGB.toFixed(1)}GB + + {#if model.is_custom} + + {/if} {:else} diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 8cf86725..01e145b4 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -50,10 +50,13 @@ from exo.shared.logging import InterceptLogger from exo.shared.models.model_cards import ( ModelCard, ModelId, + delete_custom_card, get_model_cards, + is_custom_card, ) from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file from exo.shared.types.api import ( + AddCustomModelParams, AdvancedImageParams, BenchChatCompletionRequest, BenchChatCompletionResponse, @@ -257,6 +260,8 @@ class API: self.app.delete("/instance/{instance_id}")(self.delete_instance) self.app.get("/models")(self.get_models) self.app.get("/v1/models")(self.get_models) + self.app.post("/models/add")(self.add_custom_model) + self.app.delete("/models/custom/{model_id:path}")(self.delete_custom_model) self.app.post("/v1/chat/completions", response_model=None)( self.chat_completions ) @@ -1216,11 +1221,42 @@ class API: storage_size_megabytes=int(card.storage_size.in_mb), supports_tensor=card.supports_tensor, tasks=[task.value for task in card.tasks], + is_custom=is_custom_card(card.model_id), ) for card in await get_model_cards() ] ) + async def add_custom_model(self, payload: AddCustomModelParams) -> ModelListModel: + """Fetch a model from HuggingFace and save as a custom model card.""" + try: + card = await ModelCard.fetch_from_hf(payload.model_id) + except Exception as exc: + raise HTTPException( + status_code=400, detail=f"Failed to fetch model: {exc}" + ) from exc + + return ModelListModel( + id=card.model_id, + hugging_face_id=card.model_id, + name=card.model_id.short(), + description="", + tags=[], + storage_size_megabytes=int(card.storage_size.in_mb), + supports_tensor=card.supports_tensor, + tasks=[task.value for task in card.tasks], + is_custom=True, + ) + + async def delete_custom_model(self, model_id: ModelId) -> JSONResponse: + """Delete a user-added custom model card.""" + deleted = await delete_custom_card(model_id) + if not deleted: + raise HTTPException(status_code=404, detail="Custom model card not found") + return JSONResponse( + {"message": "Model card deleted", "model_id": str(model_id)} + ) + async def run(self): cfg = Config() cfg.bind = f"0.0.0.0:{self.port}" diff --git a/src/exo/shared/constants.py b/src/exo/shared/constants.py index 39438cbd..b09d6b2d 100644 --- a/src/exo/shared/constants.py +++ b/src/exo/shared/constants.py @@ -58,6 +58,8 @@ LIBP2P_COMMANDS_TOPIC = "commands" EXO_MAX_CHUNK_SIZE = 512 * 1024 +EXO_CUSTOM_MODEL_CARDS_DIR = EXO_DATA_HOME / "custom_model_cards" + EXO_IMAGE_CACHE_DIR = EXO_CACHE_HOME / "images" EXO_TRACING_CACHE_DIR = EXO_CACHE_HOME / "traces" diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index 58d84849..47b54d38 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -18,14 +18,19 @@ from pydantic import ( ) from tomlkit.exceptions import TOMLKitError -from exo.shared.constants import EXO_ENABLE_IMAGE_MODELS, RESOURCES_DIR +from exo.shared.constants import ( + EXO_CUSTOM_MODEL_CARDS_DIR, + EXO_ENABLE_IMAGE_MODELS, + RESOURCES_DIR, +) from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory from exo.utils.pydantic_ext import CamelCaseModel # kinda ugly... # TODO: load search path from config.toml -_csp = [Path(RESOURCES_DIR) / "inference_model_cards"] +_custom_cards_dir = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR)) +_csp = [Path(RESOURCES_DIR) / "inference_model_cards", _custom_cards_dir] if EXO_ENABLE_IMAGE_MODELS: _csp.append(Path(RESOURCES_DIR) / "image_model_cards") @@ -85,8 +90,9 @@ class ModelCard(CamelCaseModel): data = tomlkit.dumps(py) # pyright: ignore[reportUnknownMemberType] await f.write(data) - async def save_to_default_path(self): - await self.save(Path(RESOURCES_DIR) / (self.model_id.normalize() + ".toml")) + async def save_to_custom_dir(self) -> None: + await aios.makedirs(str(_custom_cards_dir), exist_ok=True) + await self.save(_custom_cards_dir / (self.model_id.normalize() + ".toml")) @staticmethod async def load_from_path(path: Path) -> "ModelCard": @@ -120,11 +126,31 @@ class ModelCard(CamelCaseModel): supports_tensor=config_data.supports_tensor, tasks=[ModelTask.TextGeneration], ) - await mc.save_to_default_path() + await mc.save_to_custom_dir() _card_cache[model_id] = mc return mc +async def delete_custom_card(model_id: ModelId) -> bool: + """Delete a user-added custom model card. Returns True if deleted.""" + card_path = _custom_cards_dir / (ModelId(model_id).normalize() + ".toml") + if await card_path.exists(): + await card_path.unlink() + _card_cache.pop(model_id, None) + return True + return False + + +def is_custom_card(model_id: ModelId) -> bool: + """Check if a model card exists in the custom cards directory.""" + import os + + card_path = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR)) / ( + ModelId(model_id).normalize() + ".toml" + ) + return os.path.isfile(str(card_path)) + + # TODO: quantizing and dynamically creating model cards def _generate_image_model_quant_variants( # pyright: ignore[reportUnusedFunction] base_name: str, diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py index 40dbb288..e5710014 100644 --- a/src/exo/shared/types/api.py +++ b/src/exo/shared/types/api.py @@ -42,6 +42,7 @@ class ModelListModel(BaseModel): storage_size_megabytes: int = Field(default=0) supports_tensor: bool = Field(default=False) tasks: list[str] = Field(default=[]) + is_custom: bool = Field(default=False) class ModelList(BaseModel): @@ -201,6 +202,10 @@ class BenchChatCompletionRequest(ChatCompletionRequest): pass +class AddCustomModelParams(BaseModel): + model_id: ModelId + + class PlaceInstanceParams(BaseModel): model_id: ModelId sharding: Sharding = Sharding.Pipeline From 41ed7afb3bd6263c70cde9d29e5f440ce98782e8 Mon Sep 17 00:00:00 2001 From: Alex Cheema <41707476+AlexCheema@users.noreply.github.com> Date: Wed, 4 Feb 2026 05:56:23 -0800 Subject: [PATCH 4/9] feat: add model picker modal with grouped models and HF Hub search (#1369) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Motivation Reimplements the model picker modal from #1191 on top of the custom model support branch. Replaces the inline model dropdown with a full-featured modal that groups models by base model, supports filtering, favorites, and HuggingFace Hub search. ## Changes **Backend:** - Add `family`, `quantization`, `base_model`, `capabilities` metadata fields to `ModelCard` and all 40 TOML model cards - Pass new fields through `ModelListModel` and `get_models()` API response - Add `GET /models/search` endpoint using `huggingface_hub.list_models()` **Dashboard (7 new files):** - `ModelPickerModal.svelte` — Main modal with search, family filtering, HuggingFace Hub tab - `ModelPickerGroup.svelte` — Expandable model group row with quantization variants - `FamilySidebar.svelte` — Vertical sidebar with family icons (All, Favorites, Hub, model families) - `FamilyLogos.svelte` — SVG icons for each model family - `ModelFilterPopover.svelte` — Capability and size range filters - `HuggingFaceResultItem.svelte` — HF search result item with download/like counts - `favorites.svelte.ts` — localStorage-backed favorites store **Integration:** - Replace inline dropdown in `+page.svelte` with button that opens `ModelPickerModal` - Custom models shown in Hub tab with delete support **Polish:** - Real brand logos (Meta, Qwen, DeepSeek, OpenAI, GLM, MiniMax, Kimi, HuggingFace) from Simple Icons / LobeHub - Clean SVG stroke icons for capabilities (thinking, code, vision, image gen) - Consistent `border-exo-yellow/10` borders, descriptive tooltips throughout - Cluster memory (used/total) shown in modal header - Selected model highlight with checkmark for both single and multi-variant groups - Cursor pointer on all interactive elements, fix filter popover click-outside bug - Custom models now appear in All tab alongside built-in models ## Bug Fix: Gemma 3 EOS tokens Also included in this branch: fix for Gemma 3 models generating infinite `` tokens. The tokenizer's `eos_token_ids` was missing token ID 106 (``), so generation never stopped. The fix appends this token to the EOS list after loading the tokenizer. Also handles `eos_token_ids` being a `set` (not just a `list`). ## Why It Works Model metadata (family, capabilities, etc.) is stored directly in TOML cards rather than derived from heuristics, ensuring accuracy. The modal groups models by `base_model` field so quantization variants appear together. Custom models are separated into the Hub tab since they lack grouping metadata. ## Test Plan ### Manual Testing - Open dashboard, click model selector to open modal - Browse models by family sidebar, search, and filters - Expand model groups to see quantization variants - Star favorites and verify persistence across page reloads - Navigate to Hub tab, search and add models - Verify error messages shown for invalid model IDs - Run a Gemma 3 model and verify generation stops at `` ### Automated Testing - `uv run basedpyright` — 0 errors - `uv run ruff check` — passes - `nix fmt` — clean - `uv run pytest src/` — 173 passed - `cd dashboard && npm run build` — builds successfully --------- Co-authored-by: Claude Opus 4.5 --- .gitignore | 1 + .mlx_typings/mlx_lm/tokenizer_utils.pyi | 3 +- .../src/lib/components/FamilyLogos.svelte | 73 ++ .../src/lib/components/FamilySidebar.svelte | 142 ++++ .../components/HuggingFaceResultItem.svelte | 127 +++ .../lib/components/ModelFilterPopover.svelte | 182 +++++ .../lib/components/ModelPickerGroup.svelte | 324 ++++++++ .../lib/components/ModelPickerModal.svelte | 748 ++++++++++++++++++ dashboard/src/lib/components/index.ts | 6 + dashboard/src/lib/stores/favorites.svelte.ts | 97 +++ dashboard/src/routes/+page.svelte | 341 ++------ .../mlx-community--DeepSeek-V3.1-4bit.toml | 4 + .../mlx-community--DeepSeek-V3.1-8bit.toml | 4 + .../mlx-community--GLM-4.5-Air-8bit.toml | 4 + .../mlx-community--GLM-4.5-Air-bf16.toml | 4 + .../mlx-community--GLM-4.7-4bit.toml | 4 + .../mlx-community--GLM-4.7-6bit.toml | 4 + .../mlx-community--GLM-4.7-8bit-gs32.toml | 4 + .../mlx-community--GLM-4.7-Flash-4bit.toml | 4 + .../mlx-community--GLM-4.7-Flash-5bit.toml | 4 + .../mlx-community--GLM-4.7-Flash-6bit.toml | 4 + .../mlx-community--GLM-4.7-Flash-8bit.toml | 4 + .../mlx-community--Kimi-K2-Instruct-4bit.toml | 4 + .../mlx-community--Kimi-K2-Thinking.toml | 4 + .../mlx-community--Kimi-K2.5.toml | 4 + ...community--Llama-3.2-1B-Instruct-4bit.toml | 4 + ...community--Llama-3.2-3B-Instruct-4bit.toml | 4 + ...community--Llama-3.2-3B-Instruct-8bit.toml | 4 + ...ommunity--Llama-3.3-70B-Instruct-4bit.toml | 4 + ...ommunity--Llama-3.3-70B-Instruct-8bit.toml | 4 + ...ity--Meta-Llama-3.1-70B-Instruct-4bit.toml | 4 + ...nity--Meta-Llama-3.1-8B-Instruct-4bit.toml | 4 + ...nity--Meta-Llama-3.1-8B-Instruct-8bit.toml | 4 + ...nity--Meta-Llama-3.1-8B-Instruct-bf16.toml | 4 + .../mlx-community--MiniMax-M2.1-3bit.toml | 4 + .../mlx-community--MiniMax-M2.1-8bit.toml | 4 + .../mlx-community--Qwen3-0.6B-4bit.toml | 4 + .../mlx-community--Qwen3-0.6B-8bit.toml | 4 + ...y--Qwen3-235B-A22B-Instruct-2507-4bit.toml | 4 + ...y--Qwen3-235B-A22B-Instruct-2507-8bit.toml | 4 + .../mlx-community--Qwen3-30B-A3B-4bit.toml | 4 + .../mlx-community--Qwen3-30B-A3B-8bit.toml | 4 + ...--Qwen3-Coder-480B-A35B-Instruct-4bit.toml | 4 + ...--Qwen3-Coder-480B-A35B-Instruct-8bit.toml | 4 + ...ity--Qwen3-Next-80B-A3B-Instruct-4bit.toml | 4 + ...ity--Qwen3-Next-80B-A3B-Instruct-8bit.toml | 4 + ...ity--Qwen3-Next-80B-A3B-Thinking-4bit.toml | 4 + ...ity--Qwen3-Next-80B-A3B-Thinking-8bit.toml | 4 + .../mlx-community--gpt-oss-120b-MXFP4-Q8.toml | 4 + .../mlx-community--gpt-oss-20b-MXFP4-Q8.toml | 4 + ...ommunity--llama-3.3-70b-instruct-fp16.toml | 4 + src/exo/master/api.py | 30 + src/exo/shared/models/model_cards.py | 4 + src/exo/shared/types/api.py | 13 + src/exo/worker/engines/mlx/auto_parallel.py | 6 + src/exo/worker/engines/mlx/utils_mlx.py | 11 + 56 files changed, 1998 insertions(+), 270 deletions(-) create mode 100644 dashboard/src/lib/components/FamilyLogos.svelte create mode 100644 dashboard/src/lib/components/FamilySidebar.svelte create mode 100644 dashboard/src/lib/components/HuggingFaceResultItem.svelte create mode 100644 dashboard/src/lib/components/ModelFilterPopover.svelte create mode 100644 dashboard/src/lib/components/ModelPickerGroup.svelte create mode 100644 dashboard/src/lib/components/ModelPickerModal.svelte create mode 100644 dashboard/src/lib/stores/favorites.svelte.ts diff --git a/.gitignore b/.gitignore index d0b8299e..139fc326 100644 --- a/.gitignore +++ b/.gitignore @@ -31,3 +31,4 @@ dashboard/.svelte-kit/ # host config snapshots hosts_*.json +.swp diff --git a/.mlx_typings/mlx_lm/tokenizer_utils.pyi b/.mlx_typings/mlx_lm/tokenizer_utils.pyi index 251e3d28..83eb4e33 100644 --- a/.mlx_typings/mlx_lm/tokenizer_utils.pyi +++ b/.mlx_typings/mlx_lm/tokenizer_utils.pyi @@ -108,6 +108,7 @@ class TokenizerWrapper: _tokenizer: PreTrainedTokenizerFast eos_token_id: int | None eos_token: str | None + eos_token_ids: list[int] | set[int] | None bos_token_id: int | None bos_token: str | None vocab_size: int @@ -117,7 +118,7 @@ class TokenizerWrapper: self, tokenizer: Any, detokenizer_class: Any = ..., - eos_token_ids: list[int] | None = ..., + eos_token_ids: list[int] | set[int] | None = ..., chat_template: Any = ..., tool_parser: Any = ..., tool_call_start: str | None = ..., diff --git a/dashboard/src/lib/components/FamilyLogos.svelte b/dashboard/src/lib/components/FamilyLogos.svelte new file mode 100644 index 00000000..8e2919fe --- /dev/null +++ b/dashboard/src/lib/components/FamilyLogos.svelte @@ -0,0 +1,73 @@ + + +{#if family === "favorites"} + + + +{:else if family === "llama" || family === "meta"} + + + +{:else if family === "qwen"} + + + +{:else if family === "deepseek"} + + + +{:else if family === "openai" || family === "gpt-oss"} + + + +{:else if family === "glm"} + + + +{:else if family === "minimax"} + + + +{:else if family === "kimi"} + + + + +{:else if family === "huggingface"} + + + +{:else} + + + +{/if} diff --git a/dashboard/src/lib/components/FamilySidebar.svelte b/dashboard/src/lib/components/FamilySidebar.svelte new file mode 100644 index 00000000..886a68d5 --- /dev/null +++ b/dashboard/src/lib/components/FamilySidebar.svelte @@ -0,0 +1,142 @@ + + +
+ + + + + {#if hasFavorites} + + {/if} + + + + +
+ + + {#each families as family} + + {/each} +
diff --git a/dashboard/src/lib/components/HuggingFaceResultItem.svelte b/dashboard/src/lib/components/HuggingFaceResultItem.svelte new file mode 100644 index 00000000..566d8e17 --- /dev/null +++ b/dashboard/src/lib/components/HuggingFaceResultItem.svelte @@ -0,0 +1,127 @@ + + +
+
+
+ {modelName} + {#if isAdded} + Added + {/if} +
+
+ {model.author} + + + + + {formatNumber(model.downloads)} + + + + + + {formatNumber(model.likes)} + +
+
+ +
+ {#if isAdded} + + {:else} + + {/if} +
+
diff --git a/dashboard/src/lib/components/ModelFilterPopover.svelte b/dashboard/src/lib/components/ModelFilterPopover.svelte new file mode 100644 index 00000000..5406618a --- /dev/null +++ b/dashboard/src/lib/components/ModelFilterPopover.svelte @@ -0,0 +1,182 @@ + + + + + +
e.stopPropagation()} + role="dialog" + aria-label="Filter options" +> +
+ +
+

Capabilities

+
+ {#each availableCapabilities as cap} + {@const isSelected = filters.capabilities.includes(cap.id)} + + {/each} +
+
+ + +
+

Model Size

+
+ {#each sizeRanges as range} + {@const isSelected = + filters.sizeRange && + filters.sizeRange.min === range.min && + filters.sizeRange.max === range.max} + + {/each} +
+
+ + + +
+
diff --git a/dashboard/src/lib/components/ModelPickerGroup.svelte b/dashboard/src/lib/components/ModelPickerGroup.svelte new file mode 100644 index 00000000..b3ad425b --- /dev/null +++ b/dashboard/src/lib/components/ModelPickerGroup.svelte @@ -0,0 +1,324 @@ + + +
+ +
{ + if (group.hasMultipleVariants) { + onToggleExpand(); + } else { + const modelId = group.variants[0]?.id; + if (modelId && canModelFit(modelId)) { + onSelectModel(modelId); + } + } + }} + role="button" + tabindex="0" + onkeydown={(e) => { + if (e.key === "Enter" || e.key === " ") { + e.preventDefault(); + if (group.hasMultipleVariants) { + onToggleExpand(); + } else { + const modelId = group.variants[0]?.id; + if (modelId && canModelFit(modelId)) { + onSelectModel(modelId); + } + } + } + }} + > + + {#if group.hasMultipleVariants} + + + + {:else} +
+ {/if} + + +
+
+ + {group.name} + + + {#each group.capabilities.filter((c) => c !== "text") as cap} + {#if cap === "thinking"} + + + + {:else if cap === "code"} + + + + {:else if cap === "vision"} + + + + + {:else if cap === "image_gen"} + + + + + + {/if} + {/each} +
+
+ + + {#if !group.hasMultipleVariants && group.smallestVariant?.storage_size_megabytes} + + {formatSize(group.smallestVariant.storage_size_megabytes)} + + {/if} + + + {#if group.hasMultipleVariants} + + {group.variants.length} variants + + {/if} + + + {#if isMainSelected} + + + + {/if} + + + + + + +
+ + + {#if isExpanded && group.hasMultipleVariants} +
+ {#each group.variants as variant} + {@const modelCanFit = canModelFit(variant.id)} + {@const isSelected = selectedModelId === variant.id} + + {/each} +
+ {/if} +
diff --git a/dashboard/src/lib/components/ModelPickerModal.svelte b/dashboard/src/lib/components/ModelPickerModal.svelte new file mode 100644 index 00000000..40827a0e --- /dev/null +++ b/dashboard/src/lib/components/ModelPickerModal.svelte @@ -0,0 +1,748 @@ + + + + +{#if isOpen} + + + + + + + + {#if infoGroup} +
(infoGroup = null)} + role="presentation" + >
+ + {/if} +{/if} diff --git a/dashboard/src/lib/components/index.ts b/dashboard/src/lib/components/index.ts index dc8a7d76..b0424bf4 100644 --- a/dashboard/src/lib/components/index.ts +++ b/dashboard/src/lib/components/index.ts @@ -6,3 +6,9 @@ export { default as ChatSidebar } from "./ChatSidebar.svelte"; export { default as ModelCard } from "./ModelCard.svelte"; export { default as MarkdownContent } from "./MarkdownContent.svelte"; export { default as ImageParamsPanel } from "./ImageParamsPanel.svelte"; +export { default as FamilyLogos } from "./FamilyLogos.svelte"; +export { default as FamilySidebar } from "./FamilySidebar.svelte"; +export { default as HuggingFaceResultItem } from "./HuggingFaceResultItem.svelte"; +export { default as ModelFilterPopover } from "./ModelFilterPopover.svelte"; +export { default as ModelPickerGroup } from "./ModelPickerGroup.svelte"; +export { default as ModelPickerModal } from "./ModelPickerModal.svelte"; diff --git a/dashboard/src/lib/stores/favorites.svelte.ts b/dashboard/src/lib/stores/favorites.svelte.ts new file mode 100644 index 00000000..877b059c --- /dev/null +++ b/dashboard/src/lib/stores/favorites.svelte.ts @@ -0,0 +1,97 @@ +/** + * FavoritesStore - Manages favorite models with localStorage persistence + */ + +import { browser } from "$app/environment"; + +const FAVORITES_KEY = "exo-favorite-models"; + +class FavoritesStore { + favorites = $state>(new Set()); + + constructor() { + if (browser) { + this.loadFromStorage(); + } + } + + private loadFromStorage() { + try { + const stored = localStorage.getItem(FAVORITES_KEY); + if (stored) { + const parsed = JSON.parse(stored) as string[]; + this.favorites = new Set(parsed); + } + } catch (error) { + console.error("Failed to load favorites:", error); + } + } + + private saveToStorage() { + try { + const array = Array.from(this.favorites); + localStorage.setItem(FAVORITES_KEY, JSON.stringify(array)); + } catch (error) { + console.error("Failed to save favorites:", error); + } + } + + add(baseModelId: string) { + const next = new Set(this.favorites); + next.add(baseModelId); + this.favorites = next; + this.saveToStorage(); + } + + remove(baseModelId: string) { + const next = new Set(this.favorites); + next.delete(baseModelId); + this.favorites = next; + this.saveToStorage(); + } + + toggle(baseModelId: string) { + if (this.favorites.has(baseModelId)) { + this.remove(baseModelId); + } else { + this.add(baseModelId); + } + } + + isFavorite(baseModelId: string): boolean { + return this.favorites.has(baseModelId); + } + + getAll(): string[] { + return Array.from(this.favorites); + } + + getSet(): Set { + return new Set(this.favorites); + } + + hasAny(): boolean { + return this.favorites.size > 0; + } + + clearAll() { + this.favorites = new Set(); + this.saveToStorage(); + } +} + +export const favoritesStore = new FavoritesStore(); + +export const favorites = () => favoritesStore.favorites; +export const hasFavorites = () => favoritesStore.hasAny(); +export const isFavorite = (baseModelId: string) => + favoritesStore.isFavorite(baseModelId); +export const toggleFavorite = (baseModelId: string) => + favoritesStore.toggle(baseModelId); +export const addFavorite = (baseModelId: string) => + favoritesStore.add(baseModelId); +export const removeFavorite = (baseModelId: string) => + favoritesStore.remove(baseModelId); +export const getFavorites = () => favoritesStore.getAll(); +export const getFavoritesSet = () => favoritesStore.getSet(); +export const clearFavorites = () => favoritesStore.clearAll(); diff --git a/dashboard/src/routes/+page.svelte b/dashboard/src/routes/+page.svelte index 6ecce7cc..288c3991 100644 --- a/dashboard/src/routes/+page.svelte +++ b/dashboard/src/routes/+page.svelte @@ -5,7 +5,13 @@ ChatMessages, ChatSidebar, ModelCard, + ModelPickerModal, } from "$lib/components"; + import { + favorites, + toggleFavorite, + getFavoritesSet, + } from "$lib/stores/favorites.svelte"; import { hasStartedChat, isTopologyMinimized, @@ -101,6 +107,10 @@ tasks?: string[]; hugging_face_id?: string; is_custom?: boolean; + family?: string; + quantization?: string; + base_model?: string; + capabilities?: string[]; }> >([]); @@ -212,14 +222,11 @@ let launchingModelId = $state(null); let instanceDownloadExpandedNodes = $state>(new Set()); - // Custom dropdown state - let isModelDropdownOpen = $state(false); - let modelDropdownSearch = $state(""); + // Model picker modal state + let isModelPickerOpen = $state(false); - // Custom model add state - let customModelInput = $state(""); - let isAddingCustomModel = $state(false); - let customModelError = $state(""); + // Favorites state (reactive) + const favoritesSet = $derived(getFavoritesSet()); // Slider dragging state let isDraggingSlider = $state(false); @@ -536,41 +543,25 @@ } } - async function addCustomModel() { - const modelId = customModelInput.trim(); - if (!modelId) return; + async function addModelFromPicker(modelId: string) { + const response = await fetch("/models/add", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ model_id: modelId }), + }); - isAddingCustomModel = true; - customModelError = ""; - - try { - const response = await fetch("/models/add", { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ model_id: modelId }), - }); - - if (!response.ok) { - try { - const err = await response.json(); - customModelError = - err.detail || `Failed to add model (${response.status})`; - } catch { - customModelError = `Failed to add model (${response.status}: ${response.statusText})`; - } - return; + if (!response.ok) { + let message = `Failed to add model (${response.status}: ${response.statusText})`; + try { + const err = await response.json(); + if (err.detail) message = err.detail; + } catch { + // use default message } - - const added = await response.json(); - customModelInput = ""; - await fetchModels(); - selectPreviewModel(added.id); - isModelDropdownOpen = false; - } catch { - customModelError = "Network error"; - } finally { - isAddingCustomModel = false; + throw new Error(message); } + + await fetchModels(); } async function deleteCustomModel(modelId: string) { @@ -587,6 +578,12 @@ } } + function handleModelPickerSelect(modelId: string) { + selectPreviewModel(modelId); + saveLaunchDefaults(); + isModelPickerOpen = false; + } + async function launchInstance( modelId: string, specificPreview?: PlacementPreview | null, @@ -2417,14 +2414,12 @@ > - -
+ +
-
- - - -
- - {#if isModelDropdownOpen} - - - -
- -
- - -
{ - e.preventDefault(); - addCustomModel(); - }} - class="flex gap-1.5" - > - - -
- {#if customModelError} -
- {customModelError} -
- {/if} -
- - -
- {#each sortedModels().filter((m) => !modelDropdownSearch || (m.name || m.id) - .toLowerCase() - .includes(modelDropdownSearch.toLowerCase())) as model} - {@const sizeGB = getModelSizeGB(model)} - {@const modelCanFit = hasEnoughMemory(model)} - {@const isImageModel = modelSupportsImageGeneration( - model.id, - )} - {@const isImageEditModel = modelSupportsImageEditing( - model.id, - )} - - {/if} - - - {:else} -
- No models found -
- {/each} -
+
- {/if} +
@@ -3462,3 +3246,22 @@ {/if}
+ + m.id))} + canModelFit={(modelId) => { + const model = models.find((m) => m.id === modelId); + return model ? hasEnoughMemory(model) : false; + }} + onSelect={handleModelPickerSelect} + onClose={() => (isModelPickerOpen = false)} + onToggleFavorite={toggleFavorite} + onAddModel={addModelFromPicker} + onDeleteModel={deleteCustomModel} + totalMemoryGB={clusterMemory().total / (1024 * 1024 * 1024)} + usedMemoryGB={clusterMemory().used / (1024 * 1024 * 1024)} +/> diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml index 26de8de8..41784cf6 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "deepseek" +quantization = "4bit" +base_model = "DeepSeek V3.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 405874409472 diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml index 13cf367b..a5d77bcd 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "deepseek" +quantization = "8bit" +base_model = "DeepSeek V3.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 765577920512 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml index 288392f6..a7acea44 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml @@ -3,6 +3,10 @@ n_layers = 46 hidden_size = 4096 supports_tensor = false tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.5 Air" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 122406567936 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml index 00a19df2..4258c225 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml @@ -3,6 +3,10 @@ n_layers = 46 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "bf16" +base_model = "GLM 4.5 Air" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 229780750336 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml index 816c9657..0672d664 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "4bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 198556925568 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml index b087164b..bcf1cae4 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "6bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 286737579648 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml index 6f221cef..0f56c2f7 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 396963397248 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml index 43eb0dcd..8637cef0 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "4bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 19327352832 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml index 6a512c0a..b9a9da4d 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "5bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 22548578304 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml index 86c65489..e3cb1fa8 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "6bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 26843545600 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml index eb69183f..bd6df312 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 34359738368 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml index d7acabec..3f21d4c0 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "4bit" +base_model = "Kimi K2" +capabilities = ["text"] [storage_size] in_bytes = 620622774272 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml index 8d21727b..0a955b04 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "" +base_model = "Kimi K2" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 706522120192 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml index c44cf9b1..806c6b30 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "" +base_model = "Kimi K2.5" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 662498705408 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml index db334221..b38ec20f 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 16 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.2 1B" +capabilities = ["text"] [storage_size] in_bytes = 729808896 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml index 001b4a1d..81ce4567 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.2 3B" +capabilities = ["text"] [storage_size] in_bytes = 1863319552 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml index 358db81b..ac9a203b 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.2 3B" +capabilities = ["text"] [storage_size] in_bytes = 3501195264 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml index cf6eece0..24c7cbaa 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 40652242944 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml index 15f0c551..3bfc97dc 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 76799803392 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml index b766164d..27d0b724 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.1 70B" +capabilities = ["text"] [storage_size] in_bytes = 40652242944 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml index b6d10c40..1fe34ba8 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 4637851648 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml index 4cfe47cf..5310a2a0 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 8954839040 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml index 9e04f27b..eb6405e0 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "bf16" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 16882073600 diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml index 4bf81136..92ec6746 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "minimax" +quantization = "3bit" +base_model = "MiniMax M2.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 100086644736 diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml index 54a49f97..c1388d2f 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "minimax" +quantization = "8bit" +base_model = "MiniMax M2.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 242986745856 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml index 212cdef6..7929aaba 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 1024 supports_tensor = false tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 0.6B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 342884352 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml index ac591d6c..d9fcc368 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 1024 supports_tensor = false tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 0.6B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 698351616 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml index 020c11be..ef835c6a 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml @@ -3,6 +3,10 @@ n_layers = 94 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 235B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 141733920768 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml index 64afc366..f6e079ab 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml @@ -3,6 +3,10 @@ n_layers = 94 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 235B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 268435456000 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml index 1b9f92a6..48a6666f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 30B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 17612931072 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml index f8e59ac1..c283396f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 30B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 33279705088 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml index 25a0cf2b..b390bd21 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 62 hidden_size = 6144 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Coder 480B" +capabilities = ["text", "code"] [storage_size] in_bytes = 289910292480 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml index cdb1d6ec..1c21307c 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 62 hidden_size = 6144 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Coder 480B" +capabilities = ["text", "code"] [storage_size] in_bytes = 579820584960 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml index db55b7f9..386a3fa1 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text"] [storage_size] in_bytes = 46976204800 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml index e36e24b8..0e2bf2a5 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text"] [storage_size] in_bytes = 88814387200 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml index bc3bdf50..2a3e3c19 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 47080074240 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml index dd5512a7..65d33253 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 88814387200 diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml index c725e728..f579c618 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml @@ -3,6 +3,10 @@ n_layers = 36 hidden_size = 2880 supports_tensor = true tasks = ["TextGeneration"] +family = "gpt-oss" +quantization = "MXFP4-Q8" +base_model = "GPT-OSS 120B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 70652212224 diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml index bf8f1a60..af1e04ad 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml @@ -3,6 +3,10 @@ n_layers = 24 hidden_size = 2880 supports_tensor = true tasks = ["TextGeneration"] +family = "gpt-oss" +quantization = "MXFP4-Q8" +base_model = "GPT-OSS 20B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 12025908224 diff --git a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml index dd451015..e61660c2 100644 --- a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml +++ b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "fp16" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 144383672320 diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 01e145b4..282cb0f1 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -74,6 +74,7 @@ from exo.shared.types.api import ( ErrorResponse, FinishReason, GenerationStats, + HuggingFaceSearchResult, ImageData, ImageEditsTaskParams, ImageGenerationResponse, @@ -262,6 +263,7 @@ class API: self.app.get("/v1/models")(self.get_models) self.app.post("/models/add")(self.add_custom_model) self.app.delete("/models/custom/{model_id:path}")(self.delete_custom_model) + self.app.get("/models/search")(self.search_models) self.app.post("/v1/chat/completions", response_model=None)( self.chat_completions ) @@ -1222,6 +1224,10 @@ class API: supports_tensor=card.supports_tensor, tasks=[task.value for task in card.tasks], is_custom=is_custom_card(card.model_id), + family=card.family, + quantization=card.quantization, + base_model=card.base_model, + capabilities=card.capabilities, ) for card in await get_model_cards() ] @@ -1257,6 +1263,30 @@ class API: {"message": "Model card deleted", "model_id": str(model_id)} ) + async def search_models( + self, query: str = "", limit: int = 20 + ) -> list[HuggingFaceSearchResult]: + """Search HuggingFace Hub for mlx-community models.""" + from huggingface_hub import list_models + + results = list_models( + search=query or None, + author="mlx-community", + sort="downloads", + limit=limit, + ) + return [ + HuggingFaceSearchResult( + id=m.id, + author=m.author or "", + downloads=m.downloads or 0, + likes=m.likes or 0, + last_modified=str(m.last_modified or ""), + tags=list(m.tags or []), + ) + for m in results + ] + async def run(self): cfg = Config() cfg.bind = f"0.0.0.0:{self.port}" diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index 47b54d38..79dc95be 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -78,6 +78,10 @@ class ModelCard(CamelCaseModel): supports_tensor: bool tasks: list[ModelTask] components: list[ComponentInfo] | None = None + family: str = "" + quantization: str = "" + base_model: str = "" + capabilities: list[str] = [] @field_validator("tasks", mode="before") @classmethod diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py index e5710014..5a7bae1e 100644 --- a/src/exo/shared/types/api.py +++ b/src/exo/shared/types/api.py @@ -43,6 +43,10 @@ class ModelListModel(BaseModel): supports_tensor: bool = Field(default=False) tasks: list[str] = Field(default=[]) is_custom: bool = Field(default=False) + family: str = Field(default="") + quantization: str = Field(default="") + base_model: str = Field(default="") + capabilities: list[str] = Field(default_factory=list) class ModelList(BaseModel): @@ -206,6 +210,15 @@ class AddCustomModelParams(BaseModel): model_id: ModelId +class HuggingFaceSearchResult(BaseModel): + id: str + author: str = "" + downloads: int = 0 + likes: int = 0 + last_modified: str = "" + tags: list[str] = Field(default_factory=list) + + class PlaceInstanceParams(BaseModel): model_id: ModelId sharding: Sharding = Sharding.Pipeline diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 84ed1f29..36589b7a 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -164,6 +164,12 @@ def _inner_model(model: nn.Module) -> nn.Module: if isinstance(inner, nn.Module): return inner + inner = getattr(model, "language_model", None) + if isinstance(inner, nn.Module): + inner_inner = getattr(inner, "model", None) + if isinstance(inner_inner, nn.Module): + return inner_inner + raise ValueError("Model must either have a 'model' or 'transformer' attribute") diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index d7fb9958..5dae48a4 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -384,6 +384,17 @@ def load_tokenizer_for_model_id( eos_token_ids=eos_token_ids, ) + if "gemma-3" in model_id_lower: + gemma_3_eos_id = 1 + gemma_3_end_of_turn_id = 106 + if tokenizer.eos_token_ids is not None: + if gemma_3_end_of_turn_id not in tokenizer.eos_token_ids: + tokenizer.eos_token_ids = list(tokenizer.eos_token_ids) + [ + gemma_3_end_of_turn_id + ] + else: + tokenizer.eos_token_ids = [gemma_3_eos_id, gemma_3_end_of_turn_id] + return tokenizer From 7b6cad94c675c99899d75376f60994f314c48d6c Mon Sep 17 00:00:00 2001 From: Evan Quiney Date: Wed, 4 Feb 2026 16:38:43 +0000 Subject: [PATCH 5/9] add resources dir to nix (#1376) add resources directory to the nix exo package, and fixes the env for the dashboard dir --- .github/workflows/pipeline.yml | 4 +++- python/parts.nix | 3 ++- src/exo/shared/constants.py | 2 +- 3 files changed, 6 insertions(+), 3 deletions(-) diff --git a/.github/workflows/pipeline.yml b/.github/workflows/pipeline.yml index 0f908f1b..c2589453 100644 --- a/.github/workflows/pipeline.yml +++ b/.github/workflows/pipeline.yml @@ -142,4 +142,6 @@ jobs: # Run pytest outside sandbox (needs GPU access for MLX) export HOME="$RUNNER_TEMP" export EXO_TESTS=1 - EXO_RESOURCES_DIR="$PWD/resources" $TEST_ENV/bin/python -m pytest src -m "not slow" --import-mode=importlib + export EXO_DASHBOARD_DIR="$PWD/dashboard/" + export EXO_RESOURCES_DIR="$PWD/resources" + $TEST_ENV/bin/python -m pytest src -m "not slow" --import-mode=importlib diff --git a/python/parts.nix b/python/parts.nix index 7423ed31..9d5580ae 100644 --- a/python/parts.nix +++ b/python/parts.nix @@ -69,7 +69,8 @@ # Create wrapper scripts for script in exo exo-master exo-worker; do makeWrapper ${exoVenv}/bin/$script $out/bin/$script \ - --set DASHBOARD_DIR ${self'.packages.dashboard} \ + --set EXO_DASHBOARD_DIR ${self'.packages.dashboard} \ + --set EXO_RESOURCES_DIR ${inputs.self + "/resources"} \ ${lib.optionalString pkgs.stdenv.isDarwin "--prefix PATH : ${pkgs.macmon}/bin"} done ''; diff --git a/src/exo/shared/constants.py b/src/exo/shared/constants.py index b09d6b2d..b385b5e8 100644 --- a/src/exo/shared/constants.py +++ b/src/exo/shared/constants.py @@ -39,7 +39,7 @@ RESOURCES_DIR = ( ) _DASHBOARD_DIR_ENV = os.environ.get("EXO_DASHBOARD_DIR", None) DASHBOARD_DIR = ( - find_dashboard() if _RESOURCES_DIR_ENV is None else Path.home() / _RESOURCES_DIR_ENV + find_dashboard() if _DASHBOARD_DIR_ENV is None else Path.home() / _DASHBOARD_DIR_ENV ) # Log files (data/logs or cache) From 6177550c34ce17b2e181d96b4d561df50ceb086f Mon Sep 17 00:00:00 2001 From: ciaranbor <81697641+ciaranbor@users.noreply.github.com> Date: Wed, 4 Feb 2026 21:16:35 +0000 Subject: [PATCH 6/9] Ciaran/parallel cfg (#1361) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Motivation Enable parallel classifier-free guidance (CFG) for Qwen image models. CFG requires two forward passes (positive/negative prompts) - this allows them to run on separate nodes simultaneously, reducing latency. ## Changes - Added uses_cfg flag to ModelCard to identify CFG-based models - Extended PipelineShardMetadata with CFG topology fields (cfg_rank, cfg_world_size, peer device info) - Updated placement to create two CFG groups with reversed ordering (places CFG peers as ring neighbors) - Refactored DiffusionRunner to process CFG branches separately with exchange at last pipeline stage - Added get_cfg_branch_data() to PromptData for single-branch embeddings - Fixed seed handling in API for distributed consistency - Fixed image yield to only emit from CFG rank 0 at last stage - Increased num_sync_steps_factor from 0.125 to 0.25 for Qwen ## Why It Works - 2 nodes + CFG: Both run all layers, process different CFG branches in parallel - 4+ even nodes + CFG: Hybrid - 2 CFG groups × N/2 pipeline stages - Odd nodes or non-CFG: Falls back to pure pipeline parallelism Ring topology places CFG peers as neighbors to enable direct exchange. ## Test Plan ### Manual Testing Verified performance gain for Qwen-Image for 2 node and 4 node cluster. Non-CFG models still work ### Automated Testing Added tests in test_placement_utils.py covering 2-node CFG parallel, 4-node hybrid, odd-node fallback, and non-CFG pipeline modes. --- .../exolabs--Qwen-Image-4bit.toml | 1 + .../exolabs--Qwen-Image-8bit.toml | 1 + .../exolabs--Qwen-Image-Edit-2509-4bit.toml | 1 + .../exolabs--Qwen-Image-Edit-2509-8bit.toml | 1 + .../exolabs--Qwen-Image-Edit-2509.toml | 1 + .../exolabs--Qwen-Image.toml | 1 + src/exo/master/api.py | 17 + src/exo/master/placement_utils.py | 165 ++++- src/exo/master/tests/test_placement_utils.py | 197 +++++- src/exo/shared/models/model_cards.py | 86 +-- src/exo/shared/types/worker/shards.py | 19 +- .../worker/engines/image/distributed_model.py | 17 +- src/exo/worker/engines/image/models/base.py | 21 + .../engines/image/models/flux/adapter.py | 6 + .../engines/image/models/qwen/adapter.py | 18 + .../engines/image/models/qwen/config.py | 4 +- .../engines/image/models/qwen/edit_adapter.py | 18 + .../worker/engines/image/pipeline/runner.py | 580 ++++++++++++------ src/exo/worker/engines/mlx/utils_mlx.py | 6 + src/exo/worker/runner/runner.py | 40 +- 20 files changed, 869 insertions(+), 331 deletions(-) diff --git a/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml index 89cd0f6f..8d3a637e 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 26799533856 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml index 43951dab..ddf78c4a 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 37014734400 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml index 99a60af2..db2f5e54 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 26799533856 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml index 0f326b39..2db63265 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 37014734400 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml index 65044e6c..3b615da1 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 57445135488 diff --git a/resources/image_model_cards/exolabs--Qwen-Image.toml b/resources/image_model_cards/exolabs--Qwen-Image.toml index a39235ea..d012af50 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 57445135488 diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 282cb0f1..9bd8cbcf 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -1,6 +1,7 @@ import base64 import contextlib import json +import random import time from collections.abc import AsyncGenerator, Awaitable, Callable from datetime import datetime, timezone @@ -150,6 +151,15 @@ def _format_to_content_type(image_format: Literal["png", "jpeg", "webp"] | None) return f"image/{image_format or 'png'}" +def _ensure_seed(params: AdvancedImageParams | None) -> AdvancedImageParams: + """Ensure advanced params has a seed set for distributed consistency.""" + if params is None: + return AdvancedImageParams(seed=random.randint(0, 2**32 - 1)) + if params.seed is None: + return params.model_copy(update={"seed": random.randint(0, 2**32 - 1)}) + return params + + class API: def __init__( self, @@ -709,6 +719,9 @@ class API: with SSE-formatted events for partial and final images. """ payload.model = await self._validate_image_model(ModelId(payload.model)) + payload = payload.model_copy( + update={"advanced_params": _ensure_seed(payload.advanced_params)} + ) command = ImageGeneration( task_params=payload, @@ -957,6 +970,9 @@ class API: payload.stream = False payload.partial_images = 0 + payload = payload.model_copy( + update={"advanced_params": _ensure_seed(payload.advanced_params)} + ) command = ImageGeneration( task_params=payload, @@ -988,6 +1004,7 @@ class API: ) -> ImageEdits: """Prepare and send an image edits command with chunked image upload.""" resolved_model = await self._validate_image_model(model) + advanced_params = _ensure_seed(advanced_params) image_content = await image.read() image_data = base64.b64encode(image_content).decode("utf-8") diff --git a/src/exo/master/placement_utils.py b/src/exo/master/placement_utils.py index 309abc25..b20a39cc 100644 --- a/src/exo/master/placement_utils.py +++ b/src/exo/master/placement_utils.py @@ -10,6 +10,7 @@ from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo from exo.shared.types.topology import Cycle, RDMAConnection, SocketConnection from exo.shared.types.worker.runners import RunnerId, ShardAssignments from exo.shared.types.worker.shards import ( + CfgShardMetadata, PipelineShardMetadata, Sharding, ShardMetadata, @@ -74,40 +75,43 @@ def allocate_layers_proportionally( return result -def get_shard_assignments_for_pipeline_parallel( - model_card: ModelCard, - cycle: Cycle, - node_memory: Mapping[NodeId, MemoryUsage], -): +def _validate_cycle(cycle: Cycle) -> None: if not cycle.node_ids: raise ValueError("Cannot create shard assignments for empty node cycle") - cycle_memory = sum( - (node_memory[node_id].ram_available for node_id in cycle.node_ids), + +def _compute_total_memory( + node_ids: list[NodeId], + node_memory: Mapping[NodeId, MemoryUsage], +) -> Memory: + total_memory = sum( + (node_memory[node_id].ram_available for node_id in node_ids), start=Memory(), ) - if cycle_memory.in_bytes == 0: + if total_memory.in_bytes == 0: raise ValueError("Cannot create shard assignments: total available memory is 0") + return total_memory - total_layers = model_card.n_layers - world_size = len(cycle) - runner_to_shard: dict[RunnerId, ShardMetadata] = {} - node_to_runner: dict[NodeId, RunnerId] = {} +def _allocate_and_validate_layers( + node_ids: list[NodeId], + node_memory: Mapping[NodeId, MemoryUsage], + total_memory: Memory, + model_card: ModelCard, +) -> list[int]: layer_allocations = allocate_layers_proportionally( - total_layers=total_layers, + total_layers=model_card.n_layers, memory_fractions=[ - node_memory[node_id].ram_available.in_bytes / cycle_memory.in_bytes - for node_id in cycle.node_ids + node_memory[node_id].ram_available.in_bytes / total_memory.in_bytes + for node_id in node_ids ], ) - # Validate each node has sufficient memory for its assigned layers - memory_per_layer = model_card.storage_size.in_bytes / total_layers - for i, (node_id, node_layers) in enumerate( - zip(cycle.node_ids, layer_allocations, strict=True) - ): - required_memory = node_layers * memory_per_layer + total_storage_bytes = model_card.storage_size.in_bytes + total_layers = model_card.n_layers + for i, node_id in enumerate(node_ids): + node_layers = layer_allocations[i] + required_memory = (total_storage_bytes * node_layers) // total_layers available_memory = node_memory[node_id].ram_available.in_bytes if required_memory > available_memory: raise ValueError( @@ -116,32 +120,125 @@ def get_shard_assignments_for_pipeline_parallel( f"but only has {available_memory / (1024**3):.2f} GB available" ) - layers_assigned = 0 - for i, (node_id, node_layers) in enumerate( - zip(cycle.node_ids, layer_allocations, strict=True) - ): - runner_id = RunnerId() + return layer_allocations - shard = PipelineShardMetadata( + +def get_shard_assignments_for_pipeline_parallel( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for pipeline parallel execution.""" + world_size = len(cycle) + use_cfg_parallel = model_card.uses_cfg and world_size >= 2 and world_size % 2 == 0 + + if use_cfg_parallel: + return _get_shard_assignments_for_cfg_parallel(model_card, cycle, node_memory) + else: + return _get_shard_assignments_for_pure_pipeline(model_card, cycle, node_memory) + + +def _get_shard_assignments_for_cfg_parallel( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for CFG parallel execution. + + CFG parallel runs two independent pipelines. Group 0 processes the positive + prompt, group 1 processes the negative prompt. The ring topology places + group 1's ranks in reverse order so both "last stages" are neighbors for + efficient CFG exchange. + """ + _validate_cycle(cycle) + + world_size = len(cycle) + cfg_world_size = 2 + pipeline_world_size = world_size // cfg_world_size + + # Allocate layers for one pipeline group (both groups run the same layers) + pipeline_node_ids = cycle.node_ids[:pipeline_world_size] + pipeline_memory = _compute_total_memory(pipeline_node_ids, node_memory) + layer_allocations = _allocate_and_validate_layers( + pipeline_node_ids, node_memory, pipeline_memory, model_card + ) + + # Ring topology: group 0 ascending [0,1,2,...], group 1 descending [...,2,1,0] + # This places both last stages as neighbors for CFG exchange. + position_to_cfg_pipeline = [(0, r) for r in range(pipeline_world_size)] + [ + (1, r) for r in reversed(range(pipeline_world_size)) + ] + + runner_to_shard: dict[RunnerId, ShardMetadata] = {} + node_to_runner: dict[NodeId, RunnerId] = {} + + for device_rank, node_id in enumerate(cycle.node_ids): + cfg_rank, pipeline_rank = position_to_cfg_pipeline[device_rank] + layers_before = sum(layer_allocations[:pipeline_rank]) + node_layers = layer_allocations[pipeline_rank] + + shard = CfgShardMetadata( model_card=model_card, - device_rank=i, + device_rank=device_rank, world_size=world_size, - start_layer=layers_assigned, - end_layer=layers_assigned + node_layers, - n_layers=total_layers, + start_layer=layers_before, + end_layer=layers_before + node_layers, + n_layers=model_card.n_layers, + cfg_rank=cfg_rank, + cfg_world_size=cfg_world_size, + pipeline_rank=pipeline_rank, + pipeline_world_size=pipeline_world_size, ) + runner_id = RunnerId() runner_to_shard[runner_id] = shard node_to_runner[node_id] = runner_id - layers_assigned += node_layers - shard_assignments = ShardAssignments( + return ShardAssignments( model_id=model_card.model_id, runner_to_shard=runner_to_shard, node_to_runner=node_to_runner, ) - return shard_assignments + +def _get_shard_assignments_for_pure_pipeline( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for pure pipeline execution.""" + _validate_cycle(cycle) + total_memory = _compute_total_memory(cycle.node_ids, node_memory) + + layer_allocations = _allocate_and_validate_layers( + cycle.node_ids, node_memory, total_memory, model_card + ) + + runner_to_shard: dict[RunnerId, ShardMetadata] = {} + node_to_runner: dict[NodeId, RunnerId] = {} + + for pipeline_rank, node_id in enumerate(cycle.node_ids): + layers_before = sum(layer_allocations[:pipeline_rank]) + node_layers = layer_allocations[pipeline_rank] + + shard = PipelineShardMetadata( + model_card=model_card, + device_rank=pipeline_rank, + world_size=len(cycle), + start_layer=layers_before, + end_layer=layers_before + node_layers, + n_layers=model_card.n_layers, + ) + + runner_id = RunnerId() + runner_to_shard[runner_id] = shard + node_to_runner[node_id] = runner_id + + return ShardAssignments( + model_id=model_card.model_id, + runner_to_shard=runner_to_shard, + node_to_runner=node_to_runner, + ) def get_shard_assignments_for_tensor_parallel( diff --git a/src/exo/master/tests/test_placement_utils.py b/src/exo/master/tests/test_placement_utils.py index f2cb1067..245c4fd7 100644 --- a/src/exo/master/tests/test_placement_utils.py +++ b/src/exo/master/tests/test_placement_utils.py @@ -5,6 +5,7 @@ from exo.master.placement_utils import ( filter_cycles_by_memory, get_mlx_jaccl_coordinators, get_shard_assignments, + get_shard_assignments_for_pipeline_parallel, get_smallest_cycles, ) from exo.master.tests.conftest import ( @@ -20,7 +21,11 @@ from exo.shared.types.profiling import ( NodeNetworkInfo, ) from exo.shared.types.topology import Connection, SocketConnection -from exo.shared.types.worker.shards import Sharding +from exo.shared.types.worker.shards import ( + CfgShardMetadata, + PipelineShardMetadata, + Sharding, +) def test_filter_cycles_by_memory(): @@ -487,3 +492,193 @@ def test_get_shard_assignments_insufficient_memory_raises(): get_shard_assignments( model_card, selected_cycle, Sharding.Pipeline, node_memory ) + + +class TestCfgParallelPlacement: + def _create_ring_topology(self, node_ids: list[NodeId]) -> Topology: + topology = Topology() + for node_id in node_ids: + topology.add_node(node_id) + + for i, node_id in enumerate(node_ids): + next_node = node_ids[(i + 1) % len(node_ids)] + conn = Connection( + source=node_id, + sink=next_node, + edge=create_socket_connection(i + 1), + ) + topology.add_connection(conn) + + return topology + + def test_two_nodes_cfg_model_uses_cfg_parallel(self): + """Two nodes with CFG model should use CFG parallel (no pipeline).""" + node_a = NodeId() + node_b = NodeId() + + topology = self._create_ring_topology([node_a, node_b]) + cycles = [c for c in topology.get_cycles() if len(c) == 2] + cycle = cycles[0] + + node_memory = { + node_a: create_node_memory(1000 * 1024), + node_b: create_node_memory(1000 * 1024), + } + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 2 + + # CFG models should get CfgShardMetadata + for shard in shards: + assert isinstance(shard, CfgShardMetadata) + # Both nodes should have all layers (no pipeline split) + assert shard.start_layer == 0 + assert shard.end_layer == 60 + assert shard.cfg_world_size == 2 + # Each node is the only stage in its pipeline group + assert shard.pipeline_world_size == 1 + assert shard.pipeline_rank == 0 + + cfg_ranks = sorted( + s.cfg_rank for s in shards if isinstance(s, CfgShardMetadata) + ) + assert cfg_ranks == [0, 1] + + def test_four_nodes_cfg_model_uses_hybrid(self): + """Four nodes with CFG model should use 2 CFG groups x 2 pipeline stages.""" + nodes = [NodeId() for _ in range(4)] + + topology = self._create_ring_topology(nodes) + cycles = [c for c in topology.get_cycles() if len(c) == 4] + cycle = cycles[0] + + node_memory = {n: create_node_memory(1000 * 1024) for n in nodes} + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 4 + + # CFG models should get CfgShardMetadata + for shard in shards: + assert isinstance(shard, CfgShardMetadata) + assert shard.cfg_world_size == 2 + assert shard.pipeline_world_size == 2 + assert shard.pipeline_rank in [0, 1] + + # Check we have 2 nodes in each CFG group + cfg_0_shards = [ + s for s in shards if isinstance(s, CfgShardMetadata) and s.cfg_rank == 0 + ] + cfg_1_shards = [ + s for s in shards if isinstance(s, CfgShardMetadata) and s.cfg_rank == 1 + ] + assert len(cfg_0_shards) == 2 + assert len(cfg_1_shards) == 2 + + # Both CFG groups should have the same layer assignments + cfg_0_layers = [(s.start_layer, s.end_layer) for s in cfg_0_shards] + cfg_1_layers = [(s.start_layer, s.end_layer) for s in cfg_1_shards] + assert sorted(cfg_0_layers) == sorted(cfg_1_layers) + + def test_three_nodes_cfg_model_uses_sequential_cfg(self): + """Three nodes (odd) with CFG model should use sequential CFG (PipelineShardMetadata).""" + nodes = [NodeId() for _ in range(3)] + + topology = self._create_ring_topology(nodes) + cycles = [c for c in topology.get_cycles() if len(c) == 3] + cycle = cycles[0] + + node_memory = {n: create_node_memory(1000 * 1024) for n in nodes} + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 3 + + # Odd node count with CFG model falls back to PipelineShardMetadata (sequential CFG) + for shard in shards: + assert isinstance(shard, PipelineShardMetadata) + + def test_two_nodes_non_cfg_model_uses_pipeline(self): + """Two nodes with non-CFG model should use pure pipeline (PipelineShardMetadata).""" + node_a = NodeId() + node_b = NodeId() + + topology = self._create_ring_topology([node_a, node_b]) + cycles = [c for c in topology.get_cycles() if len(c) == 2] + cycle = cycles[0] + + node_memory = { + node_a: create_node_memory(1000 * 1024), + node_b: create_node_memory(1000 * 1024), + } + + model_card = ModelCard( + model_id=ModelId("flux-test"), + n_layers=57, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=False, # Non-CFG model + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 2 + + # Non-CFG models should get PipelineShardMetadata + for shard in shards: + assert isinstance(shard, PipelineShardMetadata) + + # Should have actual layer sharding (pipeline) + layer_ranges = sorted( + (s.start_layer, s.end_layer) + for s in shards + if isinstance(s, PipelineShardMetadata) + ) + # First shard starts at 0, last shard ends at 57 + assert layer_ranges[0][0] == 0 + assert layer_ranges[-1][1] == 57 diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index 79dc95be..f9448f11 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -65,9 +65,9 @@ class ComponentInfo(CamelCaseModel): component_name: str component_path: str storage_size: Memory - n_layers: PositiveInt | None + n_layers: PositiveInt | None = None can_shard: bool - safetensors_index_filename: str | None + safetensors_index_filename: str | None = None class ModelCard(CamelCaseModel): @@ -82,6 +82,7 @@ class ModelCard(CamelCaseModel): quantization: str = "" base_model: str = "" capabilities: list[str] = [] + uses_cfg: bool = False @field_validator("tasks", mode="before") @classmethod @@ -155,87 +156,6 @@ def is_custom_card(model_id: ModelId) -> bool: return os.path.isfile(str(card_path)) -# TODO: quantizing and dynamically creating model cards -def _generate_image_model_quant_variants( # pyright: ignore[reportUnusedFunction] - base_name: str, - base_card: ModelCard, -) -> dict[str, ModelCard]: - """Create quantized variants of an image model card. - - Only the transformer component is quantized; text encoders stay at bf16. - Sizes are calculated exactly from the base card's component sizes. - """ - if base_card.components is None: - raise ValueError(f"Image model {base_name} must have components defined") - - # quantizations = [8, 6, 5, 4, 3] - quantizations = [8, 4] - - num_transformer_bytes = next( - c.storage_size.in_bytes - for c in base_card.components - if c.component_name == "transformer" - ) - - transformer_bytes = Memory.from_bytes(num_transformer_bytes) - - remaining_bytes = Memory.from_bytes( - sum( - c.storage_size.in_bytes - for c in base_card.components - if c.component_name != "transformer" - ) - ) - - def with_transformer_size(new_size: Memory) -> list[ComponentInfo]: - assert base_card.components is not None - return [ - ComponentInfo( - component_name=c.component_name, - component_path=c.component_path, - storage_size=new_size - if c.component_name == "transformer" - else c.storage_size, - n_layers=c.n_layers, - can_shard=c.can_shard, - safetensors_index_filename=c.safetensors_index_filename, - ) - for c in base_card.components - ] - - variants = { - base_name: ModelCard( - model_id=base_card.model_id, - storage_size=transformer_bytes + remaining_bytes, - n_layers=base_card.n_layers, - hidden_size=base_card.hidden_size, - supports_tensor=base_card.supports_tensor, - tasks=base_card.tasks, - components=with_transformer_size(transformer_bytes), - ) - } - - for quant in quantizations: - quant_transformer_bytes = Memory.from_bytes( - (num_transformer_bytes * quant) // 16 - ) - total_bytes = remaining_bytes + quant_transformer_bytes - - model_id = ModelId(base_card.model_id + f"-{quant}bit") - - variants[f"{base_name}-{quant}bit"] = ModelCard( - model_id=model_id, - storage_size=total_bytes, - n_layers=base_card.n_layers, - hidden_size=base_card.hidden_size, - supports_tensor=base_card.supports_tensor, - tasks=base_card.tasks, - components=with_transformer_size(quant_transformer_bytes), - ) - - return variants - - class ConfigData(BaseModel): model_config = {"extra": "ignore"} # Allow unknown fields diff --git a/src/exo/shared/types/worker/shards.py b/src/exo/shared/types/worker/shards.py index 8bb23a57..59a6c54e 100644 --- a/src/exo/shared/types/worker/shards.py +++ b/src/exo/shared/types/worker/shards.py @@ -1,4 +1,5 @@ from enum import Enum +from typing import TypeAlias, final from pydantic import Field @@ -51,6 +52,7 @@ class BaseShardMetadata(TaggedModel): ) +@final class PipelineShardMetadata(BaseShardMetadata): """ Pipeline parallelism shard meta. @@ -60,8 +62,23 @@ class PipelineShardMetadata(BaseShardMetadata): """ +@final +class CfgShardMetadata(BaseShardMetadata): + """Shard metadata for CFG-parallel image generation models.""" + + cfg_rank: int # 0 = positive branch, 1 = negative branch + cfg_world_size: int = 2 + + # Pipeline-relative coordinates (computed at placement time) + pipeline_rank: int # rank within the pipeline group (0, 1, 2, ...) + pipeline_world_size: int # number of nodes per pipeline group + + +@final class TensorShardMetadata(BaseShardMetadata): pass -ShardMetadata = PipelineShardMetadata | TensorShardMetadata +ShardMetadata: TypeAlias = ( + PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata +) diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index bafa9319..8c9bd04c 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -9,7 +9,7 @@ from PIL import Image from exo.download.download_utils import build_model_path from exo.shared.types.api import AdvancedImageParams from exo.shared.types.worker.instances import BoundInstance -from exo.shared.types.worker.shards import PipelineShardMetadata +from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models import ( create_adapter_for_model, @@ -30,14 +30,19 @@ class DistributedImageModel: self, model_id: str, local_path: Path, - shard_metadata: PipelineShardMetadata, + shard_metadata: PipelineShardMetadata | CfgShardMetadata, group: Optional[mx.distributed.Group] = None, quantize: int | None = None, ): config = get_config_for_model(model_id) adapter = create_adapter_for_model(config, model_id, local_path, quantize) - if group is not None: + has_layer_sharding = ( + shard_metadata.start_layer != 0 + or shard_metadata.end_layer != shard_metadata.n_layers + ) + + if group is not None and has_layer_sharding: adapter.slice_transformer_blocks( start_layer=shard_metadata.start_layer, end_layer=shard_metadata.end_layer, @@ -75,8 +80,10 @@ class DistributedImageModel: model_path = build_model_path(model_id) shard_metadata = bound_instance.bound_shard - if not isinstance(shard_metadata, PipelineShardMetadata): - raise ValueError("Expected PipelineShardMetadata for image generation") + if not isinstance(shard_metadata, (PipelineShardMetadata, CfgShardMetadata)): + raise ValueError( + "Expected PipelineShardMetadata or CfgShardMetadata for image generation" + ) is_distributed = ( len(bound_instance.instance.shard_assignments.node_to_runner) > 1 diff --git a/src/exo/worker/engines/image/models/base.py b/src/exo/worker/engines/image/models/base.py index 90439823..f77ea882 100644 --- a/src/exo/worker/engines/image/models/base.py +++ b/src/exo/worker/engines/image/models/base.py @@ -86,6 +86,27 @@ class PromptData(ABC): """ ... + @abstractmethod + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + """Get embeddings for a single CFG branch (positive or negative). + + Used for sequential CFG and CFG parallel modes where we process + one branch at a time instead of batching. + + Args: + positive: True for positive prompt, False for negative prompt + + Returns: + Tuple of: + - embeds: [1, seq, hidden] prompt embeddings + - mask: [1, seq] attention mask or None + - pooled: [1, hidden] pooled embeddings or None + - conditioning_latents: [1, latent_seq, latent_dim] or None + """ + ... + class ModelAdapter(ABC, Generic[ModelT, TransformerT]): _config: ImageModelConfig diff --git a/src/exo/worker/engines/image/models/flux/adapter.py b/src/exo/worker/engines/image/models/flux/adapter.py index be9b43f8..1aa510da 100644 --- a/src/exo/worker/engines/image/models/flux/adapter.py +++ b/src/exo/worker/engines/image/models/flux/adapter.py @@ -64,6 +64,12 @@ class FluxPromptData(PromptData): ) -> tuple[mx.array, mx.array, mx.array | None, mx.array | None] | None: return None + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + """Flux doesn't use CFG, but we return positive data for compatibility.""" + return (self._prompt_embeds, None, self._pooled_prompt_embeds, None) + class FluxModelAdapter(ModelAdapter[Flux1, Transformer]): def __init__( diff --git a/src/exo/worker/engines/image/models/qwen/adapter.py b/src/exo/worker/engines/image/models/qwen/adapter.py index d9f009ec..e88d2a75 100644 --- a/src/exo/worker/engines/image/models/qwen/adapter.py +++ b/src/exo/worker/engines/image/models/qwen/adapter.py @@ -133,6 +133,24 @@ class QwenPromptData(PromptData): return batched_embeds, batched_mask, None, cond_latents + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + if positive: + return ( + self._prompt_embeds, + self._prompt_mask, + None, + self.conditioning_latents, + ) + else: + return ( + self._negative_prompt_embeds, + self._negative_prompt_mask, + None, + self.conditioning_latents, + ) + class QwenModelAdapter(ModelAdapter[QwenImage, QwenTransformer]): """Adapter for Qwen-Image model. diff --git a/src/exo/worker/engines/image/models/qwen/config.py b/src/exo/worker/engines/image/models/qwen/config.py index d5da1bac..4ec2cb35 100644 --- a/src/exo/worker/engines/image/models/qwen/config.py +++ b/src/exo/worker/engines/image/models/qwen/config.py @@ -12,7 +12,7 @@ QWEN_IMAGE_CONFIG = ImageModelConfig( ), ), default_steps={"low": 10, "medium": 25, "high": 50}, - num_sync_steps_factor=0.125, # ~3 sync steps for medium (30 steps) + num_sync_steps_factor=0.25, guidance_scale=3.5, # Set to None or < 1.0 to disable CFG ) @@ -24,6 +24,6 @@ QWEN_IMAGE_EDIT_CONFIG = ImageModelConfig( ), ), default_steps={"low": 10, "medium": 25, "high": 50}, - num_sync_steps_factor=0.125, + num_sync_steps_factor=0.25, guidance_scale=3.5, ) diff --git a/src/exo/worker/engines/image/models/qwen/edit_adapter.py b/src/exo/worker/engines/image/models/qwen/edit_adapter.py index e327eb0c..4a88a4e3 100644 --- a/src/exo/worker/engines/image/models/qwen/edit_adapter.py +++ b/src/exo/worker/engines/image/models/qwen/edit_adapter.py @@ -153,6 +153,24 @@ class QwenEditPromptData(PromptData): return batched_embeds, batched_mask, None, batched_cond_latents + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + if positive: + return ( + self._prompt_embeds, + self._prompt_mask, + None, + self._conditioning_latents, + ) + else: + return ( + self._negative_prompt_embeds, + self._negative_prompt_mask, + None, + self._conditioning_latents, + ) + class QwenEditModelAdapter(ModelAdapter[QwenImageEdit, QwenTransformer]): """Adapter for Qwen-Image-Edit model. diff --git a/src/exo/worker/engines/image/pipeline/runner.py b/src/exo/worker/engines/image/pipeline/runner.py index e1f65efd..f7054763 100644 --- a/src/exo/worker/engines/image/pipeline/runner.py +++ b/src/exo/worker/engines/image/pipeline/runner.py @@ -1,5 +1,7 @@ +from collections.abc import Iterator +from dataclasses import dataclass from math import ceil -from typing import Any, Optional +from typing import Any, Optional, final import mlx.core as mx from mflux.models.common.config.config import Config @@ -11,7 +13,7 @@ from exo.shared.tracing import ( clear_trace_buffer, trace, ) -from exo.shared.types.worker.shards import PipelineShardMetadata +from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models.base import ( ModelAdapter, @@ -25,6 +27,16 @@ from exo.worker.engines.image.pipeline.block_wrapper import ( ) +@final +@dataclass(frozen=True) +class CfgBranch: + positive: bool + embeds: mx.array + mask: mx.array | None + pooled: mx.array | None + cond_latents: mx.array | None + + def calculate_patch_heights( latent_height: int, num_patches: int ) -> tuple[list[int], int]: @@ -70,29 +82,18 @@ class DiffusionRunner: config: ImageModelConfig, adapter: ModelAdapter[Any, Any], group: Optional[mx.distributed.Group], - shard_metadata: PipelineShardMetadata, + shard_metadata: PipelineShardMetadata | CfgShardMetadata, num_patches: Optional[int] = None, ): self.config = config self.adapter = adapter self.group = group - if group is None: - self.rank = 0 - self.world_size = 1 - self.next_rank = 0 - self.prev_rank = 0 - self.start_layer = 0 - self.end_layer = config.total_blocks - else: - self.rank = shard_metadata.device_rank - self.world_size = shard_metadata.world_size - self.next_rank = (self.rank + 1) % self.world_size - self.prev_rank = (self.rank - 1 + self.world_size) % self.world_size - self.start_layer = shard_metadata.start_layer - self.end_layer = shard_metadata.end_layer + self._init_cfg_topology(shard_metadata) - self.num_patches = num_patches if num_patches else max(1, self.world_size) + self.num_patches = ( + num_patches if num_patches else max(1, self.pipeline_world_size) + ) self.total_joint = config.joint_block_count self.total_single = config.single_block_count @@ -102,6 +103,97 @@ class DiffusionRunner: self._compute_assigned_blocks() + def _init_cfg_topology( + self, shard_metadata: PipelineShardMetadata | CfgShardMetadata + ) -> None: + """Initialize CFG and pipeline topology from shard metadata. + + Both CfgShardMetadata and PipelineShardMetadata represent pipeline parallel + execution. CFG adds a second parallel pipeline for negative prompt processing, + but within each pipeline group the communication pattern is identical. + """ + if self.group is None: + # Single node - no distributed communication + self.rank = 0 + self.world_size = 1 + self.start_layer = 0 + self.end_layer = self.config.total_blocks + self.cfg_rank = 0 + self.cfg_world_size = 1 + self.cfg_parallel = False + self.pipeline_rank = 0 + self.pipeline_world_size = 1 + self.next_pipeline_rank: int | None = None + self.prev_pipeline_rank: int | None = None + self.cfg_peer_rank: int | None = None + self.first_pipeline_rank: int = 0 + self.last_pipeline_rank: int = 0 + return + + # Common fields from base metadata + self.rank = shard_metadata.device_rank + self.world_size = shard_metadata.world_size + self.start_layer = shard_metadata.start_layer + self.end_layer = shard_metadata.end_layer + + if isinstance(shard_metadata, CfgShardMetadata): + # CFG parallel: two independent pipelines + self.cfg_rank = shard_metadata.cfg_rank + self.cfg_world_size = shard_metadata.cfg_world_size + self.cfg_parallel = True + self.pipeline_rank = shard_metadata.pipeline_rank + self.pipeline_world_size = shard_metadata.pipeline_world_size + else: + # Pure pipeline: single pipeline group, sequential CFG + self.cfg_rank = 0 + self.cfg_world_size = 1 + self.cfg_parallel = False + self.pipeline_rank = shard_metadata.device_rank + self.pipeline_world_size = shard_metadata.world_size + + # Pipeline neighbor computation (same logic for both types) + is_first = self.pipeline_rank == 0 + is_last = self.pipeline_rank == self.pipeline_world_size - 1 + + self.next_pipeline_rank = ( + None + if is_last + else self._device_rank_for(self.cfg_rank, self.pipeline_rank + 1) + ) + self.prev_pipeline_rank = ( + None + if is_first + else self._device_rank_for(self.cfg_rank, self.pipeline_rank - 1) + ) + + # CFG peer is the corresponding last stage in the other CFG group + if self.cfg_parallel and is_last: + other_cfg_rank = 1 - self.cfg_rank + self.cfg_peer_rank = self._device_rank_for( + other_cfg_rank, self.pipeline_rank + ) + else: + self.cfg_peer_rank = None + + # First/last pipeline ranks for ring communication (latent broadcast) + self.first_pipeline_rank = self._device_rank_for(self.cfg_rank, 0) + self.last_pipeline_rank = self._device_rank_for( + self.cfg_rank, self.pipeline_world_size - 1 + ) + + def _device_rank_for(self, cfg_rank: int, pipeline_rank: int) -> int: + """Convert (cfg_rank, pipeline_rank) to device_rank in the ring topology. + + Ring layout: [cfg0_pipe0, cfg0_pipe1, ..., cfg1_pipeN-1, cfg1_pipeN-2, ..., cfg1_pipe0] + Group 0 is in ascending order, group 1 is reversed so last stages are neighbors. + """ + if not self.cfg_parallel: + return pipeline_rank + if cfg_rank == 0: + return pipeline_rank + else: + return self.world_size - 1 - pipeline_rank + def _compute_assigned_blocks(self) -> None: """Determine which joint/single blocks this stage owns.""" start = self.start_layer @@ -138,11 +230,11 @@ class DiffusionRunner: @property def is_first_stage(self) -> bool: - return self.rank == 0 + return self.pipeline_rank == 0 @property def is_last_stage(self) -> bool: - return self.rank == self.world_size - 1 + return self.pipeline_rank == self.pipeline_world_size - 1 @property def is_distributed(self) -> bool: @@ -153,6 +245,97 @@ class DiffusionRunner: return self._guidance_override return self.config.guidance_scale + def _get_cfg_branches(self, prompt_data: PromptData) -> Iterator[CfgBranch]: + """Yield the CFG branches this node should process. + + - No CFG: yields one branch (positive) + - CFG parallel: yields one branch (our assigned branch) + - Sequential CFG: yields two branches (positive, then negative) + """ + if not self.adapter.needs_cfg: + embeds, mask, pooled, cond = prompt_data.get_cfg_branch_data(positive=True) + yield CfgBranch( + positive=True, + embeds=embeds, + mask=mask, + pooled=pooled, + cond_latents=cond, + ) + elif self.cfg_parallel: + positive = self.cfg_rank == 0 + embeds, mask, pooled, cond = prompt_data.get_cfg_branch_data(positive) + yield CfgBranch( + positive=positive, + embeds=embeds, + mask=mask, + pooled=pooled, + cond_latents=cond, + ) + else: + pos_embeds, pos_mask, pos_pooled, pos_cond = ( + prompt_data.get_cfg_branch_data(positive=True) + ) + yield CfgBranch( + positive=True, + embeds=pos_embeds, + mask=pos_mask, + pooled=pos_pooled, + cond_latents=pos_cond, + ) + neg_embeds, neg_mask, neg_pooled, neg_cond = ( + prompt_data.get_cfg_branch_data(positive=False) + ) + yield CfgBranch( + positive=False, + embeds=neg_embeds, + mask=neg_mask, + pooled=neg_pooled, + cond_latents=neg_cond, + ) + + def _combine_cfg_results(self, results: list[tuple[bool, mx.array]]) -> mx.array: + if len(results) == 1: + positive, noise = results[0] + if self.cfg_parallel and self.is_last_stage: + # TODO(ciaran): try to remove + mx.eval(noise) + return self._exchange_and_apply_guidance(noise, positive) + return noise + + noise_neg = next(n for p, n in results if not p) + noise_pos = next(n for p, n in results if p) + return self._apply_guidance(noise_pos, noise_neg) + + def _exchange_and_apply_guidance( + self, noise: mx.array, is_positive: bool + ) -> mx.array: + assert self.group is not None + assert self.cfg_peer_rank is not None + + if is_positive: + noise = mx.distributed.send(noise, self.cfg_peer_rank, group=self.group) + mx.async_eval(noise) + noise_neg = mx.distributed.recv_like( + noise, self.cfg_peer_rank, group=self.group + ) + mx.eval(noise_neg) + noise_pos = noise + else: + noise_pos = mx.distributed.recv_like( + noise, self.cfg_peer_rank, group=self.group + ) + mx.eval(noise_pos) + noise = mx.distributed.send(noise, self.cfg_peer_rank, group=self.group) + mx.async_eval(noise) + noise_neg = noise + + return self._apply_guidance(noise_pos, noise_neg) + + def _apply_guidance(self, noise_pos: mx.array, noise_neg: mx.array) -> mx.array: + scale = self._get_effective_guidance_scale() + assert scale is not None + return self.adapter.apply_guidance(noise_pos, noise_neg, scale) + def _ensure_wrappers( self, text_seq_len: int, @@ -470,7 +653,9 @@ class DiffusionRunner: ) -> mx.array: if self.group is None: return self._single_node_step(t, config, latents, prompt_data) - elif t < config.init_time_step + num_sync_steps: + elif ( + self.pipeline_world_size == 1 or t < config.init_time_step + num_sync_steps + ): with trace(name=f"sync {t}", rank=self.rank, category="sync"): return self._sync_pipeline_step( t, @@ -496,42 +681,29 @@ class DiffusionRunner: prompt_data: PromptData, ) -> mx.array: cond_image_grid = prompt_data.cond_image_grid - needs_cfg = self.adapter.needs_cfg + results: list[tuple[bool, mx.array]] = [] + + for branch in self._get_cfg_branches(prompt_data): + # Reset caches before each branch to ensure no state contamination + self._reset_all_caches() - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, cond_latents = batched_data pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds - ) - step_latents = mx.concatenate([latents, latents], axis=0) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - cond_latents = prompt_data.conditioning_latents - step_latents = latents - - noise = self._forward_pass( - step_latents, - prompt_embeds, - pooled_embeds, - t=t, - config=config, - encoder_hidden_states_mask=encoder_mask, - cond_image_grid=cond_image_grid, - conditioning_latents=cond_latents, - ) - - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale=guidance_scale + branch.pooled if branch.pooled is not None else branch.embeds ) + noise = self._forward_pass( + latents, + branch.embeds, + pooled_embeds, + t=t, + config=config, + encoder_hidden_states_mask=branch.mask, + cond_image_grid=cond_image_grid, + conditioning_latents=branch.cond_latents, + ) + results.append((branch.positive, noise)) + + noise = self._combine_cfg_results(results) return config.scheduler.step(noise=noise, timestep=t, latents=latents) # pyright: ignore[reportAny] def _create_patches( @@ -582,7 +754,7 @@ class DiffusionRunner: ) text_embeddings = self.adapter.compute_text_embeddings( - t, config, pooled_prompt_embeds + t, config, pooled_prompt_embeds, hidden_states=hidden_states ) image_rotary_embeddings = self.adapter.compute_rotary_embeddings( prompt_embeds, @@ -594,19 +766,22 @@ class DiffusionRunner: if self.has_joint_blocks: if not self.is_first_stage: + assert self.prev_pipeline_rank is not None with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.recv( (batch_size, num_img_tokens, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) encoder_hidden_states = mx.distributed.recv( (batch_size, text_seq_len, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(hidden_states, encoder_hidden_states) @@ -639,34 +814,45 @@ class DiffusionRunner: if self.has_single_blocks or self.is_last_stage: hidden_states = concatenated else: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): concatenated = mx.distributed.send( - concatenated, self.next_rank, group=self.group + concatenated, self.next_pipeline_rank, group=self.group ) mx.async_eval(concatenated) elif self.has_joint_blocks and not self.is_last_stage: assert encoder_hidden_states is not None - with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"): + assert self.next_pipeline_rank is not None + with trace( + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", + ): hidden_states = mx.distributed.send( - hidden_states, self.next_rank, group=self.group + hidden_states, self.next_pipeline_rank, group=self.group ) encoder_hidden_states = mx.distributed.send( - encoder_hidden_states, self.next_rank, group=self.group + encoder_hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(hidden_states, encoder_hidden_states) if self.has_single_blocks: if not self.owns_concat_stage and not self.is_first_stage: + assert self.prev_pipeline_rank is not None with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.recv( (batch_size, text_seq_len + num_img_tokens, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(hidden_states) @@ -689,11 +875,14 @@ class DiffusionRunner: mx.eval(hidden_states) if not self.is_last_stage: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.send( - hidden_states, self.next_rank, group=self.group + hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(hidden_states) @@ -716,83 +905,67 @@ class DiffusionRunner: kontext_image_ids: mx.array | None = None, ) -> mx.array: prev_latents = hidden_states - needs_cfg = self.adapter.needs_cfg cond_image_grid = prompt_data.cond_image_grid scaled_hidden_states = config.scheduler.scale_model_input(hidden_states, t) # pyright: ignore[reportAny] original_latent_tokens: int = scaled_hidden_states.shape[1] # pyright: ignore[reportAny] - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, cond_latents = batched_data + results: list[tuple[bool, mx.array]] = [] + + for branch in self._get_cfg_branches(prompt_data): pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds + branch.pooled if branch.pooled is not None else branch.embeds ) - step_latents = mx.concatenate( - [scaled_hidden_states, scaled_hidden_states], axis=0 + + cond_latents = branch.cond_latents + if cond_latents is not None: + num_img_tokens: int = original_latent_tokens + cond_latents.shape[1] + else: + num_img_tokens = original_latent_tokens + + step_latents: mx.array = scaled_hidden_states # pyright: ignore[reportAny] + if self.is_first_stage and cond_latents is not None: + step_latents = mx.concatenate([step_latents, cond_latents], axis=1) + + text_seq_len = branch.embeds.shape[1] + self._ensure_wrappers(text_seq_len, branch.mask) + + noise = self._run_sync_pass( + t, + config, + step_latents, + branch.embeds, + pooled_embeds, + branch.mask, + cond_image_grid, + kontext_image_ids, + num_img_tokens, + original_latent_tokens, + cond_latents, ) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - cond_latents = prompt_data.conditioning_latents - step_latents = scaled_hidden_states # pyright: ignore[reportAny] - if cond_latents is not None: - num_img_tokens: int = original_latent_tokens + cond_latents.shape[1] - else: - num_img_tokens = original_latent_tokens - - if self.is_first_stage and cond_latents is not None: - step_latents = mx.concatenate([step_latents, cond_latents], axis=1) - - text_seq_len = prompt_embeds.shape[1] - self._ensure_wrappers(text_seq_len, encoder_mask) - - noise = self._run_sync_pass( - t, - config, - step_latents, - prompt_embeds, - pooled_embeds, - encoder_mask, - cond_image_grid, - kontext_image_ids, - num_img_tokens, - original_latent_tokens, - cond_latents, - ) + if self.is_last_stage: + assert noise is not None + results.append((branch.positive, noise)) if self.is_last_stage: - assert noise is not None - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale - ) + noise = self._combine_cfg_results(results) hidden_states = config.scheduler.step( # pyright: ignore[reportAny] noise=noise, timestep=t, latents=prev_latents ) if not self.is_first_stage: - with trace(name="send 0", rank=self.rank, category="comms"): - hidden_states = mx.distributed.send( - hidden_states, 0, group=self.group - ) - mx.async_eval(hidden_states) + hidden_states = mx.distributed.send( + hidden_states, self.first_pipeline_rank, group=self.group + ) + mx.async_eval(hidden_states) elif self.is_first_stage: - with trace( - name=f"recv {self.world_size - 1}", rank=self.rank, category="comms" - ): - hidden_states = mx.distributed.recv_like( - prev_latents, src=self.world_size - 1, group=self.group - ) - mx.eval(hidden_states) + hidden_states = mx.distributed.recv_like( + prev_latents, src=self.last_pipeline_rank, group=self.group + ) + mx.eval(hidden_states) else: hidden_states = prev_latents @@ -809,39 +982,10 @@ class DiffusionRunner: kontext_image_ids: mx.array | None = None, ) -> mx.array: patch_latents, token_indices = self._create_patches(latents, config) - needs_cfg = self.adapter.needs_cfg cond_image_grid = prompt_data.cond_image_grid - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, _ = batched_data - pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds - ) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - - text_seq_len = prompt_embeds.shape[1] - self._ensure_wrappers(text_seq_len, encoder_mask) - self._set_text_seq_len(text_seq_len) - - if self.joint_block_wrappers: - for wrapper in self.joint_block_wrappers: - wrapper.set_encoder_mask(encoder_mask) - - text_embeddings = self.adapter.compute_text_embeddings(t, config, pooled_embeds) - image_rotary_embeddings = self.adapter.compute_rotary_embeddings( - prompt_embeds, - config, - encoder_hidden_states_mask=encoder_mask, - cond_image_grid=cond_image_grid, - kontext_image_ids=kontext_image_ids, - ) - prev_patch_latents = [p for p in patch_latents] + encoder_hidden_states: mx.array | None = None for patch_idx in range(len(patch_latents)): @@ -853,34 +997,57 @@ class DiffusionRunner: and not is_first_async_step ): with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.last_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv_like( - patch, src=self.prev_rank, group=self.group + patch, src=self.last_pipeline_rank, group=self.group ) mx.eval(patch) - step_patch = mx.concatenate([patch, patch], axis=0) if needs_cfg else patch + results: list[tuple[bool, mx.array]] = [] - noise, encoder_hidden_states = self._run_single_patch_pass( - patch=step_patch, - patch_idx=patch_idx, - token_indices=token_indices[patch_idx], - prompt_embeds=prompt_embeds, - text_embeddings=text_embeddings, - image_rotary_embeddings=image_rotary_embeddings, - encoder_hidden_states=encoder_hidden_states, - ) + for branch in self._get_cfg_branches(prompt_data): + pooled_embeds = ( + branch.pooled if branch.pooled is not None else branch.embeds + ) + + text_seq_len = branch.embeds.shape[1] + self._ensure_wrappers(text_seq_len, branch.mask) + self._set_text_seq_len(text_seq_len) + + if self.joint_block_wrappers: + for wrapper in self.joint_block_wrappers: + wrapper.set_encoder_mask(branch.mask) + + text_embeddings = self.adapter.compute_text_embeddings( + t, config, pooled_embeds + ) + image_rotary_embeddings = self.adapter.compute_rotary_embeddings( + branch.embeds, + config, + encoder_hidden_states_mask=branch.mask, + cond_image_grid=cond_image_grid, + kontext_image_ids=kontext_image_ids, + ) + + noise, encoder_hidden_states = self._run_single_patch_pass( + patch=patch, + patch_idx=patch_idx, + token_indices=token_indices[patch_idx], + prompt_embeds=branch.embeds, + text_embeddings=text_embeddings, + image_rotary_embeddings=image_rotary_embeddings, + encoder_hidden_states=encoder_hidden_states, + ) + + if self.is_last_stage: + assert noise is not None + results.append((branch.positive, noise)) if self.is_last_stage: - assert noise is not None - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale - ) + noise = self._combine_cfg_results(results) patch_latents[patch_idx] = config.scheduler.step( # pyright: ignore[reportAny] noise=noise, @@ -890,10 +1057,14 @@ class DiffusionRunner: if not self.is_first_stage and t != config.num_inference_steps - 1: with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.first_pipeline_rank}", + rank=self.rank, + category="comms", ): patch_latents[patch_idx] = mx.distributed.send( - patch_latents[patch_idx], self.next_rank, group=self.group + patch_latents[patch_idx], + self.first_pipeline_rank, + group=self.group, ) mx.async_eval(patch_latents[patch_idx]) @@ -933,26 +1104,31 @@ class DiffusionRunner: if self.has_joint_blocks: if not self.is_first_stage: + assert self.prev_pipeline_rank is not None patch_len = patch.shape[1] with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv( (batch_size, patch_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(patch) if patch_idx == 0: with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): encoder_hidden_states = mx.distributed.recv( (batch_size, text_seq_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(encoder_hidden_states) @@ -988,39 +1164,54 @@ class DiffusionRunner: if self.has_single_blocks or self.is_last_stage: patch = patch_concat else: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): patch_concat = mx.distributed.send( - patch_concat, self.next_rank, group=self.group + patch_concat, self.next_pipeline_rank, group=self.group ) mx.async_eval(patch_concat) elif self.has_joint_blocks and not self.is_last_stage: - with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"): - patch = mx.distributed.send(patch, self.next_rank, group=self.group) + assert self.next_pipeline_rank is not None + with trace( + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", + ): + patch = mx.distributed.send( + patch, self.next_pipeline_rank, group=self.group + ) mx.async_eval(patch) if patch_idx == 0: assert encoder_hidden_states is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): encoder_hidden_states = mx.distributed.send( - encoder_hidden_states, self.next_rank, group=self.group + encoder_hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(encoder_hidden_states) if self.has_single_blocks: if not self.owns_concat_stage and not self.is_first_stage: + assert self.prev_pipeline_rank is not None patch_len = patch.shape[1] with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv( (batch_size, text_seq_len + patch_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(patch) @@ -1043,15 +1234,20 @@ class DiffusionRunner: mx.eval(patch) if not self.is_last_stage: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): - patch = mx.distributed.send(patch, self.next_rank, group=self.group) + patch = mx.distributed.send( + patch, self.next_pipeline_rank, group=self.group + ) mx.async_eval(patch) noise: mx.array | None = None if self.is_last_stage: - patch = patch[:, text_seq_len:, :] - noise = self.adapter.final_projection(patch, text_embeddings) + patch_img_only = patch[:, text_seq_len:, :] + noise = self.adapter.final_projection(patch_img_only, text_embeddings) return noise, encoder_hidden_states diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index 5dae48a4..e12aa185 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -48,6 +48,7 @@ from exo.shared.types.worker.instances import ( MlxRingInstance, ) from exo.shared.types.worker.shards import ( + CfgShardMetadata, PipelineShardMetadata, ShardMetadata, TensorShardMetadata, @@ -274,6 +275,11 @@ def shard_and_load( logger.info(f"loading model from {model_path} with pipeline parallelism") model = pipeline_auto_parallel(model, group, shard_metadata) eval_with_timeout(model.parameters(), timeout_seconds, on_timeout) + case CfgShardMetadata(): + raise ValueError( + "CfgShardMetadata is not supported for text model loading - " + "this metadata type is only for image generation models" + ) # TODO: Do we need this? mx.eval(model) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 0b527318..3232e0e6 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -66,7 +66,11 @@ from exo.shared.types.worker.runners import ( RunnerStatus, RunnerWarmingUp, ) -from exo.shared.types.worker.shards import ShardMetadata +from exo.shared.types.worker.shards import ( + CfgShardMetadata, + PipelineShardMetadata, + ShardMetadata, +) from exo.utils.channels import MpReceiver, MpSender from exo.worker.engines.image import ( DistributedImageModel, @@ -87,6 +91,22 @@ from exo.worker.engines.mlx.utils_mlx import ( from exo.worker.runner.bootstrap import logger +def _is_primary_output_node(shard_metadata: ShardMetadata) -> bool: + """Check if this node is the primary output node for image generation. + + For CFG models: the last pipeline stage in CFG group 0 (positive prompt). + For non-CFG models: the last pipeline stage. + """ + if isinstance(shard_metadata, CfgShardMetadata): + is_pipeline_last = ( + shard_metadata.pipeline_rank == shard_metadata.pipeline_world_size - 1 + ) + return is_pipeline_last and shard_metadata.cfg_rank == 0 + elif isinstance(shard_metadata, PipelineShardMetadata): + return shard_metadata.device_rank == shard_metadata.world_size - 1 + return False + + def main( bound_instance: BoundInstance, event_sender: MpSender[Event], @@ -367,14 +387,11 @@ def main( ) try: - # Generate images using the image generation backend - # Track image_index for final images only image_index = 0 for response in generate_image(model=model, task=task_params): - if ( - shard_metadata.device_rank - == shard_metadata.world_size - 1 - ): + is_primary_output = _is_primary_output_node(shard_metadata) + + if is_primary_output: match response: case PartialImageResponse(): logger.info( @@ -399,7 +416,7 @@ def main( image_index += 1 # can we make this more explicit? except Exception as e: - if shard_metadata.device_rank == shard_metadata.world_size - 1: + if _is_primary_output_node(shard_metadata): event_sender.send( ChunkGenerated( command_id=command_id, @@ -434,10 +451,7 @@ def main( try: image_index = 0 for response in generate_image(model=model, task=task_params): - if ( - shard_metadata.device_rank - == shard_metadata.world_size - 1 - ): + if _is_primary_output_node(shard_metadata): match response: case PartialImageResponse(): logger.info( @@ -461,7 +475,7 @@ def main( ) image_index += 1 except Exception as e: - if shard_metadata.device_rank == shard_metadata.world_size - 1: + if _is_primary_output_node(shard_metadata): event_sender.send( ChunkGenerated( command_id=command_id, From 221640a65b4554985c99300775ea5695ccfd2238 Mon Sep 17 00:00:00 2001 From: rltakashige Date: Thu, 5 Feb 2026 12:00:37 +0000 Subject: [PATCH 7/9] Acknowledge task after runner status is updated (#1381) ## Motivation Duplicate tasks are still observed. ## Changes Moved task acknowledgement to after the runner has changed its status. ## Why It Works Tasks now remain pending until the runner has updated its status. ## Test Plan ### Manual Testing Seems to work fine from manual testing. Hard to test a race condition though. ### Automated Testing Updated the event ordering test. --- src/exo/worker/runner/runner.py | 10 +++++++++- .../tests/unittests/test_runner/test_event_ordering.py | 10 +++++----- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 3232e0e6..109ea219 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -145,7 +145,6 @@ def main( event_sender.send( TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Running) ) - event_sender.send(TaskAcknowledged(task_id=task.task_id)) match task: case ConnectToGroup() if isinstance( current_status, (RunnerIdle, RunnerFailed) @@ -157,6 +156,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) group = initialize_mlx(bound_instance) logger.info("runner connected") @@ -173,6 +173,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) def on_model_load_timeout() -> None: event_sender.send( @@ -215,6 +216,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) logger.info(f"warming up inference for instance: {instance}") if ModelTask.TextGeneration in shard_metadata.model_card.tasks: @@ -254,6 +256,8 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) + assert model and not isinstance(model, DistributedImageModel) assert tokenizer @@ -385,6 +389,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) try: image_index = 0 @@ -447,6 +452,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) try: image_index = 0 @@ -502,6 +508,8 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) + current_status = RunnerShutdown() case _: raise ValueError( 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 16a43f2b..edf5ef3a 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 @@ -201,29 +201,29 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch): TaskStatusUpdated( task_id=INITIALIZATION_TASK_ID, task_status=TaskStatus.Running ), - TaskAcknowledged(task_id=INITIALIZATION_TASK_ID), RunnerStatusUpdated( runner_id=RUNNER_1_ID, runner_status=RunnerConnecting() ), + TaskAcknowledged(task_id=INITIALIZATION_TASK_ID), TaskStatusUpdated( task_id=INITIALIZATION_TASK_ID, task_status=TaskStatus.Complete ), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerConnected()), TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=LOAD_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoading()), + TaskAcknowledged(task_id=LOAD_TASK_ID), TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Complete), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoaded()), TaskStatusUpdated(task_id=WARMUP_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=WARMUP_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerWarmingUp()), + TaskAcknowledged(task_id=WARMUP_TASK_ID), TaskStatusUpdated(task_id=WARMUP_TASK_ID, task_status=TaskStatus.Complete), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerReady()), TaskStatusUpdated( task_id=CHAT_COMPLETION_TASK_ID, task_status=TaskStatus.Running ), - TaskAcknowledged(task_id=CHAT_COMPLETION_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerRunning()), + TaskAcknowledged(task_id=CHAT_COMPLETION_TASK_ID), expected_chunk, TaskStatusUpdated( task_id=CHAT_COMPLETION_TASK_ID, task_status=TaskStatus.Complete @@ -231,10 +231,10 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch): # CHAT COMPLETION TASK SHOULD COMPLETE BEFORE RUNNER READY RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerReady()), TaskStatusUpdated(task_id=SHUTDOWN_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=SHUTDOWN_TASK_ID), RunnerStatusUpdated( runner_id=RUNNER_1_ID, runner_status=RunnerShuttingDown() ), + TaskAcknowledged(task_id=SHUTDOWN_TASK_ID), TaskStatusUpdated( task_id=SHUTDOWN_TASK_ID, task_status=TaskStatus.Complete ), From 01b86a9e814da97c0ab9b70c251c9e6caf215352 Mon Sep 17 00:00:00 2001 From: Alex Cheema <41707476+AlexCheema@users.noreply.github.com> Date: Thu, 5 Feb 2026 05:21:26 -0800 Subject: [PATCH 8/9] feat: add uncertainty visualization with token-level logprobs (#1180) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Motivation Adds uncertainty visualization to the chat interface, allowing users to see token-level confidence scores and regenerate responses from any point in the generation. This enables users to: - Understand model confidence at each token - Explore alternative completions by regenerating from uncertain tokens - Debug and analyze model behavior ## Changes ### Uncertainty Visualization - Add `TokenHeatmap` component showing token-level probability coloring - Toggle uncertainty view per message with bar chart icon - Display tooltip with probability, logprob, and top alternative tokens on hover ### Regenerate from Token - Add "Regenerate from here" button in token tooltip - Use `continue_final_message` in chat template to continue within same turn (no EOS tokens) - Add `continue_from_prefix` flag to `ChatCompletionTaskParams` ### Request Cancellation - Add `AbortController` to cancel in-flight requests when regenerating mid-generation - Handle `BrokenResourceError` server-side when client disconnects gracefully ### Additional APIs - Add Claude Messages API support (`/v1/messages`) - Add OpenAI Responses API support (`/v1/responses`) ## Why It Works - **Proper continuation**: Using `continue_final_message=True` instead of `add_generation_prompt=True` keeps the assistant turn open, allowing the model to continue naturally from the prefix without end-of-turn markers - **Clean cancellation**: AbortController aborts the HTTP request, and server catches `BrokenResourceError` to avoid crashes - **Stable hover during generation**: TokenHeatmap tracks hover by index (stable across re-renders) with longer hide delay during generation ## Test Plan ### Manual Testing - Send a message and verify logprobs are collected - Enable uncertainty view and verify token coloring based on probability - Hover over tokens to see tooltip with alternatives - Click "Regenerate from here" on a token mid-response - Verify the response continues naturally from that point - Verify aborting mid-generation and regenerating works without server crash ### Automated Testing - Added tests for Claude Messages API adapter - Added tests for OpenAI Responses API adapter 🤖 Generated with [Claude Code](https://claude.com/claude-code) --------- Co-authored-by: Claude Opus 4.5 Co-authored-by: Evan --- .../src/lib/components/ChatMessages.svelte | 72 +++- .../src/lib/components/TokenHeatmap.svelte | 236 +++++++++++++ dashboard/src/lib/stores/app.svelte.ts | 319 +++++++++++++++++- src/exo/master/adapters/chat_completions.py | 30 ++ src/exo/master/api.py | 15 + src/exo/shared/types/chunks.py | 9 +- src/exo/shared/types/text_generation.py | 2 + .../shared/types/worker/runner_response.py | 4 +- src/exo/worker/engines/mlx/constants.py | 2 + .../worker/engines/mlx/generator/generate.py | 75 +++- src/exo/worker/engines/mlx/utils_mlx.py | 9 + src/exo/worker/runner/runner.py | 2 + 12 files changed, 760 insertions(+), 15 deletions(-) create mode 100644 dashboard/src/lib/components/TokenHeatmap.svelte diff --git a/dashboard/src/lib/components/ChatMessages.svelte b/dashboard/src/lib/components/ChatMessages.svelte index 15ea088d..44b9ec0d 100644 --- a/dashboard/src/lib/components/ChatMessages.svelte +++ b/dashboard/src/lib/components/ChatMessages.svelte @@ -6,11 +6,13 @@ deleteMessage, editAndRegenerate, regenerateLastResponse, + regenerateFromToken, setEditingImage, } from "$lib/stores/app.svelte"; import type { Message } from "$lib/stores/app.svelte"; import type { MessageAttachment } from "$lib/stores/app.svelte"; import MarkdownContent from "./MarkdownContent.svelte"; + import TokenHeatmap from "./TokenHeatmap.svelte"; interface Props { class?: string; @@ -99,6 +101,23 @@ let copiedMessageId = $state(null); let expandedThinkingMessageIds = $state>(new Set()); + // Uncertainty heatmap toggle + let heatmapMessageIds = $state>(new Set()); + + function toggleHeatmap(messageId: string) { + const next = new Set(heatmapMessageIds); + if (next.has(messageId)) { + next.delete(messageId); + } else { + next.add(messageId); + } + heatmapMessageIds = next; + } + + function isHeatmapVisible(messageId: string): boolean { + return heatmapMessageIds.has(messageId); + } + function formatTimestamp(timestamp: number): string { return new Date(timestamp).toLocaleTimeString("en-US", { hour12: false, @@ -548,13 +567,23 @@ > {:else if message.content || (loading && !message.attachments?.some((a) => a.type === "generated-image"))} - - {#if loading && !message.content} - + {#if isHeatmapVisible(message.id) && message.tokens && message.tokens.length > 0} + + regenerateFromToken(message.id, tokenIndex)} + /> + {:else} + + {#if loading && !message.content} + + {/if} {/if} {/if} @@ -629,6 +658,35 @@ {/if} + + {#if message.role === "assistant" && message.tokens && message.tokens.length > 0} + + {/if} + {#if message.role === "assistant" && isLastAssistantMessage(message.id) && !loading} + {/if} + + +
+
+
+ +{/if} + + diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index 51de6c66..6fdb0c7c 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -242,6 +242,19 @@ export interface MessageAttachment { mimeType?: string; } +export interface TopLogprob { + token: string; + logprob: number; + bytes: number[] | null; +} + +export interface TokenData { + token: string; + logprob: number; + probability: number; + topLogprobs: TopLogprob[]; +} + export interface Message { id: string; role: "user" | "assistant" | "system"; @@ -253,6 +266,7 @@ export interface Message { tps?: number; // Tokens per second (for assistant messages) requestType?: "chat" | "image-generation" | "image-editing"; sourceImageDataUrl?: string; // For image editing regeneration + tokens?: TokenData[]; } export interface Conversation { @@ -540,7 +554,18 @@ class AppStore { */ private saveConversationsToStorage() { try { - localStorage.setItem(STORAGE_KEY, JSON.stringify(this.conversations)); + // Strip tokens from messages before saving to avoid bloating localStorage + const stripped = this.conversations.map((conv) => ({ + ...conv, + messages: conv.messages.map((msg) => { + if (msg.tokens) { + const { tokens: _, ...rest } = msg; + return rest; + } + return msg; + }), + })); + localStorage.setItem(STORAGE_KEY, JSON.stringify(stripped)); } catch (error) { console.error("Failed to save conversations:", error); } @@ -1445,6 +1470,213 @@ class AppStore { } } + /** + * Regenerate response from a specific token index. + * Truncates the assistant message at the given token and re-generates from there. + */ + async regenerateFromToken( + messageId: string, + tokenIndex: number, + ): Promise { + if (this.isLoading) return; + + const targetConversationId = this.activeConversationId; + if (!targetConversationId) return; + + const msgIndex = this.messages.findIndex((m) => m.id === messageId); + if (msgIndex === -1) return; + + const msg = this.messages[msgIndex]; + if ( + msg.role !== "assistant" || + !msg.tokens || + tokenIndex >= msg.tokens.length + ) + return; + + // Keep tokens up to (not including) the specified index + const tokensToKeep = msg.tokens.slice(0, tokenIndex); + const prefixText = tokensToKeep.map((t) => t.token).join(""); + + // Remove all messages after this assistant message + this.messages = this.messages.slice(0, msgIndex + 1); + + // Update the message to show the prefix + this.messages[msgIndex].content = prefixText; + this.messages[msgIndex].tokens = tokensToKeep; + this.updateActiveConversation(); + + // Set up for continuation - modify the existing message in place + this.isLoading = true; + this.currentResponse = prefixText; + this.ttftMs = null; + this.tps = null; + this.totalTokens = tokensToKeep.length; + + try { + // Build messages for API - include the partial assistant message + const systemPrompt = { + role: "system" as const, + content: + "You are a helpful AI assistant. Respond directly and concisely. Do not show your reasoning or thought process.", + }; + + const apiMessages = [ + systemPrompt, + ...this.messages.map((m) => { + let msgContent = m.content; + if (m.attachments) { + for (const attachment of m.attachments) { + if (attachment.type === "text" && attachment.content) { + msgContent += `\n\n[File: ${attachment.name}]\n\`\`\`\n${attachment.content}\n\`\`\``; + } + } + } + return { role: m.role, content: msgContent }; + }), + ]; + + const modelToUse = this.getModelForRequest(); + if (!modelToUse) { + throw new Error("No model available"); + } + + const requestStartTime = performance.now(); + let firstTokenTime: number | null = null; + let tokenCount = tokensToKeep.length; + + const response = await fetch("/v1/chat/completions", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + model: modelToUse, + messages: apiMessages, + stream: true, + logprobs: true, + top_logprobs: 5, + }), + }); + + if (!response.ok) { + const errorText = await response.text(); + throw new Error(`API error: ${response.status} - ${errorText}`); + } + + const reader = response.body?.getReader(); + if (!reader) throw new Error("No response body"); + + let fullContent = prefixText; + const collectedTokens: TokenData[] = [...tokensToKeep]; + + interface ChatCompletionChunk { + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; + } + + await this.parseSSEStream( + reader, + targetConversationId, + (parsed) => { + const choice = parsed.choices?.[0]; + const delta = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + + if (delta) { + if (firstTokenTime === null) { + firstTokenTime = performance.now(); + this.ttftMs = firstTokenTime - requestStartTime; + } + + tokenCount += 1; + this.totalTokens = tokenCount; + + if (firstTokenTime !== null && tokenCount > tokensToKeep.length) { + const elapsed = performance.now() - firstTokenTime; + this.tps = ((tokenCount - tokensToKeep.length) / elapsed) * 1000; + } + + fullContent += delta; + const { displayContent, thinkingContent } = + this.stripThinkingTags(fullContent); + + if (this.activeConversationId === targetConversationId) { + this.currentResponse = displayContent; + } + + // Update existing message in place + this.updateConversationMessage( + targetConversationId, + messageId, + (m) => { + m.content = displayContent; + m.thinking = thinkingContent || undefined; + m.tokens = [...collectedTokens]; + }, + ); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + }, + ); + + // Final update + if (this.conversationExists(targetConversationId)) { + const { displayContent, thinkingContent } = + this.stripThinkingTags(fullContent); + this.updateConversationMessage(targetConversationId, messageId, (m) => { + m.content = displayContent; + m.thinking = thinkingContent || undefined; + m.tokens = [...collectedTokens]; + if (this.ttftMs !== null) m.ttftMs = this.ttftMs; + if (this.tps !== null) m.tps = this.tps; + }); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + } catch (error) { + console.error("Error regenerating from token:", error); + if (this.conversationExists(targetConversationId)) { + this.updateConversationMessage(targetConversationId, messageId, (m) => { + m.content = `${prefixText}\n\nError: ${error instanceof Error ? error.message : "Unknown error"}`; + }); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + } finally { + this.isLoading = false; + this.currentResponse = ""; + this.saveConversationsToStorage(); + } + } + /** * Helper method to regenerate a chat completion response */ @@ -1513,6 +1745,8 @@ class AppStore { model: modelToUse, messages: apiMessages, stream: true, + logprobs: true, + top_logprobs: 5, }), }); @@ -1527,16 +1761,49 @@ class AppStore { } let streamedContent = ""; + const collectedTokens: TokenData[] = []; interface ChatCompletionChunk { - choices?: Array<{ delta?: { content?: string } }>; + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; } await this.parseSSEStream( reader, targetConversationId, (parsed) => { - const delta = parsed.choices?.[0]?.delta?.content; + const choice = parsed.choices?.[0]; + const delta = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + if (delta) { streamedContent += delta; const { displayContent, thinkingContent } = @@ -1554,6 +1821,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1572,6 +1840,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1914,6 +2183,8 @@ class AppStore { messages: apiMessages, temperature: 0.7, stream: true, + logprobs: true, + top_logprobs: 5, }), }); @@ -1930,14 +2201,48 @@ class AppStore { let streamedContent = ""; interface ChatCompletionChunk { - choices?: Array<{ delta?: { content?: string } }>; + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; } + const collectedTokens: TokenData[] = []; + await this.parseSSEStream( reader, targetConversationId, (parsed) => { - const tokenContent = parsed.choices?.[0]?.delta?.content; + const choice = parsed.choices?.[0]; + const tokenContent = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + if (tokenContent) { // Track first token for TTFT if (firstTokenTime === null) { @@ -1973,6 +2278,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1997,6 +2303,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; // Store performance metrics on the message if (this.ttftMs !== null) { msg.ttftMs = this.ttftMs; @@ -2693,6 +3000,8 @@ export const editMessage = (messageId: string, newContent: string) => export const editAndRegenerate = (messageId: string, newContent: string) => appStore.editAndRegenerate(messageId, newContent); export const regenerateLastResponse = () => appStore.regenerateLastResponse(); +export const regenerateFromToken = (messageId: string, tokenIndex: number) => + appStore.regenerateFromToken(messageId, tokenIndex); // Conversation actions export const conversations = () => appStore.conversations; diff --git a/src/exo/master/adapters/chat_completions.py b/src/exo/master/adapters/chat_completions.py index e144696b..3e013079 100644 --- a/src/exo/master/adapters/chat_completions.py +++ b/src/exo/master/adapters/chat_completions.py @@ -14,6 +14,8 @@ from exo.shared.types.api import ( ErrorInfo, ErrorResponse, FinishReason, + Logprobs, + LogprobsContentItem, StreamingChoiceResponse, ToolCall, ) @@ -81,6 +83,8 @@ def chat_request_to_text_generation( chat_template_messages=chat_template_messages if chat_template_messages else None, + logprobs=request.logprobs or False, + top_logprobs=request.top_logprobs, ) @@ -88,6 +92,19 @@ def chunk_to_response( chunk: TokenChunk, command_id: CommandId ) -> ChatCompletionResponse: """Convert a TokenChunk to a streaming ChatCompletionResponse.""" + # Build logprobs if available + logprobs: Logprobs | None = None + if chunk.logprob is not None: + logprobs = Logprobs( + content=[ + LogprobsContentItem( + token=chunk.text, + logprob=chunk.logprob, + top_logprobs=chunk.top_logprobs or [], + ) + ] + ) + return ChatCompletionResponse( id=command_id, created=int(time.time()), @@ -96,6 +113,7 @@ def chunk_to_response( StreamingChoiceResponse( index=0, delta=ChatCompletionMessage(role="assistant", content=chunk.text), + logprobs=logprobs, finish_reason=chunk.finish_reason, ) ], @@ -162,6 +180,7 @@ async def collect_chat_response( """Collect all token chunks and return a single ChatCompletionResponse.""" text_parts: list[str] = [] tool_calls: list[ToolCall] = [] + logprobs_content: list[LogprobsContentItem] = [] model: str | None = None finish_reason: FinishReason | None = None error_message: str | None = None @@ -176,6 +195,14 @@ async def collect_chat_response( if isinstance(chunk, TokenChunk): text_parts.append(chunk.text) + if chunk.logprob is not None: + logprobs_content.append( + LogprobsContentItem( + token=chunk.text, + logprob=chunk.logprob, + top_logprobs=chunk.top_logprobs or [], + ) + ) if isinstance(chunk, ToolCallChunk): tool_calls.extend( @@ -208,6 +235,9 @@ async def collect_chat_response( content=combined_text, tool_calls=tool_calls if tool_calls else None, ), + logprobs=Logprobs(content=logprobs_content) + if logprobs_content + else None, finish_reason=finish_reason, ) ], diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 9bd8cbcf..0ad5454c 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -627,6 +627,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_chat_response( @@ -1183,6 +1188,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_claude_response( @@ -1210,6 +1220,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_responses_response( diff --git a/src/exo/shared/types/chunks.py b/src/exo/shared/types/chunks.py index e96dbc9d..5fe9eb1c 100644 --- a/src/exo/shared/types/chunks.py +++ b/src/exo/shared/types/chunks.py @@ -2,7 +2,12 @@ from collections.abc import Generator from typing import Any, Literal from exo.shared.models.model_cards import ModelId -from exo.shared.types.api import GenerationStats, ImageGenerationStats, Usage +from exo.shared.types.api import ( + GenerationStats, + ImageGenerationStats, + TopLogprobItem, + Usage, +) from exo.utils.pydantic_ext import TaggedModel from .api import FinishReason @@ -20,6 +25,8 @@ class TokenChunk(BaseChunk): usage: Usage | None finish_reason: Literal["stop", "length", "content_filter"] | None = None stats: GenerationStats | None = None + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None class ErrorChunk(BaseChunk): diff --git a/src/exo/shared/types/text_generation.py b/src/exo/shared/types/text_generation.py index 31f97a70..3e7b89fd 100644 --- a/src/exo/shared/types/text_generation.py +++ b/src/exo/shared/types/text_generation.py @@ -40,3 +40,5 @@ class TextGenerationTaskParams(BaseModel, frozen=True): stop: str | list[str] | None = None seed: int | None = None chat_template_messages: list[dict[str, Any]] | None = None + logprobs: bool = False + top_logprobs: int | None = None diff --git a/src/exo/shared/types/worker/runner_response.py b/src/exo/shared/types/worker/runner_response.py index 5dfbe547..d1bea77e 100644 --- a/src/exo/shared/types/worker/runner_response.py +++ b/src/exo/shared/types/worker/runner_response.py @@ -6,6 +6,7 @@ from exo.shared.types.api import ( GenerationStats, ImageGenerationStats, ToolCallItem, + TopLogprobItem, Usage, ) from exo.utils.pydantic_ext import TaggedModel @@ -22,7 +23,8 @@ class TokenizedResponse(BaseRunnerResponse): class GenerationResponse(BaseRunnerResponse): text: str token: int - # logprobs: list[float] | None = None # too big. we can change to be top-k + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None finish_reason: FinishReason | None = None stats: GenerationStats | None = None usage: Usage | None diff --git a/src/exo/worker/engines/mlx/constants.py b/src/exo/worker/engines/mlx/constants.py index dbffdfa0..86a663e4 100644 --- a/src/exo/worker/engines/mlx/constants.py +++ b/src/exo/worker/engines/mlx/constants.py @@ -11,5 +11,7 @@ QUANTIZE_MODEL_MODE: str | None = "affine" CACHE_GROUP_SIZE: int = 64 KV_CACHE_BITS: int | None = None +DEFAULT_TOP_LOGPROBS: int = 5 + # TODO: We should really make this opt-in, but Kimi requires trust_remote_code=True TRUST_REMOTE_CODE: bool = True diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index a38d70c5..67a31ae0 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -12,6 +12,7 @@ from exo.shared.types.api import ( FinishReason, GenerationStats, PromptTokensDetails, + TopLogprobItem, Usage, ) from exo.shared.types.common import ModelId @@ -23,7 +24,12 @@ from exo.shared.types.worker.runner_response import ( ) from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.cache import KVPrefixCache, encode_prompt, make_kv_cache -from exo.worker.engines.mlx.constants import KV_BITS, KV_GROUP_SIZE, MAX_TOKENS +from exo.worker.engines.mlx.constants import ( + DEFAULT_TOP_LOGPROBS, + KV_BITS, + KV_GROUP_SIZE, + MAX_TOKENS, +) from exo.worker.engines.mlx.utils_mlx import ( apply_chat_template, mx_barrier, @@ -155,6 +161,60 @@ def eos_ids_from_tokenizer(tokenizer: TokenizerWrapper) -> list[int]: return eos +def extract_top_logprobs( + logprobs: mx.array, + tokenizer: TokenizerWrapper, + top_logprobs: int, + selected_token: int, +) -> tuple[float, list[TopLogprobItem]]: + """Extract the selected token's logprob and top alternative tokens. + + Args: + logprobs: Full vocabulary logprobs array from MLX + tokenizer: Tokenizer for decoding token IDs to strings + top_logprobs: Number of top alternatives to return + selected_token: The token ID that was actually sampled + + Returns: + Tuple of (selected_token_logprob, list of TopLogprobItem for top alternatives) + """ + # Get the logprob of the selected token + selected_logprob = float(logprobs[selected_token].item()) + + # Get top indices (most probable tokens) + # mx.argpartition gives indices that would partition the array + # We negate logprobs since argpartition finds smallest, and we want largest + top_logprobs = min(top_logprobs, logprobs.shape[0]) # Don't exceed vocab size + top_indices = mx.argpartition(-logprobs, top_logprobs)[:top_logprobs] + + # Get the actual logprob values for these indices + top_values = logprobs[top_indices] + + # Sort by logprob (descending) for consistent ordering + sort_order = mx.argsort(-top_values) + top_indices = top_indices[sort_order] + top_values = top_values[sort_order] + + # Convert to list of TopLogprobItem + top_logprob_items: list[TopLogprobItem] = [] + for i in range(top_logprobs): + token_id = int(top_indices[i].item()) + token_logprob = float(top_values[i].item()) + # Decode token ID to string + token_str = tokenizer.decode([token_id]) + # Get byte representation + token_bytes = list(token_str.encode("utf-8")) + top_logprob_items.append( + TopLogprobItem( + token=token_str, + logprob=token_logprob, + bytes=token_bytes, + ) + ) + + return selected_logprob, top_logprob_items + + def mlx_generate( model: Model, tokenizer: TokenizerWrapper, @@ -296,9 +356,22 @@ def mlx_generate( ), ) + # Extract logprobs from the full vocabulary logprobs array + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None + if task.logprobs: + logprob, top_logprobs = extract_top_logprobs( + logprobs=out.logprobs, + tokenizer=tokenizer, + top_logprobs=task.top_logprobs or DEFAULT_TOP_LOGPROBS, + selected_token=out.token, + ) + yield GenerationResponse( text=text, token=out.token, + logprob=logprob, + top_logprobs=top_logprobs, finish_reason=finish_reason, stats=stats, usage=usage, diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index e12aa185..4f1140fb 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -459,6 +459,12 @@ def apply_chat_template( continue formatted_messages.append({"role": msg.role, "content": msg.content}) + # For assistant prefilling, append content after templating to avoid a closing turn token. + partial_assistant_content: str | None = None + if formatted_messages and formatted_messages[-1].get("role") == "assistant": + partial_assistant_content = cast(str, formatted_messages[-1].get("content", "")) + formatted_messages = formatted_messages[:-1] + prompt: str = tokenizer.apply_chat_template( formatted_messages, tokenize=False, @@ -466,6 +472,9 @@ def apply_chat_template( tools=task_params.tools, ) + if partial_assistant_content: + prompt += partial_assistant_content + logger.info(prompt) return prompt diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 109ea219..b0e655cc 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -344,6 +344,8 @@ def main( usage=response.usage, finish_reason=response.finish_reason, stats=response.stats, + logprob=response.logprob, + top_logprobs=response.top_logprobs, ), ) ) From 3a9baeb9db469975f9c760960bf0d1fcb187d2ce Mon Sep 17 00:00:00 2001 From: Jake Hillion Date: Tue, 3 Feb 2026 22:49:47 +0000 Subject: [PATCH 9/9] EXO: add CLI flags for root install/uninstall The macOS app required user interaction via AppleScript prompts to install or uninstall network configuration components, making automated deployments difficult. Added --install and --uninstall command line flags that execute the network setup scripts directly when running as root, bypassing GUI prompts. Created a new main.swift entry point that parses CLI arguments and delegates to NetworkSetupHelper's new direct execution methods. This enables headless installation via `sudo EXO --install` for automated deployment scenarios while preserving the existing GUI behavior when launched normally. Test plan: - Deployed to a machine that didn't have the content installed. Got blocked on the popup and EXO never launched. - Relaunched EXO, confirmed it still never starts because of the popup. - Ran `sudo /Applications/EXO.app/Contents/MacOS/EXO --install` - Launched EXO - the API started as expected. - Ran `sudo /Applications/EXO.app/Contents/MacOS/EXO --uninstall` - Launched EXO - got the popup. --- app/EXO/EXO/EXOApp.swift | 1 - app/EXO/EXO/Services/NetworkSetupHelper.swift | 55 ++++++++++++ app/EXO/EXO/main.swift | 85 +++++++++++++++++++ 3 files changed, 140 insertions(+), 1 deletion(-) create mode 100644 app/EXO/EXO/main.swift diff --git a/app/EXO/EXO/EXOApp.swift b/app/EXO/EXO/EXOApp.swift index 7669c408..3aff58c3 100644 --- a/app/EXO/EXO/EXOApp.swift +++ b/app/EXO/EXO/EXOApp.swift @@ -14,7 +14,6 @@ import SwiftUI import UserNotifications import os.log -@main struct EXOApp: App { @StateObject private var controller: ExoProcessController @StateObject private var stateService: ClusterStateService diff --git a/app/EXO/EXO/Services/NetworkSetupHelper.swift b/app/EXO/EXO/Services/NetworkSetupHelper.swift index 82cb82d4..5428ee8e 100644 --- a/app/EXO/EXO/Services/NetworkSetupHelper.swift +++ b/app/EXO/EXO/Services/NetworkSetupHelper.swift @@ -288,6 +288,61 @@ enum NetworkSetupHelper { """ } + /// Direct install without GUI (requires root). + /// Returns true on success, false on failure. + static func installDirectly() -> Bool { + let script = makeInstallerScript() + return runShellDirectly(script) + } + + /// Direct uninstall without GUI (requires root). + /// Returns true on success, false on failure. + static func uninstallDirectly() -> Bool { + let script = makeUninstallScript() + return runShellDirectly(script) + } + + /// Run a shell script directly via Process (no AppleScript, requires root). + /// Returns true on success, false on failure. + private static func runShellDirectly(_ script: String) -> Bool { + let process = Process() + process.executableURL = URL(fileURLWithPath: "/bin/bash") + process.arguments = ["-c", script] + + let outputPipe = Pipe() + let errorPipe = Pipe() + process.standardOutput = outputPipe + process.standardError = errorPipe + + do { + try process.run() + process.waitUntilExit() + + let outputData = outputPipe.fileHandleForReading.readDataToEndOfFile() + let errorData = errorPipe.fileHandleForReading.readDataToEndOfFile() + + if let output = String(data: outputData, encoding: .utf8), !output.isEmpty { + print(output) + } + if let errorOutput = String(data: errorData, encoding: .utf8), !errorOutput.isEmpty { + fputs(errorOutput, stderr) + } + + if process.terminationStatus == 0 { + logger.info("Shell script completed successfully") + return true + } else { + logger.error("Shell script failed with exit code \(process.terminationStatus)") + return false + } + } catch { + logger.error( + "Failed to run shell script: \(error.localizedDescription, privacy: .public)") + fputs("Error: \(error.localizedDescription)\n", stderr) + return false + } + } + private static func runShellAsAdmin(_ script: String) throws { let escapedScript = script diff --git a/app/EXO/EXO/main.swift b/app/EXO/EXO/main.swift new file mode 100644 index 00000000..9383981f --- /dev/null +++ b/app/EXO/EXO/main.swift @@ -0,0 +1,85 @@ +// +// main.swift +// EXO +// +// Created by Jake Hillion on 2026-02-03. +// + +import Foundation + +/// Command line options for the EXO app +enum CLICommand { + case install + case uninstall + case help + case none +} + +/// Parse command line arguments to determine the CLI command +func parseArguments() -> CLICommand { + let args = CommandLine.arguments + if args.contains("--help") || args.contains("-h") { + return .help + } + if args.contains("--install") { + return .install + } + if args.contains("--uninstall") { + return .uninstall + } + return .none +} + +/// Print usage information +func printUsage() { + let programName = (CommandLine.arguments.first as NSString?)?.lastPathComponent ?? "EXO" + print( + """ + Usage: \(programName) [OPTIONS] + + Options: + --install Install EXO network configuration (requires root) + --uninstall Uninstall EXO network configuration (requires root) + --help, -h Show this help message + + When run without options, starts the normal GUI application. + + Examples: + sudo \(programName) --install Install network components as root + sudo \(programName) --uninstall Remove network components as root + """) +} + +/// Check if running as root +func isRunningAsRoot() -> Bool { + return getuid() == 0 +} + +// Main entry point +let command = parseArguments() + +switch command { +case .help: + printUsage() + exit(0) + +case .install: + if !isRunningAsRoot() { + fputs("Error: --install requires root privileges. Run with sudo.\n", stderr) + exit(1) + } + let success = NetworkSetupHelper.installDirectly() + exit(success ? 0 : 1) + +case .uninstall: + if !isRunningAsRoot() { + fputs("Error: --uninstall requires root privileges. Run with sudo.\n", stderr) + exit(1) + } + let success = NetworkSetupHelper.uninstallDirectly() + exit(success ? 0 : 1) + +case .none: + // Start normal GUI application + EXOApp.main() +}