Compare commits

..
Author SHA1 Message Date
Evan 7d949a08cf workin on it 2026-04-22 14:15:23 +01:00
Evan b7ed1f9dc6 implement engine interface for mlx and mflux 2026-04-22 12:51:07 +01:00
Evan 0bc29c6c9d remove worker responses from more places
runner responses are in a weird state, sharing a lot in common with
chunks. this begins consolidating them into a more coherent collection
of types
2026-04-22 12:51:01 +01:00
Evan 2229b99bec fix test 2026-04-22 12:51:01 +01:00
Evan 0a77426276 remove layer loading callback 2026-04-22 12:51:01 +01:00
37 changed files with 1203 additions and 1345 deletions

No files matched your search

-1
View File
@@ -30,7 +30,6 @@ dependencies = [
"zstandard>=0.23.0",
"mlx-vlm>=0.3.11",
"transformers>=5.0.0,<5.4.0",
"pydantic-settings>=2.13.1",
]
[project.scripts]
+55 -36
View File
@@ -119,7 +119,7 @@ from exo.api.types.openai_responses import (
)
from exo.master.image_store import ImageStore
from exo.master.placement import place_instance as get_instance_placements
from exo.shared.apply import apply
from exo.routing.state_manager import StateManager
from exo.shared.constants import (
DASHBOARD_DIR,
EXO_CACHE_HOME,
@@ -223,8 +223,9 @@ class API:
download_command_sender: Sender[ForwarderDownloadCommand],
# This lets us pause the API if an election is running
election_receiver: Receiver[ElectionMessage],
state_manager: StateManager[State],
) -> None:
self.state = State()
self.state_manager = state_manager
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self._system_id = SystemId()
self.command_sender = command_sender
@@ -271,11 +272,16 @@ class API:
self._image_store = ImageStore(EXO_IMAGE_CACHE_DIR)
self._tg: TaskGroup = TaskGroup()
def reset(self, result_clock: int, event_receiver: Receiver[IndexedEvent]):
def reset(
self,
result_clock: int,
event_receiver: Receiver[IndexedEvent],
state_manager: StateManager[State],
):
logger.info("Resetting API State")
self._event_log.close()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self.state = State()
self.state_manager = state_manager
self._system_id = SystemId()
self._text_generation_queues = {}
self._image_generation_queues = {}
@@ -372,11 +378,12 @@ class API:
self.app.get("/onboarding")(self.get_onboarding)
self.app.post("/onboarding")(self.complete_onboarding)
def get_state(self, path: str = ""):
def get_state(self, path: str = "") -> Any: # pyright: ignore[reportAny]
state = self.state_manager.get()
if path == "":
return self.state
return state
try:
x = self.state.model_dump(by_alias=True)
x = state.model_dump(by_alias=True)
for attr in path.split("/"):
if attr != "":
if isinstance(x, dict):
@@ -438,6 +445,7 @@ class API:
min_nodes: int = 1,
) -> Instance:
model_card = await ModelCard.load(model_id)
state = self.state_manager.get()
try:
placements = get_instance_placements(
@@ -447,16 +455,16 @@ class API:
instance_meta=instance_meta,
min_nodes=min_nodes,
),
node_memory=self.state.node_memory,
node_network=self.state.node_network,
topology=self.state.topology,
current_instances=self.state.instances,
download_status=self.state.downloads,
node_memory=state.node_memory,
node_network=state.node_network,
topology=state.topology,
current_instances=state.instances,
download_status=state.downloads,
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
current_ids = set(self.state.instances.keys())
current_ids = set(state.instances.keys())
new_ids = [
instance_id for instance_id in placements if instance_id not in current_ids
]
@@ -476,8 +484,9 @@ class API:
seen: set[tuple[ModelId, Sharding, InstanceMeta, int]] = set()
previews: list[PlacementPreview] = []
required_nodes = set(node_ids) if node_ids else None
state = self.state_manager.get()
if len(list(self.state.topology.list_nodes())) == 0:
if len(list(state.topology.list_nodes())) == 0:
return PlacementPreviewResponse(previews=[])
try:
@@ -492,9 +501,7 @@ class API:
instance_combinations.extend(
[
(sharding, instance_meta, i)
for i in range(
1, len(list(self.state.topology.list_nodes())) + 1
)
for i in range(1, len(list(state.topology.list_nodes())) + 1)
]
)
# TODO: PDD
@@ -509,12 +516,12 @@ class API:
instance_meta=instance_meta,
min_nodes=min_nodes,
),
node_memory=self.state.node_memory,
node_network=self.state.node_network,
topology=self.state.topology,
current_instances=self.state.instances,
node_memory=state.node_memory,
node_network=state.node_network,
topology=state.topology,
current_instances=state.instances,
required_nodes=required_nodes,
download_status=self.state.downloads,
download_status=state.downloads,
)
except ValueError as exc:
if (model_card.model_id, sharding, instance_meta, 0) not in seen:
@@ -530,7 +537,7 @@ class API:
seen.add((model_card.model_id, sharding, instance_meta, 0))
continue
current_ids = set(self.state.instances.keys())
current_ids = set(state.instances.keys())
new_instances = [
instance
for instance_id, instance in placements.items()
@@ -592,12 +599,14 @@ class API:
return PlacementPreviewResponse(previews=previews)
def get_instance(self, instance_id: InstanceId) -> Instance:
if instance_id not in self.state.instances:
state = self.state_manager.get()
if instance_id not in state.instances:
raise HTTPException(status_code=404, detail="Instance not found")
return self.state.instances[instance_id]
return state.instances[instance_id]
async def delete_instance(self, instance_id: InstanceId) -> DeleteInstanceResponse:
if instance_id not in self.state.instances:
state = self.state_manager.get()
if instance_id not in state.instances:
raise HTTPException(status_code=404, detail="Instance not found")
command = DeleteInstance(
@@ -666,7 +675,9 @@ class API:
async def _collect_text_generation_with_stats(
self, command_id: CommandId
) -> BenchChatCompletionResponse:
sampler = PowerSampler(get_node_system=lambda: self.state.node_system)
sampler = PowerSampler(
get_node_system=lambda: self.state_manager.get().node_system
)
text_parts: list[str] = []
tool_calls: list[ToolCall] = []
model: ModelId | None = None
@@ -864,9 +875,10 @@ class API:
Raises HTTPException 404 if no instance is found for the model.
"""
state = self.state_manager.get()
if not any(
instance.shard_assignments.model_id == model_id
for instance in self.state.instances.values()
for instance in state.instances.values()
):
await self._trigger_notify_user_to_download_model(model_id)
raise HTTPException(
@@ -882,9 +894,10 @@ class API:
"""
model_card = await ModelCard.load(model)
resolved_model = model_card.model_id
state = self.state_manager.get()
if not any(
instance.shard_assignments.model_id == resolved_model
for instance in self.state.instances.values()
for instance in state.instances.values()
):
await self._trigger_notify_user_to_download_model(resolved_model)
raise HTTPException(
@@ -1194,7 +1207,9 @@ class API:
num_images: int,
response_format: str,
) -> BenchImageGenerationResponse:
sampler = PowerSampler(get_node_system=lambda: self.state.node_system)
sampler = PowerSampler(
get_node_system=lambda: self.state_manager.get().node_system
)
images: list[ImageData] = []
stats: ImageGenerationStats | None = None
async with anyio.create_task_group() as tg:
@@ -1564,8 +1579,9 @@ class API:
def none_if_empty(value: str) -> str | None:
return value or None
state = self.state_manager.get()
downloaded_model_ids: set[str] = set()
for node_downloads in self.state.downloads.values():
for node_downloads in state.downloads.values():
for dl in node_downloads:
if isinstance(dl, DownloadCompleted):
downloaded_model_ids.add(dl.shard_metadata.model_card.model_id)
@@ -1619,7 +1635,8 @@ class API:
"""Returns list of running models (active instances)."""
models: list[OllamaPsModel] = []
seen: set[str] = set()
for instance in self.state.instances.values():
state = self.state_manager.get()
for instance in state.instances.values():
model_id = str(instance.shard_assignments.model_id)
if model_id in seen:
continue
@@ -1641,7 +1658,8 @@ class API:
"""Calculate total available memory across all nodes in bytes."""
total_available = Memory()
for memory in self.state.node_memory.values():
state = self.state_manager.get()
for memory in state.node_memory.values():
total_available += memory.ram_available
return total_available
@@ -1649,10 +1667,11 @@ class API:
async def get_models(self, status: str | None = Query(default=None)) -> ModelList:
"""Returns list of available models, optionally filtered by being downloaded."""
cards = await get_model_cards()
state = self.state_manager.get()
if status == "downloaded":
downloaded_model_ids: set[str] = set()
for node_downloads in self.state.downloads.values():
for node_downloads in state.downloads.values():
for dl in node_downloads:
if isinstance(dl, DownloadCompleted):
downloaded_model_ids.add(dl.shard_metadata.model_card.model_id)
@@ -1808,7 +1827,6 @@ class API:
with self.event_receiver as events:
async for i_event in events:
self._event_log.append(i_event.event)
self.state = apply(self.state, i_event)
event = i_event.event
if isinstance(event, ChunkGenerated):
@@ -1835,7 +1853,8 @@ class API:
def _close_streams_for_instance(self, instance_id: InstanceId) -> None:
"""Close any active generation streams for commands running on the given instance."""
for task in self.state.tasks.values():
state = self.state_manager.get()
for task in state.tasks.values():
if task.instance_id != instance_id:
continue
if not isinstance(
+19 -1
View File
@@ -17,6 +17,7 @@ from exo.download.impl_shard_downloader import exo_shard_downloader
from exo.master.main import Master
from exo.routing.event_router import EventRouter
from exo.routing.router import Router, get_node_id_keypair
from exo.routing.state_manager import state_manager_from_routers
from exo.shared.constants import EXO_LOG
from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
@@ -64,6 +65,7 @@ class Node:
command_sender=router.sender(topics.COMMANDS),
external_outbound=router.sender(topics.LOCAL_EVENTS),
external_inbound=router.receiver(topics.GLOBAL_EVENTS),
state_receiver=router.receiver(topics.STATE_SNAPSHOTS),
)
logger.info(f"Starting node {node_id}")
@@ -88,6 +90,7 @@ class Node:
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
election_receiver=router.receiver(topics.ELECTION_MESSAGES),
state_manager=state_manager_from_routers(router, event_router),
)
else:
api = None
@@ -100,6 +103,7 @@ class Node:
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
api_port=args.api_port,
state_manager=state_manager_from_routers(router, event_router),
)
else:
worker = None
@@ -113,6 +117,7 @@ class Node:
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
command_receiver=router.receiver(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
state_manager=state_manager_from_routers(router, event_router),
)
er_send, er_recv = channel[ElectionResult]()
@@ -191,6 +196,7 @@ class Node:
self.router.sender(topics.COMMANDS),
self.router.receiver(topics.GLOBAL_EVENTS),
self.router.sender(topics.LOCAL_EVENTS),
self.router.receiver(topics.STATE_SNAPSHOTS),
)
if (
@@ -213,6 +219,9 @@ class Node:
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
state_manager=state_manager_from_routers(
self.router, self.event_router
),
)
self._tg.start_soon(self.master.run)
elif (
@@ -253,10 +262,19 @@ class Node:
topics.DOWNLOAD_COMMANDS
),
api_port=self._api_port,
state_manager=state_manager_from_routers(
self.router, self.event_router
),
)
self._tg.start_soon(self.worker.run)
if self.api:
self.api.reset(result.won_clock, self.event_router.receiver())
self.api.reset(
result.won_clock,
self.event_router.receiver(),
state_manager=state_manager_from_routers(
self.router, self.event_router
),
)
self._tg.start_soon(self.event_router.run)
else:
if self.api:
+28 -26
View File
@@ -10,7 +10,7 @@ from exo.master.placement import (
get_transition_events,
place_instance,
)
from exo.shared.apply import apply
from exo.routing.state_manager import StateManager
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
from exo.shared.types.commands import (
AddCustomModelCard,
@@ -80,10 +80,11 @@ class Master:
local_event_receiver: Receiver[LocalForwarderEvent],
global_event_sender: Sender[GlobalForwarderEvent],
download_command_sender: Sender[ForwarderDownloadCommand],
state_manager: StateManager[State],
):
self.node_id = node_id
self.session_id = session_id
self.state = State()
self.state_manager = state_manager
self._tg: TaskGroup = TaskGroup()
self.command_task_mapping: dict[CommandId, TaskId] = {}
self.command_receiver = command_receiver
@@ -118,6 +119,7 @@ class Master:
async def _command_processor(self) -> None:
with self.command_receiver as commands:
async for forwarder_command in commands:
state = self.state_manager.get()
try:
logger.info(f"Executing command: {forwarder_command.command}")
@@ -128,14 +130,14 @@ class Master:
case TestCommand():
pass
case TextGeneration():
for instance in self.state.instances.values():
for instance in state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
):
task_count = sum(
1
for task in self.state.tasks.values()
for task in state.tasks.values()
if task.instance_id == instance.instance_id
)
instance_task_counts[instance.instance_id] = (
@@ -170,14 +172,14 @@ class Master:
self.command_task_mapping[command.command_id] = task_id
case ImageGeneration():
for instance in self.state.instances.values():
for instance in state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
):
task_count = sum(
1
for task in self.state.tasks.values()
for task in state.tasks.values()
if task.instance_id == instance.instance_id
)
instance_task_counts[instance.instance_id] = (
@@ -214,7 +216,7 @@ class Master:
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance = state.instances.get(
selected_instance_id
)
if selected_instance:
@@ -224,14 +226,14 @@ class Master:
)
self._expected_ranks[task_id] = ranks
case ImageEdits():
for instance in self.state.instances.values():
for instance in state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
):
task_count = sum(
1
for task in self.state.tasks.values()
for task in state.tasks.values()
if task.instance_id == instance.instance_id
)
instance_task_counts[instance.instance_id] = (
@@ -268,7 +270,7 @@ class Master:
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance = state.instances.get(
selected_instance_id
)
if selected_instance:
@@ -278,12 +280,12 @@ class Master:
)
self._expected_ranks[task_id] = ranks
case DeleteInstance():
placement = delete_instance(command, self.state.instances)
placement = delete_instance(command, state.instances)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
state.instances, placement, state.tasks
)
for cmd in cancel_unnecessary_downloads(
placement, self.state.downloads
placement, state.downloads
):
await self.download_command_sender.send(
ForwarderDownloadCommand(
@@ -294,24 +296,24 @@ class Master:
case PlaceInstance():
placement = place_instance(
command,
self.state.topology,
self.state.instances,
self.state.node_memory,
self.state.node_network,
download_status=self.state.downloads,
state.topology,
state.instances,
state.node_memory,
state.node_network,
download_status=state.downloads,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
state.instances, placement, state.tasks
)
generated_events.extend(transition_events)
case CreateInstance():
placement = add_instance_to_placements(
command,
self.state.topology,
self.state.instances,
state.topology,
state.instances,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
state.instances, placement, state.tasks
)
generated_events.extend(transition_events)
case SendInputChunk(chunk=chunk):
@@ -374,9 +376,10 @@ class Master:
# These plan loops are the cracks showing in our event sourcing architecture - more things could be commands
async def _plan(self) -> None:
while True:
state = self.state_manager.get()
# kill broken instances
connected_node_ids = set(self.state.topology.list_nodes())
for instance_id, instance in self.state.instances.items():
connected_node_ids = set(state.topology.list_nodes())
for instance_id, instance in state.instances.items():
for node_id in instance.shard_assignments.node_to_runner:
if node_id not in connected_node_ids:
await self.event_sender.send(
@@ -385,7 +388,7 @@ class Master:
break
# time out dead nodes
for node_id, time in self.state.last_seen.items():
for node_id, time in state.last_seen.items():
now = datetime.now(tz=timezone.utc)
if now - time > timedelta(seconds=30):
logger.info(f"Manually removing node {node_id} due to inactivity")
@@ -411,7 +414,6 @@ class Master:
logger.debug(f"Master indexing event: {str(event)[:100]}")
indexed = IndexedEvent(event=event, idx=len(self._event_log))
self.state = apply(self.state, indexed)
event = event.model_copy(
update={"_master_time_stamp": datetime.now(tz=timezone.utc)}
+2
View File
@@ -15,6 +15,7 @@ from exo.shared.types.events import (
IndexedEvent,
LocalForwarderEvent,
)
from exo.shared.types.state import ForwarderState
from exo.utils.channels import Receiver, Sender, channel
from exo.utils.event_buffer import OrderedBuffer
from exo.utils.task_group import TaskGroup
@@ -26,6 +27,7 @@ class EventRouter:
command_sender: Sender[ForwarderCommand]
external_inbound: Receiver[GlobalForwarderEvent]
external_outbound: Sender[LocalForwarderEvent]
state_receiver: Receiver[ForwarderState]
_system_id: SystemId = field(init=False, default_factory=SystemId)
internal_outbound: list[Sender[IndexedEvent]] = field(
init=False, default_factory=list
+67
View File
@@ -0,0 +1,67 @@
from collections.abc import Callable, Iterable
from dataclasses import dataclass
import exo.routing.topics as topics
from exo.shared.apply import apply
from exo.shared.types.common import SessionId
from exo.shared.types.events import (
IndexedEvent,
)
from exo.shared.types.state import BaseState, ForwarderState, State
from exo.utils import fmap, fold
from exo.utils.channels import Receiver
from .event_router import EventRouter
from .router import Router
@dataclass
class StateManager[T: BaseState]:
_apply: Callable[[T, IndexedEvent], T]
_event_recv: Receiver[IndexedEvent]
_state_recv: Receiver[T]
_state: T
def get(self) -> T:
return self._state
async def run(self):
async with self._event_recv, self._state_recv:
async for event in self._event_recv:
# apply new states eagerly
def order_state(current: T, other: T) -> T:
return (
current
if other.last_event_idx() < current.last_event_idx()
else other
)
self._state = fold(self._state, order_state, self._state_recv.collect())
# catch up / ignore stale
if event.idx <= self._state.last_event_idx():
continue
# apply state
self._state = self._apply(self._state, event)
@dataclass
class _Hack:
recv: Receiver[ForwarderState]
session_id: SessionId
def collect(self) -> Iterable[State]:
return fmap(self._matches, self.recv.collect())
def _matches(self, s: ForwarderState) -> State | None:
return s.state if s.session_id == self.session_id else None
def state_manager_from_routers(
router: Router, event_router: EventRouter
) -> StateManager[State]:
return StateManager(
apply,
event_router.receiver(),
_Hack(router.receiver(topics.STATE_SNAPSHOTS), event_router.session_id), # type: ignore
State(),
)
+2
View File
@@ -8,6 +8,7 @@ from exo.shared.types.events import (
GlobalForwarderEvent,
LocalForwarderEvent,
)
from exo.shared.types.state import ForwarderState
from exo.utils.pydantic_ext import FrozenModel
@@ -49,3 +50,4 @@ CONNECTION_MESSAGES = TypedTopic(
DOWNLOAD_COMMANDS = TypedTopic(
"download_commands", PublishPolicy.Always, ForwarderDownloadCommand
)
STATE_SNAPSHOTS = TypedTopic("state_snapshots", PublishPolicy.Always, ForwarderState)
+1
View File
@@ -68,6 +68,7 @@ DASHBOARD_DIR = (
# Log files (data/logs or cache)
EXO_LOG_DIR = EXO_CACHE_HOME / "exo_log"
EXO_LOG = EXO_LOG_DIR / "exo.log"
EXO_TEST_LOG = EXO_CACHE_HOME / "exo_test.log"
# Identity (config)
EXO_NODE_ID_KEYPAIR = EXO_CONFIG_HOME / "node_id.keypair"
-158
View File
@@ -1,158 +0,0 @@
from pathlib import Path
from collections.abc import Sequence
import tomlkit
from exo.utils.pydantic_ext import FrozenModel
from typing import Self, Any
from pydantic import Field, BaseModel, model_validator
from pydantic_settings import BaseSettings, SettingsConfigDict, PydanticBaseSettingsSource, TomlConfigSettingsSource
from exo.shared.types.common import NodeId, ModelId
from exo.shared.types.worker.instances import InstanceId
from exo.shared.constants import EXO_CONFIG_HOME, EXO_DATA_HOME, EXO_CACHE_HOME
from exo.utils.dashboard_path import find_dashboard, find_resources
def default_merge[T: BaseModel](left: T, right: T) -> T:
if left == right:
return left
merged_dict = {}
for key in type(left).model_fields:
try:
merged_dict[key] = getattr(left, key).merge( # pyright: ignore[reportAny]
getattr(right, key, None)
)
except AttributeError:
raise NotImplementedError("Cluster Option using default implementation incorrectly")
return type(left).model_validate(merged_dict)
def _parse_colon_separated_dirs(obj: Any) -> set[Path]: # pyright: ignore[reportAny]
if isinstance(obj, (list, set)):
return set(Path(d).expanduser() for d in obj) # pyright: ignore[reportUnknownArgumentType, reportUnknownVariableType]
else:
return set(Path(d).expanduser() for d in str(obj).split(":")) # pyright: ignore[reportAny]
class ModelDirsSettings(BaseModel, frozen=True):
# env: EXO_MODEL_DIRS_DEFAULT prepends to WRITEABLE, defaults to EXO_DATA_HOME/models
# env: EXO_MODEL_DIRS_WRITEABLE, defaults to []
writeable: list[Path] = []
# env: EXO_MODEL_DIRS_READONLY, defaults to []
readonly: list[Path] = []
@model_validator(mode="before")
@classmethod
def build_defaults(cls, data: Any) -> Any: # pyright: ignore[reportAny]
if not isinstance(data, dict):
return data # pyright: ignore[reportAny]
default = Path(data.get("default", EXO_DATA_HOME / "models")).expanduser() # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType]
readonly = _parse_colon_separated_dirs(data.get("readonly", [])) # pyright: ignore[reportUnknownMemberType]
writeable = _parse_colon_separated_dirs(data.get("writeable", [])).difference(readonly) # pyright: ignore[reportUnknownMemberType]
if default not in readonly:
writeable = [default, *writeable]
return {**data, "writeable": writeable, "readonly": readonly} # pyright: ignore[reportUnknownVariableType]
class RuntimeDirsSettings(BaseModel, frozen=True):
dashboard: Path = Field(default_factory=find_dashboard)
resources: Path = Field(default_factory=find_resources)
logs: Path = EXO_CACHE_HOME / "log"
log_file: str = "latest.log"
def log_file_path(self):
return self.logs / self.log_file
# doesnt require merge
class LocalSettings(FrozenModel):
runtime_dirs: RuntimeDirsSettings
model_dirs: ModelDirsSettings
class InstanceSettings(FrozenModel):
# env: EXO_INSTANCE_DEFAULTS_BATCH_CONCURRENCY
batch_concurrency: int
def merge(self, other: Self) -> Self:
return type(self)(batch_concurrency=min(self.batch_concurrency, other.batch_concurrency))
class ClusterSettings(FrozenModel):
instance_defaults: InstanceSettings = InstanceSettings(batch_concurrency=8)
model_settings_overrides: dict[ModelId, InstanceSettings] = {}
def merge(self, other: Self) -> Self:
return default_merge(self, other)
class SettingsFile(BaseSettings):
model_config = SettingsConfigDict(
extra='ignore',
frozen=True,
toml_file=EXO_CONFIG_HOME / "config.toml",
env_prefix="EXO_",
env_nested_delimiter="_",
env_ignore_empty=True,
)
model_dirs: ModelDirsSettings
runtime_dirs: RuntimeDirsSettings
model_settings_overrides: dict[ModelId, InstanceSettings] = {}
instance_defaults: InstanceSettings
def get_local(self) -> LocalSettings:
...
def get_cluster(self) -> ClusterSettings:
...
@classmethod
def settings_customise_sources(
cls,
settings_cls: type[BaseSettings],
init_settings: PydanticBaseSettingsSource,
env_settings: PydanticBaseSettingsSource,
dotenv_settings: PydanticBaseSettingsSource,
file_secret_settings: PydanticBaseSettingsSource,
) -> tuple[PydanticBaseSettingsSource, ...]:
return (init_settings, env_settings, TomlConfigSettingsSource(settings_cls),)
def sync(self):
"""nb: only call this once per save"""
cfg_path = type(self).model_config.get("toml_file", None)
if isinstance(cfg_path, Sequence):
cfg_path=cfg_path[0]
if cfg_path:
with open(cfg_path, "w") as fp:
tomlkit.dump(self.model_dump(exclude_defaults=True), fp) # pyright: ignore[reportUnknownMemberType]
class StateSettings(FrozenModel):
per_node: dict[NodeId, LocalSettings]
per_instance: dict[InstanceId, InstanceSettings]
cluster: ClusterSettings
def model_merge_local(self, node_id: NodeId, settings: LocalSettings) -> Self:
return self.model_copy(update={
"per_node": {
**self.per_node,
node_id: settings
}
})
def model_merge_cluster(self, settings: ClusterSettings) -> Self:
return self.model_copy(update={
"cluster": self.cluster.merge(settings)
})
def settings_for(self, node_id: NodeId) -> StoredSettings:
merged = {}
for key, val in self.cluster.model_dump(exclude_defaults=True).items(): # pyright: ignore[reportAny]
merged[key] = val
if (local := self.per_node.get(node_id, None)) is not None:
for key, val in local.model_dump(exclude_defaults=True).items(): # pyright: ignore[reportAny]
merged[key] = val
return StoredSettings.model_validate(merged)
def sync(self, node_id: NodeId):
"""nb: only call this once per save"""
toml_file=EXO_CONFIG_HOME / "config.toml"
with open(toml_file, "w") as fp:
tomlkit.dump(self.settings_for(node_id).model_dump(exclude_defaults=True), fp) # pyright: ignore[reportUnknownMemberType]
+3 -3
View File
@@ -85,6 +85,6 @@ class PrefillProgressChunk(BaseChunk):
total_tokens: int
GenerationChunk = (
TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk | PrefillProgressChunk
)
StatusChunk = PrefillProgressChunk
GenerationChunk = TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk
Chunk = StatusChunk | GenerationChunk
+2 -2
View File
@@ -5,7 +5,7 @@ from pydantic import Field
from exo.shared.models.model_cards import ModelCard
from exo.shared.topology import Connection
from exo.shared.types.chunks import GenerationChunk, InputImageChunk
from exo.shared.types.chunks import Chunk, InputImageChunk
from exo.shared.types.common import CommandId, Id, ModelId, NodeId, SessionId, SystemId
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.worker.downloads import DownloadProgress
@@ -91,7 +91,7 @@ class NodeDownloadProgress(BaseEvent):
class ChunkGenerated(BaseEvent):
command_id: CommandId
chunk: GenerationChunk
chunk: Chunk
class InputChunkReceived(BaseEvent):
+18 -14
View File
@@ -1,12 +1,11 @@
from abc import ABC, abstractmethod
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, cast
from pydantic import ConfigDict, Field, field_serializer, field_validator
from pydantic.alias_generators import to_camel
from pydantic import Field, field_serializer, field_validator
from exo.shared.topology import Topology, TopologySnapshot
from exo.shared.types.common import NodeId
from exo.shared.types.common import NodeId, SessionId
from exo.shared.types.profiling import (
DiskUsage,
MemoryUsage,
@@ -24,7 +23,12 @@ from exo.shared.types.worker.runners import RunnerId, RunnerStatus
from exo.utils.pydantic_ext import FrozenModel
class State(FrozenModel):
class BaseState(ABC):
@abstractmethod
def last_event_idx(self) -> int: ...
class State(BaseState, FrozenModel, arbitrary_types_allowed=True):
"""Global system state.
The :class:`Topology` instance is encoded/decoded via an immutable
@@ -32,14 +36,6 @@ class State(FrozenModel):
standard JSON serialisation.
"""
model_config = ConfigDict(
alias_generator=to_camel,
validate_by_name=True,
extra="forbid",
# I want to reenable this ASAP, but it's causing an issue with TaskStatus
strict=True,
arbitrary_types_allowed=True,
)
instances: Mapping[InstanceId, Instance] = {}
runners: Mapping[RunnerId, RunnerStatus] = {}
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
@@ -61,6 +57,9 @@ class State(FrozenModel):
# Detected cycles where all nodes have Thunderbolt bridge enabled (>2 nodes)
thunderbolt_bridge_cycles: Sequence[Sequence[NodeId]] = []
def last_event_idx(self) -> int:
return self.last_event_applied_idx
@field_serializer("topology", mode="plain")
def _encode_topology(self, value: Topology) -> TopologySnapshot:
return value.to_snapshot()
@@ -78,7 +77,12 @@ class State(FrozenModel):
return value
if isinstance(value, Mapping): # likely a snapshot-dict coming from JSON
snapshot = TopologySnapshot(**cast(dict[str, Any], value)) # type: ignore[arg-type]
snapshot = TopologySnapshot.model_validate(value)
return Topology.from_snapshot(snapshot)
raise TypeError("Invalid representation for Topology field in State")
class ForwarderState(FrozenModel):
state: State
session_id: SessionId
+3
View File
@@ -101,3 +101,6 @@ Task = (
| ImageEdits
| Shutdown
)
TextTask = TextGeneration
ImageTask = ImageGeneration | ImageEdits
GenerationTask = TextTask | ImageTask
Whitespace-only changes.
@@ -16,10 +16,6 @@ class BaseRunnerResponse(TaggedModel):
pass
class TokenizedResponse(BaseRunnerResponse):
prompt_tokens: int
class GenerationResponse(BaseRunnerResponse):
text: str
token: int
@@ -70,6 +66,15 @@ class FinishedResponse(BaseRunnerResponse):
pass
class ModelLoadingResponse(BaseRunnerResponse):
layers_loaded: int
total: int
class CancelledResponse(BaseRunnerResponse):
pass
class PrefillProgressResponse(BaseRunnerResponse):
processed_tokens: int
total_tokens: int
+16 -1
View File
@@ -1,4 +1,5 @@
from typing import Any, Type
from collections.abc import Callable, Iterable
from typing import Any, Type, TypeGuard
from .phantom import PhantomData
@@ -14,3 +15,17 @@ def todo[T](
_phantom: PhantomData[T] = None,
) -> T:
raise NotImplementedError(msg)
def fold[T, U](acc: T, fn: Callable[[T, U], T], iterator: Iterable[U]) -> T:
for it in iterator:
acc = fn(acc, it)
return acc
def _filter_none[U](item: U | None) -> TypeGuard[U]:
return item is not None
def fmap[T, U](fn: Callable[[T], U | None], iterator: Iterable[T]) -> Iterable[U]:
return filter(_filter_none, map(fn, iterator))
-3
View File
@@ -1,10 +1,8 @@
import sys
from pathlib import Path
from typing import cast
from functools import cache
@cache
def find_resources() -> Path:
resources = _find_resources_in_repo() or _find_resources_in_bundle()
if resources is None:
@@ -33,7 +31,6 @@ def _find_resources_in_bundle() -> Path | None:
return None
@cache
def find_dashboard() -> Path:
dashboard = _find_dashboard_in_repo() or _find_dashboard_in_bundle()
if not dashboard:
+55
View File
@@ -0,0 +1,55 @@
from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from exo.shared.types.chunks import Chunk
from exo.shared.types.tasks import CANCEL_ALL_TASKS, GenerationTask, TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
ModelLoadingResponse,
)
class Engine(ABC):
_cancelled_tasks: set[TaskId]
def should_cancel(self, task_id: TaskId) -> bool:
return (
task_id in self._cancelled_tasks
or CANCEL_ALL_TASKS in self._cancelled_tasks
)
@abstractmethod
def warmup(self) -> None: ...
@abstractmethod
def submit(
self,
task: GenerationTask,
) -> None: ...
@abstractmethod
def step(
self,
) -> Iterable[tuple[TaskId, Chunk | CancelledResponse | FinishedResponse]]: ...
@abstractmethod
def close(self) -> None: ...
class Builder(ABC):
@abstractmethod
def connect(self, bound_instance: BoundInstance) -> None: ...
@abstractmethod
def load(
self,
bound_instance: BoundInstance,
) -> Generator[ModelLoadingResponse]: ...
@abstractmethod
def build(self) -> Engine: ...
@abstractmethod
def close(self) -> None: ...
+6 -2
View File
@@ -1,12 +1,16 @@
from exo.worker.engines.image.builder import (
ImageEngine,
MfluxBuilder,
)
from exo.worker.engines.image.distributed_model import (
DistributedImageModel,
initialize_image_model,
)
from exo.worker.engines.image.generate import generate_image, warmup_image_generator
__all__ = [
"MfluxBuilder",
"ImageEngine",
"DistributedImageModel",
"generate_image",
"initialize_image_model",
"warmup_image_generator",
]
+212
View File
@@ -0,0 +1,212 @@
import contextlib
from collections import deque
from collections.abc import Generator, Iterable
from dataclasses import dataclass, field
import mlx.core as mx
from loguru import logger
from exo.api.types import ImageEditsTaskParams, ImageGenerationTaskParams
from exo.shared.constants import EXO_TRACING_ENABLED
from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
from exo.shared.types.chunks import Chunk, ErrorChunk
from exo.shared.types.events import (
Event,
TraceEventData,
TracesCollected,
)
from exo.shared.types.tasks import (
GenerationTask,
ImageEdits,
ImageGeneration,
ImageTask,
TaskId,
)
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
ModelLoadingResponse,
)
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
ShardMetadata,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.base import Builder, Engine
from exo.worker.engines.image.distributed_model import (
DistributedImageModel,
)
from exo.worker.engines.image.generate import (
generate_image,
warmup_image_generator,
)
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
)
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 _send_traces_if_enabled(
event_sender: MpSender[Event],
task_id: TaskId,
rank: int,
) -> None:
if not EXO_TRACING_ENABLED:
return
traces = get_trace_buffer()
if traces:
trace_data = [
TraceEventData(
name=t.name,
start_us=t.start_us,
duration_us=t.duration_us,
rank=t.rank,
category=t.category,
)
for t in traces
]
event_sender.send(
TracesCollected(
task_id=task_id,
rank=rank,
traces=trace_data,
)
)
clear_trace_buffer()
@dataclass
class MfluxBuilder(Builder):
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
shard_metadata: ShardMetadata | None = None
image_model: DistributedImageModel | None = None
group: mx.distributed.Group | None = None
def connect(self, bound_instance: BoundInstance) -> None:
self.group = initialize_mlx(bound_instance)
def load(self, bound_instance: BoundInstance) -> Generator[ModelLoadingResponse]:
self.shard_metadata = bound_instance.bound_shard
self.image_model = DistributedImageModel.from_shard_metadata(
bound_instance.bound_shard, self.group
)
return
# very important!
yield
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.image_model, self.group
def build(
self,
) -> Engine:
assert self.image_model
assert self.shard_metadata
return ImageEngine(
self.image_model,
self.shard_metadata,
self.event_sender,
self.cancel_receiver,
)
@dataclass
class ImageEngine(Engine):
image_model: DistributedImageModel
shard_metadata: ShardMetadata
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
current_gen: Generator[tuple[TaskId, Chunk]] | None = field(
init=False, default=None
)
queue: deque[ImageTask] = field(init=False, default_factory=deque)
def warmup(self) -> None:
image = warmup_image_generator(model=self.image_model)
if image is not None:
logger.info(f"warmed up by generating {image.size} image")
else:
logger.info("warmup completed (non-primary node)")
def submit(
self,
task: GenerationTask,
) -> None:
assert isinstance(task, (ImageGeneration, ImageEdits))
self.queue.append(task)
def step(
self,
) -> Iterable[tuple[TaskId, Chunk | CancelledResponse | FinishedResponse]]:
resp = None
if self.current_gen is not None:
resp = next(self.current_gen, None)
if resp is None and len(self.queue) > 0:
task = self.queue.popleft()
self.current_gen = self._run_image_task(task.task_id, task.task_params)
resp = next(self.current_gen, None)
return (resp,) if resp is not None else ()
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.image_model
def _run_image_task(
self,
task_id: TaskId,
task_params: ImageGenerationTaskParams | ImageEditsTaskParams,
) -> Generator[tuple[TaskId, Chunk]]:
assert self.image_model
logger.info(f"received image task: {str(task_params)[:500]}")
def cancel_checker() -> bool:
for cancel_id in self.cancel_receiver.collect():
self._cancelled_tasks.add(cancel_id)
return self.should_cancel(task_id)
try:
for response in generate_image(
model=self.image_model,
task=task_params,
cancel_checker=cancel_checker,
):
if _is_primary_output_node(self.shard_metadata):
yield (task_id, response)
except Exception as e:
if _is_primary_output_node(self.shard_metadata):
yield (
task_id,
ErrorChunk(
model=self.shard_metadata.model_card.model_id,
finish_reason="error",
error_message=str(e),
),
)
raise
finally:
_send_traces_if_enabled(
self.event_sender, task_id, self.shard_metadata.device_rank
)
return
@@ -1,6 +1,6 @@
from collections.abc import Callable, Generator
from pathlib import Path
from typing import Any, Literal, Optional
from typing import Any, Literal
import mlx.core as mx
from mflux.models.common.config.config import Config
@@ -8,8 +8,12 @@ from PIL import Image
from exo.api.types import AdvancedImageParams
from exo.download.download_utils import build_model_path
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata
from exo.shared.types.common import ModelId
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
ShardMetadata,
)
from exo.worker.engines.image.config import ImageModelConfig
from exo.worker.engines.image.models import (
create_adapter_for_model,
@@ -17,21 +21,22 @@ from exo.worker.engines.image.models import (
)
from exo.worker.engines.image.models.base import ModelAdapter
from exo.worker.engines.image.pipeline import DiffusionRunner
from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init, mx_barrier
from exo.worker.engines.mlx.utils_mlx import mx_barrier
from exo.worker.runner.bootstrap import logger
class DistributedImageModel:
model_id: ModelId
_config: ImageModelConfig
_adapter: ModelAdapter[Any, Any]
_runner: DiffusionRunner
def __init__(
self,
model_id: str,
model_id: ModelId,
local_path: Path,
shard_metadata: PipelineShardMetadata | CfgShardMetadata,
group: Optional[mx.distributed.Group] = None,
group: mx.distributed.Group | None,
quantize: int | None = None,
):
config = get_config_for_model(model_id)
@@ -68,37 +73,27 @@ class DistributedImageModel:
else:
logger.info("Single-node initialization")
self.model_id = model_id
self._config = config
self._adapter = adapter
self._runner = runner
@classmethod
def from_bound_instance(
cls, bound_instance: BoundInstance
def from_shard_metadata(
cls, shard: ShardMetadata, group: mx.distributed.Group | None
) -> "DistributedImageModel":
model_id = bound_instance.bound_shard.model_card.model_id
model_id = shard.model_card.model_id
model_path = build_model_path(model_id)
shard_metadata = bound_instance.bound_shard
if not isinstance(shard_metadata, (PipelineShardMetadata, CfgShardMetadata)):
if not isinstance(shard, (PipelineShardMetadata, CfgShardMetadata)):
raise ValueError(
"Expected PipelineShardMetadata or CfgShardMetadata for image generation"
)
is_distributed = (
len(bound_instance.instance.shard_assignments.node_to_runner) > 1
)
if is_distributed:
logger.info("Starting distributed init for image model")
group = mlx_distributed_init(bound_instance)
else:
group = None
return cls(
model_id=model_id,
local_path=model_path,
shard_metadata=shard_metadata,
shard_metadata=shard,
group=group,
)
@@ -173,7 +168,3 @@ class DistributedImageModel:
else:
logger.info("generated image")
yield result
def initialize_image_model(bound_instance: BoundInstance) -> DistributedImageModel:
return DistributedImageModel.from_bound_instance(bound_instance)
+57 -19
View File
@@ -3,7 +3,7 @@ import io
import random
import tempfile
import time
from collections.abc import Callable
from collections.abc import Callable, Iterator
from pathlib import Path
from typing import Generator, Literal
@@ -17,11 +17,10 @@ from exo.api.types import (
ImageGenerationTaskParams,
ImageSize,
)
from exo.shared.constants import EXO_MAX_CHUNK_SIZE
from exo.shared.types.chunks import ImageChunk
from exo.shared.types.common import ModelId
from exo.shared.types.memory import Memory
from exo.shared.types.worker.runner_response import (
ImageGenerationResponse,
PartialImageResponse,
)
from exo.worker.engines.image.distributed_model import DistributedImageModel
@@ -71,16 +70,8 @@ def generate_image(
model: DistributedImageModel,
task: ImageGenerationTaskParams | ImageEditsTaskParams,
cancel_checker: Callable[[], bool] | None = None,
) -> Generator[ImageGenerationResponse | PartialImageResponse, None, None]:
"""Generate image(s), optionally yielding partial results.
When partial_images > 0 or stream=True, yields PartialImageResponse for
intermediate images, then ImageGenerationResponse for the final image.
Yields:
PartialImageResponse for intermediate images (if partial_images > 0, first image only)
ImageGenerationResponse for final complete images
"""
) -> Generator[ImageChunk, None, None]:
"""Generate image(s), optionally yielding partial results."""
width, height = parse_size(task.size)
quality: Literal["low", "medium", "high"] = task.quality or "medium"
@@ -142,12 +133,14 @@ def generate_image(
image = image.convert("RGB")
image.save(buffer, format=image_format)
yield PartialImageResponse(
yield from _process_image_response(
image_data=buffer.getvalue(),
format=task.output_format,
image_format=task.output_format,
partial_index=partial_idx,
total_partials=total_partials,
image_index=image_num,
model_id=model.model_id,
stats=None,
)
else:
image = result
@@ -189,9 +182,54 @@ def generate_image(
image = image.convert("RGB")
image.save(buffer, format=image_format)
yield ImageGenerationResponse(
yield from _process_image_response(
image_data=buffer.getvalue(),
format=task.output_format,
image_format=task.output_format,
stats=stats,
image_index=image_num,
model_id=model.model_id,
partial_index=None,
total_partials=None,
)
def _process_image_response(
image_data: bytes,
image_index: int,
image_format: Literal["png", "jpeg", "webp"],
partial_index: int | None,
total_partials: int | None,
stats: ImageGenerationStats | None,
model_id: ModelId,
) -> Iterator[ImageChunk]:
"""Process a single image response and send chunks."""
is_partial = partial_index is not None
encoded_data = base64.b64encode(image_data).decode("utf-8")
# Extract stats from final ImageGenerationResponse if available
data_chunks = [
encoded_data[i : i + EXO_MAX_CHUNK_SIZE]
for i in range(0, len(encoded_data), EXO_MAX_CHUNK_SIZE)
]
total_chunks = len(data_chunks)
def _data_to_chunk(item: tuple[int, str]) -> ImageChunk:
chunk_index, chunk_data = item
# Only include stats on the last chunk of the final image
chunk_stats = (
stats if chunk_index == total_chunks - 1 and not is_partial else None
)
return ImageChunk(
model=model_id,
data=chunk_data,
chunk_index=chunk_index,
total_chunks=total_chunks,
image_index=image_index,
is_partial=is_partial,
partial_index=partial_index,
total_partials=total_partials,
stats=chunk_stats,
format=image_format,
)
return map(_data_to_chunk, enumerate(data_chunks))
+35 -52
View File
@@ -1,5 +1,5 @@
from abc import ABC, abstractmethod
from collections.abc import Callable
from collections.abc import Callable, Generator
from functools import partial
from inspect import signature
from typing import TYPE_CHECKING, Literal, Protocol, cast
@@ -59,14 +59,13 @@ from mlx_lm.models.step3p5 import Model as Step35Model
from mlx_lm.models.step3p5 import Step3p5MLP as Step35MLP
from mlx_lm.models.step3p5 import Step3p5Model as Step35InnerModel
from exo.shared.types.worker.runner_response import ModelLoadingResponse
from exo.shared.types.worker.shards import PipelineShardMetadata
from exo.worker.runner.bootstrap import logger
if TYPE_CHECKING:
from mlx_lm.models.cache import Cache
LayerLoadedCallback = Callable[[int, int], None] # (layers_loaded, total_layers)
_pending_prefill_sends: list[tuple[mx.array, int, mx.distributed.Group]] = []
@@ -276,8 +275,7 @@ def pipeline_auto_parallel(
model: nn.Module,
group: mx.distributed.Group,
model_shard_meta: PipelineShardMetadata,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
"""
Automatically parallelize a model across multiple devices.
Args:
@@ -297,8 +295,7 @@ def pipeline_auto_parallel(
total = len(layers)
for i, layer in enumerate(layers):
mx.eval(layer) # type: ignore
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
layers[0] = PipelineFirstLayer(layers[0], device_rank, group=group)
layers[-1] = PipelineLastLayer(
@@ -460,8 +457,7 @@ def patch_tensor_model[T](model: T) -> T:
def tensor_auto_parallel(
model: nn.Module,
group: mx.distributed.Group,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
all_to_sharded_linear = partial(
shard_linear,
sharding="all-to-sharded",
@@ -595,7 +591,7 @@ def tensor_auto_parallel(
else:
raise ValueError(f"Unsupported model type: {type(model)}")
model = tensor_parallel_sharding_strategy.shard_model(model, on_layer_loaded)
model = yield from tensor_parallel_sharding_strategy.shard_model(model)
return patch_tensor_model(model)
@@ -619,16 +615,14 @@ class TensorParallelShardingStrategy(ABC):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module: ...
) -> Generator[ModelLoadingResponse, None, nn.Module]: ...
class LlamaShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(LlamaModel, model)
total = len(model.layers)
for i, layer in enumerate(model.layers):
@@ -646,8 +640,8 @@ class LlamaShardingStrategy(TensorParallelShardingStrategy):
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -681,8 +675,7 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(DeepseekV3Model, model)
total = len(model.layers)
@@ -738,8 +731,8 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
layer.mlp.sharding_group = self.group
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -764,8 +757,7 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(GLM4MoeLiteModel, model)
total = len(model.layers) # type: ignore
for i, layer in enumerate(model.layers): # type: ignore
@@ -816,8 +808,8 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
layer.mlp.sharding_group = self.group # type: ignore
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -904,8 +896,7 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(MiniMaxModel, model)
total = len(model.layers)
for i, layer in enumerate(model.layers):
@@ -934,8 +925,8 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
layer.block_sparse_moe = ShardedMoE(layer.block_sparse_moe) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.block_sparse_moe.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -943,8 +934,7 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(
Qwen3Model
| Qwen3MoeModel
@@ -1099,8 +1089,8 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -1108,8 +1098,7 @@ class Glm4MoeShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(Glm4MoeModel, model)
total = len(model.layers)
for i, layer in enumerate(model.layers):
@@ -1145,8 +1134,8 @@ class Glm4MoeShardingStrategy(TensorParallelShardingStrategy):
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -1154,8 +1143,7 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(GptOssMoeModel, model)
total = len(model.layers)
@@ -1186,8 +1174,8 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
layer.mlp.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -1195,8 +1183,7 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(Step35Model, model)
total = len(model.layers)
@@ -1229,8 +1216,8 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
@@ -1238,8 +1225,7 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(NemotronHModel, model)
rank = self.group.rank()
total = len(model.layers)
@@ -1272,8 +1258,7 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
layer.mixer = mixer # pyright: ignore[reportAttributeAccessIssue]
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
def _shard_mamba2_mixer(self, mixer: NemotronHMamba2Mixer, rank: int) -> None:
@@ -1380,8 +1365,7 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
def shard_model(
self,
model: nn.Module,
on_layer_loaded: LayerLoadedCallback | None,
) -> nn.Module:
) -> Generator[ModelLoadingResponse, None, nn.Module]:
model = cast(Gemma4Model, model)
layers = model.language_model.model.layers
total = len(layers)
@@ -1409,6 +1393,5 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
layer.experts.sharding_group = self.group
mx.eval(layer)
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
return model
+108
View File
@@ -0,0 +1,108 @@
import contextlib
import os
from collections.abc import Generator
from dataclasses import dataclass
import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.types.common import ModelId
from exo.shared.types.events import Event
from exo.shared.types.mlx import Model
from exo.shared.types.tasks import TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import ModelLoadingResponse
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.base import Builder, Engine
from exo.worker.engines.mlx.cache import KVPrefixCache
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
load_mlx_items,
)
from exo.worker.engines.mlx.vision import VisionProcessor
from exo.worker.runner.bootstrap import logger
from exo.worker.runner.llm_inference.batch_generator import (
BatchGenerator,
SequentialGenerator,
)
from exo.worker.runner.llm_inference.tool_parsers import make_mlx_parser
@dataclass
class MlxBuilder(Builder):
model_id: ModelId
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
inference_model: Model | None = None
tokenizer: TokenizerWrapper | None = None
group: mx.distributed.Group | None = None
vision_processor: VisionProcessor | None = None
def connect(self, bound_instance: BoundInstance) -> None:
self.group = initialize_mlx(bound_instance)
def load(self, bound_instance: BoundInstance) -> Generator[ModelLoadingResponse]:
(
self.inference_model,
self.tokenizer,
self.vision_processor,
) = yield from load_mlx_items(bound_instance, self.group)
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.inference_model, self.tokenizer, self.group
def build(
self,
) -> Engine:
assert self.inference_model
assert self.tokenizer
vision_processor = self.vision_processor
tool_parser = None
logger.info(
f"model has_tool_calling={self.tokenizer.has_tool_calling} using tokens {self.tokenizer.tool_call_start}, {self.tokenizer.tool_call_end}"
)
if (
self.tokenizer.tool_call_start
and self.tokenizer.tool_call_end
and self.tokenizer.tool_parser # type: ignore
):
tool_parser = make_mlx_parser(
self.tokenizer.tool_call_start,
self.tokenizer.tool_call_end,
self.tokenizer.tool_parser, # type: ignore
)
kv_prefix_cache = KVPrefixCache(self.group)
device_rank = 0 if self.group is None else self.group.rank()
if os.environ.get("EXO_NO_BATCH"):
logger.info("using SequentialGenerator (batching disabled)")
return SequentialGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
else:
logger.info("using BatchGenerator")
return BatchGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
@@ -457,6 +457,7 @@ class ExoBatchGenerator:
def close(self) -> None:
self._mlx_gen.close()
mx.clear_cache()
def _save_prefix_cache(
self,
+19 -22
View File
@@ -4,6 +4,7 @@ import re
import sys
import tempfile
import time
from collections.abc import Generator
from pathlib import Path
from typing import TYPE_CHECKING, Any, cast
@@ -51,6 +52,7 @@ from exo.shared.types.worker.instances import (
MlxJacclInstance,
MlxRingInstance,
)
from exo.shared.types.worker.runner_response import ModelLoadingResponse
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
@@ -58,7 +60,6 @@ from exo.shared.types.worker.shards import (
TensorShardMetadata,
)
from exo.worker.engines.mlx.auto_parallel import (
LayerLoadedCallback,
get_inner_model,
get_layers,
pipeline_auto_parallel,
@@ -66,8 +67,6 @@ from exo.worker.engines.mlx.auto_parallel import (
)
from exo.worker.runner.bootstrap import logger
Group = mx.distributed.Group
def get_weights_size(model_shard_meta: ShardMetadata) -> Memory:
return Memory.from_float_kb(
@@ -90,7 +89,7 @@ class HostList(RootModel[list[str]]):
def mlx_distributed_init(
bound_instance: BoundInstance,
) -> Group:
) -> mx.distributed.Group:
"""
Initialize MLX distributed.
"""
@@ -149,7 +148,7 @@ def mlx_distributed_init(
def initialize_mlx(
bound_instance: BoundInstance,
) -> Group:
) -> mx.distributed.Group:
# should we unseed it?
# TODO: pass in seed from params
mx.random.seed(42)
@@ -162,9 +161,10 @@ def initialize_mlx(
def load_mlx_items(
bound_instance: BoundInstance,
group: Group | None,
on_layer_loaded: LayerLoadedCallback | None,
) -> "tuple[Model, TokenizerWrapper, VisionProcessor | None]":
group: mx.distributed.Group | None,
) -> Generator[
ModelLoadingResponse, None, tuple[Model, TokenizerWrapper, "VisionProcessor | None"]
]:
if group is None:
logger.info(f"Single device used for {bound_instance.instance}")
model_path = build_model_path(bound_instance.bound_shard.model_card.model_id)
@@ -177,8 +177,7 @@ def load_mlx_items(
total = len(layers)
for i, layer in enumerate(layers):
mx.eval(layer) # type: ignore
if on_layer_loaded is not None:
on_layer_loaded(i, total)
yield ModelLoadingResponse(layers_loaded=i, total=total)
except ValueError as e:
logger.opt(exception=e).debug(
"Model architecture doesn't support layer-by-layer progress tracking",
@@ -191,10 +190,9 @@ def load_mlx_items(
else:
logger.info("Starting distributed init")
start_time = time.perf_counter()
model, tokenizer = shard_and_load(
model, tokenizer = yield from shard_and_load(
bound_instance.bound_shard,
group=group,
on_layer_loaded=on_layer_loaded,
)
end_time = time.perf_counter()
logger.info(
@@ -221,9 +219,8 @@ def load_mlx_items(
def shard_and_load(
shard_metadata: ShardMetadata,
group: Group,
on_layer_loaded: LayerLoadedCallback | None,
) -> tuple[nn.Module, TokenizerWrapper]:
group: mx.distributed.Group,
) -> Generator[ModelLoadingResponse, None, tuple[nn.Module, TokenizerWrapper]]:
model_path = build_model_path(shard_metadata.model_card.model_id)
model, _ = load_model(model_path, lazy=True, strict=False)
@@ -254,12 +251,10 @@ def shard_and_load(
match shard_metadata:
case TensorShardMetadata():
logger.info(f"loading model from {model_path} with tensor parallelism")
model = tensor_auto_parallel(model, group, on_layer_loaded)
model = yield from tensor_auto_parallel(model, group)
case PipelineShardMetadata():
logger.info(f"loading model from {model_path} with pipeline parallelism")
model = pipeline_auto_parallel(
model, group, shard_metadata, on_layer_loaded=on_layer_loaded
)
model = yield from pipeline_auto_parallel(model, group, shard_metadata)
mx.eval(model.parameters())
case CfgShardMetadata():
raise ValueError(
@@ -748,7 +743,9 @@ def set_wired_limit_for_model(model_size: Memory):
def mlx_cleanup(
model: Model | None, tokenizer: TokenizerWrapper | None, group: Group | None
model: Model | None,
tokenizer: TokenizerWrapper | None,
group: mx.distributed.Group | None,
) -> None:
del model, tokenizer, group
mx.clear_cache()
@@ -757,7 +754,7 @@ def mlx_cleanup(
gc.collect()
def mx_any(bool_: bool, group: Group | None) -> bool:
def mx_any(bool_: bool, group: mx.distributed.Group | None) -> bool:
if group is None:
return bool_
num_true = mx.distributed.all_sum(
@@ -767,7 +764,7 @@ def mx_any(bool_: bool, group: Group | None) -> bool:
return num_true.item() > 0
def mx_barrier(group: Group | None):
def mx_barrier(group: mx.distributed.Group | None):
if group is None:
return
mx.eval(
+21 -22
View File
@@ -8,7 +8,7 @@ from loguru import logger
from exo.api.types import ImageEditsTaskParams
from exo.download.download_utils import is_read_only_model_dir, resolve_existing_model
from exo.shared.apply import apply
from exo.routing.state_manager import StateManager
from exo.shared.constants import EXO_MAX_INSTANCE_RETRIES
from exo.shared.models.model_cards import ModelId, add_to_card_cache, delete_custom_card
from exo.shared.types.chunks import InputImageChunk
@@ -57,7 +57,7 @@ from exo.utils.info_gatherer.net_profile import check_reachable
from exo.utils.keyed_backoff import KeyedBackoff
from exo.utils.task_group import TaskGroup
from exo.worker.plan import plan
from exo.worker.runner.runner_supervisor import RunnerSupervisor
from exo.worker.runner.supervisor import RunnerSupervisor
class Worker:
@@ -71,6 +71,7 @@ class Worker:
# but I think it's the correct way to be thinking about commands
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
state_manager: StateManager[State],
api_port: int,
):
self.node_id: NodeId = node_id
@@ -80,7 +81,7 @@ class Worker:
self.download_command_sender = download_command_sender
self.api_port = api_port
self.state: State = State()
self.state_manager = state_manager
self.runners: dict[RunnerId, RunnerSupervisor] = {}
self._tg: TaskGroup = TaskGroup()
@@ -134,8 +135,6 @@ class Worker:
async def _event_applier(self):
with self.event_receiver as events:
async for event in events:
# 2. for each event, apply it to the state
self.state = apply(self.state, event=event)
event = event.event
if isinstance(event, InstanceDeleted):
@@ -161,14 +160,15 @@ class Worker:
async def plan_step(self):
while True:
state = self.state_manager.get()
await anyio.sleep(0.1)
task: Task | None = plan(
self.node_id,
self.runners,
self.state.downloads,
self.state.instances,
self.state.runners,
self.state.tasks,
state.downloads,
state.instances,
state.runners,
state.tasks,
self.input_chunk_buffer,
self._instance_backoff,
self._download_backoff,
@@ -305,7 +305,7 @@ class Worker:
del self.input_chunk_buffer[cmd_id]
if cmd_id in self.input_chunk_counts:
del self.input_chunk_counts[cmd_id]
await self._start_runner_task(modified_task)
await self._start_runner_task(state, modified_task)
case TextGeneration() if (
task.task_params.image_hashes
@@ -355,22 +355,22 @@ class Worker:
del self.input_chunk_buffer[cmd_id]
if cmd_id in self.input_chunk_counts:
del self.input_chunk_counts[cmd_id]
await self._start_runner_task(modified_task)
await self._start_runner_task(state, modified_task)
case LoadModel(instance_id=instance_id):
if (instance := self.state.instances.get(instance_id)) is not None:
if (instance := state.instances.get(instance_id)) is not None:
model_id = instance.shard_assignments.model_id
self._download_backoff.reset(model_id)
await self._start_runner_task(task)
await self._start_runner_task(state, task)
case task:
await self._start_runner_task(task)
await self._start_runner_task(state, task)
async def shutdown(self):
self._tg.cancel_tasks()
await self._stopped.wait()
async def _start_runner_task(self, task: Task):
if (instance := self.state.instances.get(task.instance_id)) is not None:
async def _start_runner_task(self, state: State, task: Task):
if (instance := state.instances.get(task.instance_id)) is not None:
await self.runners[
instance.shard_assignments.node_to_runner[self.node_id]
].start_task(task)
@@ -387,14 +387,13 @@ class Worker:
async def _poll_connection_updates(self):
while True:
edges = set(
conn.edge for conn in self.state.topology.out_edges(self.node_id)
)
state = self.state_manager.get()
edges = set(conn.edge for conn in state.topology.out_edges(self.node_id))
conns: defaultdict[NodeId, set[str]] = defaultdict(set)
async for ip, nid in check_reachable(
self.state.topology,
state.topology,
self.node_id,
self.state.node_network,
state.node_network,
api_port=self.api_port,
):
if ip in conns[nid]:
@@ -415,7 +414,7 @@ class Worker:
)
)
for conn in self.state.topology.out_edges(self.node_id):
for conn in state.topology.out_edges(self.node_id):
if not isinstance(conn.edge, SocketConnection):
continue
# ignore mDNS discovered connections
+1 -1
View File
@@ -40,7 +40,7 @@ from exo.shared.types.worker.runners import (
RunnerWarmingUp,
)
from exo.utils.keyed_backoff import KeyedBackoff
from exo.worker.runner.runner_supervisor import RunnerSupervisor
from exo.worker.runner.supervisor import RunnerSupervisor
def plan(
+19 -9
View File
@@ -8,6 +8,7 @@ from exo.shared.types.tasks import Task, TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runners import RunnerFailed
from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
from exo.worker.engines.base import Builder
logger: "loguru.Logger" = loguru.logger
@@ -35,23 +36,32 @@ def entrypoint(
# Import main after setting global logger - this lets us just import logger from this module
try:
if bound_instance.is_image_model:
from exo.worker.runner.image_models.runner import Runner as ImageRunner
from exo.worker.runner.runner import Runner
runner = ImageRunner(
bound_instance, event_sender, task_receiver, cancel_receiver
builder: Builder
if bound_instance.is_image_model:
from exo.worker.engines.image.builder import MfluxBuilder
builder = MfluxBuilder(
event_sender, cancel_receiver, bound_instance.bound_shard
)
runner.main()
else:
from exo.worker.engines.mlx.patches import apply_mlx_patches
from exo.worker.runner.llm_inference.runner import Runner
apply_mlx_patches()
runner = Runner(
bound_instance, event_sender, task_receiver, cancel_receiver
from exo.worker.engines.mlx.builder import MlxBuilder
# evil sharing of the event sender
builder = MlxBuilder(
model_id=bound_instance.bound_shard.model_card.model_id,
event_sender=event_sender,
cancel_receiver=cancel_receiver,
)
runner.main()
runner = Runner(bound_instance, builder, event_sender, task_receiver)
runner.main()
except ClosedResourceError:
logger.warning("Runner communication closed unexpectedly")
@@ -1,403 +0,0 @@
import base64
import time
from typing import TYPE_CHECKING, Literal
import mlx.core as mx
from exo.api.types import (
ImageEditsTaskParams,
ImageGenerationStats,
ImageGenerationTaskParams,
)
from exo.shared.constants import EXO_MAX_CHUNK_SIZE, EXO_TRACING_ENABLED
from exo.shared.models.model_cards import ModelTask
from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
from exo.shared.types.chunks import ErrorChunk, ImageChunk
from exo.shared.types.common import CommandId, ModelId
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
TraceEventData,
TracesCollected,
)
from exo.shared.types.tasks import (
CANCEL_ALL_TASKS,
ConnectToGroup,
ImageEdits,
ImageGeneration,
LoadModel,
Shutdown,
StartWarmup,
Task,
TaskId,
TaskStatus,
)
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
ImageGenerationResponse,
PartialImageResponse,
)
from exo.shared.types.worker.runners import (
RunnerConnected,
RunnerConnecting,
RunnerIdle,
RunnerLoaded,
RunnerLoading,
RunnerReady,
RunnerRunning,
RunnerShutdown,
RunnerShuttingDown,
RunnerStatus,
RunnerWarmingUp,
)
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
ShardMetadata,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.image import (
DistributedImageModel,
generate_image,
initialize_image_model,
warmup_image_generator,
)
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
)
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 _process_image_response(
response: ImageGenerationResponse | PartialImageResponse,
command_id: CommandId,
shard_metadata: ShardMetadata,
event_sender: MpSender[Event],
image_index: int,
) -> None:
"""Process a single image response and send chunks."""
encoded_data = base64.b64encode(response.image_data).decode("utf-8")
is_partial = isinstance(response, PartialImageResponse)
# Extract stats from final ImageGenerationResponse if available
stats = response.stats if isinstance(response, ImageGenerationResponse) else None
_send_image_chunk(
encoded_data=encoded_data,
command_id=command_id,
model_id=shard_metadata.model_card.model_id,
event_sender=event_sender,
image_index=response.image_index,
is_partial=is_partial,
partial_index=response.partial_index if is_partial else None,
total_partials=response.total_partials if is_partial else None,
stats=stats,
image_format=response.format,
)
def _send_traces_if_enabled(
event_sender: MpSender[Event],
task_id: TaskId,
rank: int,
) -> None:
if not EXO_TRACING_ENABLED:
return
traces = get_trace_buffer()
if traces:
trace_data = [
TraceEventData(
name=t.name,
start_us=t.start_us,
duration_us=t.duration_us,
rank=t.rank,
category=t.category,
)
for t in traces
]
event_sender.send(
TracesCollected(
task_id=task_id,
rank=rank,
traces=trace_data,
)
)
clear_trace_buffer()
def _send_image_chunk(
encoded_data: str,
command_id: CommandId,
model_id: ModelId,
event_sender: MpSender[Event],
image_index: int,
is_partial: bool,
partial_index: int | None = None,
total_partials: int | None = None,
stats: ImageGenerationStats | None = None,
image_format: Literal["png", "jpeg", "webp"] | None = None,
) -> None:
"""Send base64-encoded image data as chunks via events."""
data_chunks = [
encoded_data[i : i + EXO_MAX_CHUNK_SIZE]
for i in range(0, len(encoded_data), EXO_MAX_CHUNK_SIZE)
]
total_chunks = len(data_chunks)
for chunk_index, chunk_data in enumerate(data_chunks):
# Only include stats on the last chunk of the final image
chunk_stats = (
stats if chunk_index == total_chunks - 1 and not is_partial else None
)
event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=ImageChunk(
model=model_id,
data=chunk_data,
chunk_index=chunk_index,
total_chunks=total_chunks,
image_index=image_index,
is_partial=is_partial,
partial_index=partial_index,
total_partials=total_partials,
stats=chunk_stats,
format=image_format,
),
)
)
class Runner:
def __init__(
self,
bound_instance: BoundInstance,
event_sender: MpSender[Event],
task_receiver: MpReceiver[Task],
cancel_receiver: MpReceiver[TaskId],
):
self.event_sender = event_sender
self.task_receiver = task_receiver
self.cancel_receiver = cancel_receiver
self.bound_instance = bound_instance
self.instance, self.runner_id, self.shard_metadata = (
bound_instance.instance,
bound_instance.bound_runner_id,
bound_instance.bound_shard,
)
self.device_rank = self.shard_metadata.device_rank
logger.info("hello from the runner")
if getattr(self.shard_metadata, "immediate_exception", False):
raise Exception("Fake exception - runner failed to spin up.")
if timeout := getattr(self.shard_metadata, "should_timeout", 0):
time.sleep(timeout)
self.setup_start_time = time.time()
self.cancelled_tasks = set[TaskId]()
self.image_model: DistributedImageModel | None = None
self.group = None
self.current_status: RunnerStatus = RunnerIdle()
logger.info("runner created")
self.update_status(RunnerIdle())
self.seen = set[TaskId]()
def update_status(self, status: RunnerStatus):
self.current_status = status
self.event_sender.send(
RunnerStatusUpdated(
runner_id=self.runner_id, runner_status=self.current_status
)
)
def send_task_status(self, task: Task, status: TaskStatus):
self.event_sender.send(
TaskStatusUpdated(task_id=task.task_id, task_status=status)
)
def acknowledge_task(self, task: Task):
self.event_sender.send(TaskAcknowledged(task_id=task.task_id))
def _check_cancelled(self, task_id: TaskId) -> bool:
for cancel_id in self.cancel_receiver.collect():
self.cancelled_tasks.add(cancel_id)
return (
task_id in self.cancelled_tasks or CANCEL_ALL_TASKS in self.cancelled_tasks
)
def _run_image_task(
self,
task: Task,
task_params: ImageGenerationTaskParams | ImageEditsTaskParams,
command_id: CommandId,
) -> None:
assert self.image_model
logger.info(f"received image task: {str(task)[:500]}")
logger.info("runner running")
self.update_status(RunnerRunning())
self.acknowledge_task(task)
def cancel_checker() -> bool:
return self._check_cancelled(task.task_id)
try:
image_index = 0
for response in generate_image(
model=self.image_model,
task=task_params,
cancel_checker=cancel_checker,
):
if _is_primary_output_node(self.shard_metadata):
match response:
case PartialImageResponse():
logger.info(
f"sending partial ImageChunk {response.partial_index}/{response.total_partials}"
)
_process_image_response(
response,
command_id,
self.shard_metadata,
self.event_sender,
image_index,
)
case ImageGenerationResponse():
logger.info("sending final ImageChunk")
_process_image_response(
response,
command_id,
self.shard_metadata,
self.event_sender,
image_index,
)
image_index += 1
except Exception as e:
if _is_primary_output_node(self.shard_metadata):
self.event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=ErrorChunk(
model=self.shard_metadata.model_card.model_id,
finish_reason="error",
error_message=str(e),
),
)
)
raise
finally:
_send_traces_if_enabled(self.event_sender, task.task_id, self.device_rank)
self.current_status = RunnerReady()
logger.info("runner ready")
def main(self):
with self.task_receiver as tasks:
for task in tasks:
if task.task_id in self.seen:
logger.warning("repeat task - potential error")
self.seen.add(task.task_id)
self.cancelled_tasks.discard(CANCEL_ALL_TASKS)
self.send_task_status(task, TaskStatus.Running)
self.handle_task(task)
was_cancelled = (task.task_id in self.cancelled_tasks) or (
CANCEL_ALL_TASKS in self.cancelled_tasks
)
if not was_cancelled:
self.send_task_status(task, TaskStatus.Complete)
self.update_status(self.current_status)
if isinstance(self.current_status, RunnerShutdown):
break
def handle_task(self, task: Task):
match task:
case ConnectToGroup() if isinstance(self.current_status, RunnerIdle):
logger.info("runner connecting")
self.update_status(RunnerConnecting())
self.acknowledge_task(task)
self.group = initialize_mlx(self.bound_instance)
logger.info("runner connected")
self.current_status = RunnerConnected()
# we load the model if it's connected with a group, or idle without a group. we should never tell a model to connect if it doesn't need to
case LoadModel() if (
isinstance(self.current_status, RunnerConnected)
and self.group is not None
) or (isinstance(self.current_status, RunnerIdle) and self.group is None):
logger.info("runner loading")
self.update_status(RunnerLoading())
self.acknowledge_task(task)
assert (
ModelTask.TextToImage in self.shard_metadata.model_card.tasks
or ModelTask.ImageToImage in self.shard_metadata.model_card.tasks
), f"Incorrect model task(s): {self.shard_metadata.model_card.tasks}"
self.image_model = initialize_image_model(self.bound_instance)
self.current_status = RunnerLoaded()
logger.info("runner loaded")
case StartWarmup() if isinstance(self.current_status, RunnerLoaded):
logger.info("runner warming up")
self.update_status(RunnerWarmingUp())
self.acknowledge_task(task)
logger.info(f"warming up inference for instance: {self.instance}")
assert self.image_model
image = warmup_image_generator(model=self.image_model)
if image is not None:
logger.info(f"warmed up by generating {image.size} image")
else:
logger.info("warmup completed (non-primary node)")
logger.info(
f"runner initialized in {time.time() - self.setup_start_time} seconds"
)
self.current_status = RunnerReady()
logger.info("runner ready")
case (
ImageGeneration(task_params=task_params, command_id=command_id)
| ImageEdits(task_params=task_params, command_id=command_id)
) if isinstance(self.current_status, RunnerReady):
self._run_image_task(task, task_params, command_id)
case Shutdown():
logger.info("runner shutting down")
if not TYPE_CHECKING:
del self.image_model, self.group
mx.clear_cache()
import gc
gc.collect()
self.update_status(RunnerShuttingDown())
self.acknowledge_task(task)
self.current_status = RunnerShutdown()
case _:
raise ValueError(
f"Received {task.__class__.__name__} outside of state machine in {self.current_status=}"
)
@@ -1,22 +1,31 @@
import itertools
import time
from abc import ABC, abstractmethod
from collections import deque
from collections.abc import Generator, Iterable
from collections.abc import Generator, Iterator
from dataclasses import dataclass, field
import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.constants import EXO_MAX_CONCURRENT_REQUESTS
from exo.shared.types.chunks import ErrorChunk, PrefillProgressChunk
from exo.shared.types.chunks import ErrorChunk, GenerationChunk, PrefillProgressChunk
from exo.shared.types.common import ModelId
from exo.shared.types.events import ChunkGenerated, Event
from exo.shared.types.mlx import Model
from exo.shared.types.tasks import CANCEL_ALL_TASKS, TaskId, TextGeneration
from exo.shared.types.tasks import (
CANCEL_ALL_TASKS,
GenerationTask,
TaskId,
TextGeneration,
)
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.runner_response import GenerationResponse, ToolCallResponse
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
GenerationResponse,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.base import Engine
from exo.worker.engines.mlx.cache import KVPrefixCache
from exo.worker.engines.mlx.generator.batch_generate import ExoBatchGenerator
from exo.worker.engines.mlx.generator.generate import (
@@ -32,18 +41,10 @@ from exo.worker.engines.mlx.utils_mlx import (
from exo.worker.engines.mlx.vision import VisionProcessor
from exo.worker.runner.bootstrap import logger
from .model_output_parsers import apply_all_parsers
from .model_output_parsers import apply_all_parsers, map_responses_to_chunks
from .tool_parsers import ToolParser
class Cancelled:
pass
class Finished:
pass
class GeneratorQueue[T]:
def __init__(self):
self._q = deque[T]()
@@ -59,35 +60,6 @@ class GeneratorQueue[T]:
yield self._q.popleft()
class InferenceGenerator(ABC):
_cancelled_tasks: set[TaskId]
def should_cancel(self, task_id: TaskId) -> bool:
return (
task_id in self._cancelled_tasks
or CANCEL_ALL_TASKS in self._cancelled_tasks
)
@abstractmethod
def warmup(self) -> None: ...
@abstractmethod
def submit(
self,
task: TextGeneration,
) -> None: ...
@abstractmethod
def step(
self,
) -> Iterable[
tuple[TaskId, ToolCallResponse | GenerationResponse | Cancelled | Finished]
]: ...
@abstractmethod
def close(self) -> None: ...
EXO_RUNNER_MUST_FAIL = "EXO RUNNER MUST FAIL"
EXO_RUNNER_MUST_OOM = "EXO RUNNER MUST OOM"
EXO_RUNNER_MUST_TIMEOUT = "EXO RUNNER MUST TIMEOUT"
@@ -111,7 +83,7 @@ def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None:
@dataclass(eq=False)
class SequentialGenerator(InferenceGenerator):
class SequentialGenerator(Engine):
model: Model
tokenizer: TokenizerWrapper
group: mx.distributed.Group | None
@@ -137,7 +109,7 @@ class SequentialGenerator(InferenceGenerator):
# queue that the 1st generator should push to and 3rd generator should pull from
GeneratorQueue[GenerationResponse],
# generator to get parsed outputs
Generator[GenerationResponse | ToolCallResponse | None],
Iterator[GenerationChunk | None],
]
| None
) = field(default=None, init=False)
@@ -152,8 +124,9 @@ class SequentialGenerator(InferenceGenerator):
def submit(
self,
task: TextGeneration,
task: GenerationTask,
) -> None:
assert isinstance(task, TextGeneration)
self._cancelled_tasks.discard(CANCEL_ALL_TASKS)
self._all_tasks[task.task_id] = task
self._maybe_queue.append(task)
@@ -183,8 +156,8 @@ class SequentialGenerator(InferenceGenerator):
def step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
) -> Iterator[
tuple[TaskId, GenerationChunk | FinishedResponse | CancelledResponse]
]:
if self._active is None:
self.agree_on_tasks()
@@ -192,23 +165,25 @@ class SequentialGenerator(InferenceGenerator):
if self._queue:
self._start_next()
else:
return map(lambda task: (task, Cancelled()), self._cancelled_tasks)
return map(
lambda task: (task, CancelledResponse()), self._cancelled_tasks
)
assert self._active is not None
task, mlx_gen, queue, output_generator = self._active
task, gen, queue, output_generator = self._active
output: list[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
tuple[TaskId, GenerationChunk | CancelledResponse | FinishedResponse]
] = []
try:
response = next(mlx_gen)
response = next(gen)
queue.push(response)
# drain potentially many responses every time
while (parsed := next(output_generator, None)) is not None:
output.append((task.task_id, parsed))
except (StopIteration, PrefillCancelled):
output.append((task.task_id, Finished()))
output.append((task.task_id, FinishedResponse()))
self._active = None
if self._queue:
self._start_next()
@@ -220,20 +195,22 @@ class SequentialGenerator(InferenceGenerator):
return itertools.chain(
output,
map(lambda task: (task, Cancelled()), self._cancelled_tasks),
map(lambda task: (task, CancelledResponse()), self._cancelled_tasks),
)
def _start_next(self) -> None:
task = self._queue.popleft()
try:
mlx_gen = self._build_generator(task)
gen = self._build_generator(task)
except Exception as e:
self._send_error(task, e)
raise
queue = GeneratorQueue[GenerationResponse]()
if task.task_params.bench:
output_generator = queue.gen()
output_generator: Iterator[GenerationChunk | None] = map(
lambda r: map_responses_to_chunks(r, self.model_id), queue.gen()
)
else:
output_generator = apply_all_parsers(
queue.gen(),
@@ -244,7 +221,7 @@ class SequentialGenerator(InferenceGenerator):
self.model_id,
task.task_params.tools,
)
self._active = (task, mlx_gen, queue, output_generator)
self._active = (task, gen, queue, output_generator)
def _send_error(self, task: TextGeneration, e: Exception) -> None:
if self.device_rank == 0:
@@ -314,7 +291,7 @@ class SequentialGenerator(InferenceGenerator):
@dataclass(eq=False)
class BatchGenerator(InferenceGenerator):
class BatchGenerator(Engine):
model: Model
tokenizer: TokenizerWrapper
group: mx.distributed.Group | None
@@ -332,18 +309,18 @@ class BatchGenerator(InferenceGenerator):
_maybe_cancel: list[TextGeneration] = field(default_factory=list, init=False)
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
_mlx_gen: ExoBatchGenerator = field(init=False)
_gen: ExoBatchGenerator = field(init=False)
_active_tasks: dict[
int,
tuple[
TextGeneration,
GeneratorQueue[GenerationResponse],
Generator[GenerationResponse | ToolCallResponse | None],
Iterator[GenerationChunk | None],
],
] = field(default_factory=dict, init=False)
def __post_init__(self) -> None:
self._mlx_gen = ExoBatchGenerator(
self._gen = ExoBatchGenerator(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
@@ -361,8 +338,9 @@ class BatchGenerator(InferenceGenerator):
def submit(
self,
task: TextGeneration,
task: GenerationTask,
) -> None:
assert isinstance(task, TextGeneration)
self._cancelled_tasks.discard(CANCEL_ALL_TASKS)
self._all_tasks[task.task_id] = task
self._maybe_queue.append(task)
@@ -392,8 +370,8 @@ class BatchGenerator(InferenceGenerator):
def step(
self,
) -> Iterable[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
) -> Iterator[
tuple[TaskId, GenerationChunk | CancelledResponse | FinishedResponse]
]:
if not self._queue:
self.agree_on_tasks()
@@ -411,7 +389,9 @@ class BatchGenerator(InferenceGenerator):
queue = GeneratorQueue[GenerationResponse]()
if task.task_params.bench:
output_generator = queue.gen()
output_generator: Iterator[GenerationChunk | None] = map(
lambda r: map_responses_to_chunks(r, self.model_id), queue.gen()
)
else:
output_generator = apply_all_parsers(
queue.gen(),
@@ -424,13 +404,13 @@ class BatchGenerator(InferenceGenerator):
)
self._active_tasks[uid] = (task, queue, output_generator)
if not self._mlx_gen.has_work:
if not self._gen.has_work:
return self._apply_cancellations()
results = self._mlx_gen.step()
results = self._gen.step()
output: list[
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
tuple[TaskId, GenerationChunk | CancelledResponse | FinishedResponse]
] = []
for uid, response in results:
if uid not in self._active_tasks:
@@ -446,38 +426,38 @@ class BatchGenerator(InferenceGenerator):
# check if original response was terminal and append a Finished()
if response.finish_reason is not None:
output.append((task.task_id, Finished()))
output.append((task.task_id, FinishedResponse()))
del self._active_tasks[uid]
return itertools.chain(output, self._apply_cancellations())
def _apply_cancellations(
self,
) -> list[tuple[TaskId, Cancelled]]:
) -> Iterator[tuple[TaskId, CancelledResponse]]:
if not self._cancelled_tasks:
return []
return iter([])
cancel_all = CANCEL_ALL_TASKS in self._cancelled_tasks
uids_to_cancel: list[int] = []
results: list[tuple[TaskId, Cancelled]] = []
results: list[tuple[TaskId, CancelledResponse]] = []
for uid, (task, _, _) in list(self._active_tasks.items()):
if task.task_id in self._cancelled_tasks or cancel_all:
uids_to_cancel.append(uid)
results.append((task.task_id, Cancelled()))
results.append((task.task_id, CancelledResponse()))
del self._active_tasks[uid]
if uids_to_cancel:
self._mlx_gen.cancel(uids_to_cancel)
self._gen.cancel(uids_to_cancel)
already_cancelled = {tid for tid, _ in results}
for tid in self._cancelled_tasks:
if tid != CANCEL_ALL_TASKS and tid not in already_cancelled:
results.append((tid, Cancelled()))
results.append((tid, CancelledResponse()))
self._cancelled_tasks.clear()
return results
return iter(results)
def _send_error(self, task: TextGeneration, e: Exception) -> None:
if self.device_rank == 0:
@@ -529,7 +509,7 @@ class BatchGenerator(InferenceGenerator):
self.agree_on_tasks()
return self._mlx_gen.submit(
return self._gen.submit(
task_params=task.task_params,
prompt=prompt,
on_prefill_progress=on_prefill_progress,
@@ -538,5 +518,5 @@ class BatchGenerator(InferenceGenerator):
)
def close(self) -> None:
self._mlx_gen.close()
self._gen.close()
del self.model, self.tokenizer, self.group
@@ -1,4 +1,4 @@
from collections.abc import Generator
from collections.abc import Generator, Iterator
from functools import cache
from typing import Any
@@ -14,6 +14,12 @@ from openai_harmony import ( # pyright: ignore[reportMissingTypeStubs]
)
from exo.api.types import ToolCallItem
from exo.shared.types.chunks import (
ErrorChunk,
GenerationChunk,
TokenChunk,
ToolCallChunk,
)
from exo.shared.types.common import ModelId
from exo.shared.types.mlx import Model
from exo.shared.types.worker.runner_response import GenerationResponse, ToolCallResponse
@@ -64,29 +70,70 @@ def apply_all_parsers(
model_type: type[Model],
model_id: ModelId,
tools: list[dict[str, Any]] | None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
mlx_generator = receiver
) -> Iterator[GenerationChunk | None]:
generator = receiver
if issubclass(model_type, GptOssModel):
mlx_generator = parse_gpt_oss(mlx_generator)
generator = parse_gpt_oss(generator)
elif (
issubclass(model_type, DeepseekV32Model)
and "deepseek" in model_id.normalize().lower()
):
mlx_generator = parse_deepseek_v32(mlx_generator)
generator = parse_deepseek_v32(generator)
else:
if tokenizer.has_thinking:
mlx_generator = parse_thinking_models(
mlx_generator,
generator = parse_thinking_models(
generator,
tokenizer.think_start,
tokenizer.think_end,
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
)
if tool_parser:
mlx_generator = parse_tool_calls(mlx_generator, tool_parser, tools)
generator = parse_tool_calls(generator, tool_parser, tools)
return count_reasoning_tokens(mlx_generator)
generator = count_reasoning_tokens(generator)
return map(lambda r: map_responses_to_chunks(r, model_id), generator)
def map_responses_to_chunks(
response: GenerationResponse | ToolCallResponse | None, model_id: ModelId
) -> GenerationChunk | None:
match response:
case None:
return None
case GenerationResponse():
if response.finish_reason == "error":
return ErrorChunk(
error_message=response.text,
model=model_id,
)
else:
finish_reason = response.finish_reason
assert finish_reason not in (
"error",
"tool_calls",
"function_call",
)
return TokenChunk(
model=model_id,
text=response.text,
token_id=response.token,
usage=response.usage,
finish_reason=finish_reason,
stats=response.stats,
logprob=response.logprob,
top_logprobs=response.top_logprobs,
is_thinking=response.is_thinking,
)
case ToolCallResponse():
return ToolCallChunk(
tool_calls=response.tool_calls,
model=model_id,
usage=response.usage,
stats=response.stats,
)
def parse_gpt_oss(
@@ -1,434 +0,0 @@
import os
import time
from dataclasses import dataclass
from enum import Enum
import mlx.core as mx
from anyio import WouldBlock
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.models.model_cards import ModelTask
from exo.shared.types.chunks import (
ErrorChunk,
TokenChunk,
ToolCallChunk,
)
from exo.shared.types.common import CommandId, ModelId
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
)
from exo.shared.types.mlx import Model
from exo.shared.types.tasks import (
ConnectToGroup,
LoadModel,
Shutdown,
StartWarmup,
Task,
TaskId,
TaskStatus,
TextGeneration,
)
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
GenerationResponse,
ToolCallResponse,
)
from exo.shared.types.worker.runners import (
RunnerConnected,
RunnerConnecting,
RunnerIdle,
RunnerLoaded,
RunnerLoading,
RunnerReady,
RunnerRunning,
RunnerShutdown,
RunnerShuttingDown,
RunnerStatus,
RunnerWarmingUp,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.mlx.cache import KVPrefixCache
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
load_mlx_items,
)
from exo.worker.engines.mlx.vision import VisionProcessor
from exo.worker.runner.bootstrap import logger
from exo.worker.runner.llm_inference.batch_generator import (
BatchGenerator,
InferenceGenerator,
SequentialGenerator,
)
from .batch_generator import Cancelled, Finished
from .tool_parsers import make_mlx_parser
class ExitCode(str, Enum):
AllTasksComplete = "AllTasksComplete"
Shutdown = "Shutdown"
class Runner:
def __init__(
self,
bound_instance: BoundInstance,
event_sender: MpSender[Event],
task_receiver: MpReceiver[Task],
cancel_receiver: MpReceiver[TaskId],
):
self.event_sender = event_sender
self.task_receiver = task_receiver
self.cancel_receiver = cancel_receiver
self.bound_instance = bound_instance
self.instance, self.runner_id, self.shard_metadata = (
self.bound_instance.instance,
self.bound_instance.bound_runner_id,
self.bound_instance.bound_shard,
)
self.model_id = self.shard_metadata.model_card.model_id
self.device_rank = self.shard_metadata.device_rank
logger.info("hello from the runner")
if getattr(self.shard_metadata, "immediate_exception", False):
raise Exception("Fake exception - runner failed to spin up.")
if timeout := getattr(self.shard_metadata, "should_timeout", 0):
time.sleep(timeout)
self.setup_start_time = time.time()
self.generator: Builder | InferenceGenerator = Builder(
self.model_id,
self.event_sender,
self.cancel_receiver,
)
self.seen: set[TaskId] = set()
self.active_tasks: dict[
TaskId,
TextGeneration,
] = {}
logger.info("runner created")
self.update_status(RunnerIdle())
def update_status(self, status: RunnerStatus):
self.current_status = status
self.event_sender.send(
RunnerStatusUpdated(
runner_id=self.runner_id, runner_status=self.current_status
)
)
def send_task_status(self, task_id: TaskId, task_status: TaskStatus):
self.event_sender.send(
TaskStatusUpdated(task_id=task_id, task_status=task_status)
)
def acknowledge_task(self, task: Task):
self.event_sender.send(TaskAcknowledged(task_id=task.task_id))
def main(self):
with self.task_receiver:
for task in self.task_receiver:
if task.task_id in self.seen:
logger.warning("repeat task - potential error")
continue
self.seen.add(task.task_id)
self.handle_first_task(task)
if isinstance(self.current_status, RunnerShutdown):
break
def handle_first_task(self, task: Task):
self.send_task_status(task.task_id, TaskStatus.Running)
match task:
case ConnectToGroup() if isinstance(self.current_status, RunnerIdle):
assert isinstance(self.generator, Builder)
logger.info("runner connecting")
self.update_status(RunnerConnecting())
self.acknowledge_task(task)
self.generator.group = initialize_mlx(self.bound_instance)
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerConnected())
logger.info("runner connected")
# we load the model if it's connected with a group, or idle without a group. we should never tell a model to connect if it doesn't need to
case LoadModel() if isinstance(self.generator, Builder) and (
(
isinstance(self.current_status, RunnerConnected)
and self.generator.group is not None
)
or (
isinstance(self.current_status, RunnerIdle)
and self.generator.group is None
)
):
total_layers = (
self.shard_metadata.end_layer - self.shard_metadata.start_layer
)
logger.info("runner loading")
self.update_status(
RunnerLoading(layers_loaded=0, total_layers=total_layers)
)
self.acknowledge_task(task)
def on_layer_loaded(layers_loaded: int, total: int) -> None:
self.update_status(
RunnerLoading(layers_loaded=layers_loaded, total_layers=total)
)
assert (
ModelTask.TextGeneration in self.shard_metadata.model_card.tasks
), f"Incorrect model task(s): {self.shard_metadata.model_card.tasks}"
(
self.generator.inference_model,
self.generator.tokenizer,
self.generator.vision_processor,
) = load_mlx_items(
self.bound_instance,
self.generator.group,
on_layer_loaded=on_layer_loaded,
)
self.generator = self.generator.build()
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerLoaded())
logger.info("runner loaded")
case StartWarmup() if isinstance(self.current_status, RunnerLoaded):
assert isinstance(self.generator, InferenceGenerator)
logger.info("runner warming up")
self.update_status(RunnerWarmingUp())
self.acknowledge_task(task)
self.generator.warmup()
logger.info(
f"runner initialized in {time.time() - self.setup_start_time} seconds"
)
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerReady())
logger.info("runner ready")
case TextGeneration() if isinstance(self.current_status, RunnerReady):
return_code = self.handle_generation_tasks(starting_task=task)
if return_code == ExitCode.Shutdown:
return
case Shutdown():
self.shutdown(task)
return
case _:
raise ValueError(
f"Received {task.__class__.__name__} outside of state machine in {self.current_status=}"
)
def shutdown(self, task: Task):
logger.info("runner shutting down")
self.update_status(RunnerShuttingDown())
self.acknowledge_task(task)
if isinstance(self.generator, InferenceGenerator):
self.generator.close()
mx.clear_cache()
import gc
gc.collect()
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerShutdown())
def submit_text_generation(self, task: TextGeneration):
assert isinstance(self.generator, InferenceGenerator)
self.active_tasks[task.task_id] = task
self.generator.submit(task)
def handle_generation_tasks(self, starting_task: TextGeneration):
assert isinstance(self.current_status, RunnerReady)
assert isinstance(self.generator, InferenceGenerator)
logger.info(f"received chat request: {starting_task}")
self.update_status(RunnerRunning())
logger.info("runner running")
self.acknowledge_task(starting_task)
self.seen.add(starting_task.task_id)
self.submit_text_generation(starting_task)
while self.active_tasks:
results = self.generator.step()
finished: list[TaskId] = []
for task_id, result in results:
match result:
case Cancelled():
finished.append(task_id)
case Finished():
self.send_task_status(task_id, TaskStatus.Complete)
finished.append(task_id)
case _:
self.send_response(
result, self.active_tasks[task_id].command_id
)
for task_id in finished:
self.active_tasks.pop(task_id, None)
try:
task = self.task_receiver.receive_nowait()
if task.task_id in self.seen:
logger.warning("repeat task - potential error")
continue
self.seen.add(task.task_id)
match task:
case TextGeneration():
self.acknowledge_task(task)
self.submit_text_generation(task)
case Shutdown():
self.shutdown(task)
return ExitCode.Shutdown
case _:
raise ValueError(
f"Received {task.__class__.__name__} outside of state machine in {self.current_status=}"
)
except WouldBlock:
pass
self.update_status(RunnerReady())
logger.info("runner ready")
return ExitCode.AllTasksComplete
def send_response(
self,
response: GenerationResponse | ToolCallResponse,
command_id: CommandId,
):
match response:
case GenerationResponse():
if self.device_rank == 0 and response.finish_reason == "error":
self.event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=ErrorChunk(
error_message=response.text,
model=self.model_id,
),
)
)
elif self.device_rank == 0:
assert response.finish_reason not in (
"error",
"tool_calls",
"function_call",
)
self.event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=TokenChunk(
model=self.model_id,
text=response.text,
token_id=response.token,
usage=response.usage,
finish_reason=response.finish_reason,
stats=response.stats,
logprob=response.logprob,
top_logprobs=response.top_logprobs,
is_thinking=response.is_thinking,
),
)
)
case ToolCallResponse():
if self.device_rank == 0:
self.event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=ToolCallChunk(
tool_calls=response.tool_calls,
model=self.model_id,
usage=response.usage,
stats=response.stats,
),
)
)
@dataclass
class Builder:
model_id: ModelId
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
inference_model: Model | None = None
tokenizer: TokenizerWrapper | None = None
group: mx.distributed.Group | None = None
vision_processor: VisionProcessor | None = None
def build(
self,
) -> InferenceGenerator:
assert self.model_id
assert self.inference_model
assert self.tokenizer
vision_processor = self.vision_processor
tool_parser = None
logger.info(
f"model has_tool_calling={self.tokenizer.has_tool_calling} using tokens {self.tokenizer.tool_call_start}, {self.tokenizer.tool_call_end}"
)
if (
self.tokenizer.tool_call_start
and self.tokenizer.tool_call_end
and self.tokenizer.tool_parser # type: ignore
):
tool_parser = make_mlx_parser(
self.tokenizer.tool_call_start,
self.tokenizer.tool_call_end,
self.tokenizer.tool_parser, # type: ignore
)
kv_prefix_cache = KVPrefixCache(self.group)
device_rank = 0 if self.group is None else self.group.rank()
if os.environ.get("EXO_NO_BATCH"):
logger.info("using SequentialGenerator (batching disabled)")
return SequentialGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
logger.info("using BatchGenerator")
return BatchGenerator(
model=self.inference_model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
+279
View File
@@ -0,0 +1,279 @@
import time
from enum import Enum
from anyio import WouldBlock
from exo.shared.types.chunks import Chunk
from exo.shared.types.common import CommandId
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
)
from exo.shared.types.tasks import (
ConnectToGroup,
GenerationTask,
ImageEdits,
ImageGeneration,
LoadModel,
Shutdown,
StartWarmup,
Task,
TaskId,
TaskStatus,
TextGeneration,
)
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
)
from exo.shared.types.worker.runners import (
RunnerConnected,
RunnerConnecting,
RunnerIdle,
RunnerLoaded,
RunnerLoading,
RunnerReady,
RunnerRunning,
RunnerShutdown,
RunnerShuttingDown,
RunnerStatus,
RunnerWarmingUp,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.base import Builder, Engine
from .bootstrap import logger
class ExitCode(str, Enum):
AllTasksComplete = "AllTasksComplete"
Shutdown = "Shutdown"
class Runner:
def __init__(
self,
bound_instance: BoundInstance,
builder: Builder,
event_sender: MpSender[Event],
task_receiver: MpReceiver[Task],
):
self.event_sender = event_sender
self.task_receiver = task_receiver
self.bound_instance = bound_instance
self.instance, self.runner_id, self.shard_metadata = (
self.bound_instance.instance,
self.bound_instance.bound_runner_id,
self.bound_instance.bound_shard,
)
self.model_id = self.shard_metadata.model_card.model_id
self.device_rank = self.shard_metadata.device_rank
logger.info("hello from the runner")
if getattr(self.shard_metadata, "immediate_exception", False):
raise Exception("Fake exception - runner failed to spin up.")
if timeout := getattr(self.shard_metadata, "should_timeout", 0):
time.sleep(timeout)
self.setup_start_time = time.time()
self.generator: Builder | Engine = builder
self.seen: set[TaskId] = set()
self.active_tasks: dict[
TaskId,
GenerationTask,
] = {}
logger.info("runner created")
self.update_status(RunnerIdle())
def update_status(self, status: RunnerStatus):
self.current_status = status
self.event_sender.send(
RunnerStatusUpdated(
runner_id=self.runner_id, runner_status=self.current_status
)
)
def send_task_status(self, task_id: TaskId, task_status: TaskStatus):
self.event_sender.send(
TaskStatusUpdated(task_id=task_id, task_status=task_status)
)
def acknowledge_task(self, task: Task):
self.event_sender.send(TaskAcknowledged(task_id=task.task_id))
def main(self):
with self.task_receiver:
for task in self.task_receiver:
if task.task_id in self.seen:
logger.warning("repeat task - potential error")
continue
self.seen.add(task.task_id)
self.handle_first_task(task)
if isinstance(self.current_status, RunnerShutdown):
break
def handle_first_task(self, task: Task):
self.send_task_status(task.task_id, TaskStatus.Running)
match task:
case ConnectToGroup() if isinstance(self.current_status, RunnerIdle):
assert isinstance(self.generator, Builder)
logger.info("runner connecting")
self.update_status(RunnerConnecting())
self.acknowledge_task(task)
self.generator.connect(self.bound_instance)
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerConnected())
logger.info("runner connected")
# we load the model if it's connected with a group, or idle without a group. we should never tell a model to connect if it doesn't need to
case LoadModel() if isinstance(self.generator, Builder) and (
isinstance(self.current_status, (RunnerConnected, RunnerIdle))
):
total_layers = (
self.shard_metadata.end_layer - self.shard_metadata.start_layer
)
logger.info("runner loading")
self.update_status(
RunnerLoading(layers_loaded=0, total_layers=total_layers)
)
self.acknowledge_task(task)
for load_progress in self.generator.load(self.bound_instance):
self.update_status(
RunnerLoading(
layers_loaded=load_progress.layers_loaded,
total_layers=load_progress.total,
)
)
self.generator = self.generator.build()
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerLoaded())
logger.info("runner loaded")
case StartWarmup() if isinstance(self.current_status, RunnerLoaded):
assert isinstance(self.generator, Engine)
logger.info("runner warming up")
self.update_status(RunnerWarmingUp())
self.acknowledge_task(task)
self.generator.warmup()
logger.info(
f"runner initialized in {time.time() - self.setup_start_time} seconds"
)
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerReady())
logger.info("runner ready")
case TextGeneration() | ImageEdits() | ImageGeneration() if isinstance(
self.current_status, RunnerReady
):
return_code = self.handle_generation_tasks(starting_task=task)
if return_code == ExitCode.Shutdown:
return
case Shutdown():
self.shutdown(task)
return
case _:
raise ValueError(
f"Received {task.__class__.__name__} outside of state machine in {self.current_status=}"
)
def shutdown(self, task: Task):
logger.info("runner shutting down")
self.update_status(RunnerShuttingDown())
self.acknowledge_task(task)
self.generator.close()
import gc
gc.collect()
self.send_task_status(task.task_id, TaskStatus.Complete)
self.update_status(RunnerShutdown())
def submit_generation(self, task: GenerationTask):
assert isinstance(self.generator, Engine)
self.active_tasks[task.task_id] = task
self.generator.submit(task)
def handle_generation_tasks(self, starting_task: GenerationTask):
assert isinstance(self.current_status, RunnerReady)
assert isinstance(self.generator, Engine)
logger.info(f"received chat request: {starting_task}")
self.update_status(RunnerRunning())
logger.info("runner running")
self.acknowledge_task(starting_task)
self.seen.add(starting_task.task_id)
self.submit_generation(starting_task)
while self.active_tasks:
results = self.generator.step()
finished: list[TaskId] = []
for task_id, result in results:
match result:
case CancelledResponse():
finished.append(task_id)
case FinishedResponse():
self.send_task_status(task_id, TaskStatus.Complete)
finished.append(task_id)
case other:
self.send_chunk(other, self.active_tasks[task_id].command_id)
for task_id in finished:
self.active_tasks.pop(task_id, None)
try:
task = self.task_receiver.receive_nowait()
if task.task_id in self.seen:
logger.warning("repeat task - potential error")
continue
self.seen.add(task.task_id)
match task:
case TextGeneration() | ImageEdits() | ImageGeneration():
self.acknowledge_task(task)
self.submit_generation(task)
case Shutdown():
self.shutdown(task)
return ExitCode.Shutdown
case _:
raise ValueError(
f"Received {task.__class__.__name__} outside of state machine in {self.current_status=}"
)
except WouldBlock:
pass
self.update_status(RunnerReady())
logger.info("runner ready")
return ExitCode.AllTasksComplete
def send_chunk(
self,
chunk: Chunk,
command_id: CommandId,
):
if self.device_rank == 0:
self.event_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))
@@ -1,14 +1,13 @@
# Check tasks are complete before runner is ever ready.
import unittest.mock
from collections.abc import Iterable
from dataclasses import dataclass
from typing import Callable
import mlx.core as mx
import pytest
import exo.worker.engines.mlx.builder as mlx_builder
import exo.worker.runner.llm_inference.batch_generator as mlx_batch_generator
import exo.worker.runner.llm_inference.model_output_parsers as mlx_model_output_parsers
import exo.worker.runner.llm_inference.runner as mlx_runner
from exo.shared.types.chunks import TokenChunk
from exo.shared.types.events import (
ChunkGenerated,
@@ -46,6 +45,8 @@ from exo.shared.types.worker.runners import (
RunnerWarmingUp,
)
from exo.utils.channels import mp_channel
from exo.worker.engines.mlx.builder import MlxBuilder
from exo.worker.runner.runner import Runner
from ...constants import (
CHAT_COMPLETION_TASK_ID,
@@ -115,13 +116,22 @@ def assert_events_equal(test_events: Iterable[Event], true_events: Iterable[Even
assert test_event == true_event, f"{test_event} != {true_event}"
@dataclass
class MockLoadOutput:
layers_loaded: int
total: int
@pytest.fixture
def patch_out_mlx(monkeypatch: pytest.MonkeyPatch):
# initialize_mlx returns a mock group
monkeypatch.setattr(mlx_runner, "initialize_mlx", make_nothin(MockGroup()))
monkeypatch.setattr(
mlx_runner, "load_mlx_items", make_nothin((1, MockTokenizer, None))
)
monkeypatch.setattr(mlx_builder, "initialize_mlx", make_nothin(MockGroup()))
def lmi_gen():
yield MockLoadOutput(1, 1)
return (1, MockTokenizer, None)
monkeypatch.setattr(mlx_builder, "load_mlx_items", make_nothin(lmi_gen()))
monkeypatch.setattr(mlx_batch_generator, "warmup_inference", make_nothin(1))
monkeypatch.setattr(mlx_batch_generator, "_check_for_debug_prompts", nothin)
monkeypatch.setattr(mlx_batch_generator, "mx_any", make_nothin(False))
@@ -264,17 +274,18 @@ def _run(tasks: Iterable[Task], send_after_ready: list[Task] | None = None):
# this is some c++ nonsense
task_receiver.close = nothin
task_receiver.join = nothin
with unittest.mock.patch(
"exo.worker.runner.llm_inference.runner.mx.distributed.all_gather",
make_nothin(mx.array([1])),
):
runner = mlx_runner.Runner(
bound_instance,
event_sender, # pyright: ignore[reportArgumentType]
task_receiver,
cancel_receiver,
)
runner.main()
builder = MlxBuilder(
bound_instance.bound_shard.model_card.model_id,
event_sender, # pyright: ignore[reportArgumentType]
cancel_receiver,
)
runner = Runner(
bound_instance,
builder,
event_sender, # pyright: ignore[reportArgumentType]
task_receiver,
)
runner.main()
return event_sender.events
@@ -318,6 +329,10 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch):
runner_status=RunnerLoading(layers_loaded=0, total_layers=32),
),
TaskAcknowledged(task_id=LOAD_TASK_ID),
RunnerStatusUpdated(
runner_id=RUNNER_1_ID,
runner_status=RunnerLoading(layers_loaded=1, total_layers=1),
),
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),
@@ -17,7 +17,7 @@ from exo.shared.types.text_generation import (
from exo.shared.types.worker.instances import BoundInstance, InstanceId
from exo.shared.types.worker.runners import RunnerFailed, RunnerId
from exo.utils.channels import channel, mp_channel
from exo.worker.runner.runner_supervisor import RunnerSupervisor
from exo.worker.runner.supervisor import RunnerSupervisor
from exo.worker.tests.unittests.conftest import get_bound_mlx_ring_instance