mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-08 11:35:40 -04:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7d949a08cf | ||
|
|
b7ed1f9dc6 | ||
|
|
0bc29c6c9d | ||
|
|
2229b99bec | ||
|
|
0a77426276 |
No files matched your search
@@ -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
@@ -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
@@ -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
@@ -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)}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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:
|
||||
|
||||
@@ -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: ...
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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))
|
||||
File renamed without changes.
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user