Add initial image edits spec

This commit is contained in:
ciaranbor
2026-01-06 10:48:26 +00:00
parent 8b7d8ef394
commit dc661e4b5e
4 changed files with 114 additions and 3 deletions
+48
View File
@@ -16,6 +16,7 @@ from exo.shared.types.commands import (
CreateInstance,
DeleteInstance,
ForwarderCommand,
ImageEdits,
ImageGeneration,
PlaceInstance,
RequestEventLog,
@@ -36,6 +37,9 @@ from exo.shared.types.state import State
from exo.shared.types.tasks import (
ChatCompletion as ChatCompletionTask,
)
from exo.shared.types.tasks import (
ImageEdits as ImageEditsTask,
)
from exo.shared.types.tasks import (
ImageGeneration as ImageGenerationTask,
)
@@ -194,6 +198,50 @@ class Master:
)
)
self.command_task_mapping[command.command_id] = task_id
case ImageEdits():
# TODO: refactor with ChatCompletion
instance_task_counts: dict[InstanceId, int] = {}
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.request_params.model
):
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
)
instance_task_counts[instance.instance_id] = (
task_count
)
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.request_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[
instance_id
],
)
task_id = TaskId()
generated_events.append(
TaskCreated(
task_id=task_id,
task=ImageEditsTask(
task_id=task_id,
command_id=command.command_id,
instance_id=available_instance_ids[0],
task_status=TaskStatus.Pending,
task_params=command.request_params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
case DeleteInstance():
placement = delete_instance(command, self.state.instances)
+14 -1
View File
@@ -2,7 +2,11 @@ from enum import Enum
from pydantic import Field
from exo.shared.types.api import ChatCompletionTaskParams, ImageGenerationTaskParams
from exo.shared.types.api import (
ChatCompletionTaskParams,
ImageEditsTaskParams,
ImageGenerationTaskParams,
)
from exo.shared.types.common import CommandId, Id
from exo.shared.types.worker.instances import BoundInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId
@@ -64,6 +68,14 @@ class ImageGeneration(BaseTask): # emitted by Master
error_message: str | None = Field(default=None)
class ImageEdits(BaseTask): # emitted by Master
command_id: CommandId
task_params: ImageEditsTaskParams
error_type: str | None = Field(default=None)
error_message: str | None = Field(default=None)
class Shutdown(BaseTask): # emitted by Worker
runner_id: RunnerId
@@ -76,5 +88,6 @@ Task = (
| StartWarmup
| ChatCompletion
| ImageGeneration
| ImageEdits
| Shutdown
)
+6 -2
View File
@@ -9,6 +9,7 @@ from exo.shared.types.tasks import (
ConnectToGroup,
CreateRunner,
DownloadModel,
ImageEdits,
ImageGeneration,
LoadModel,
Shutdown,
@@ -266,8 +267,11 @@ def _pending_tasks(
) -> Task | None:
for task in tasks.values():
# for now, just forward chat completions
if not isinstance(task, ChatCompletion) and not isinstance(
task, ImageGeneration
# TODO: do this better!
if (
not isinstance(task, ChatCompletion)
and not isinstance(task, ImageGeneration)
and not isinstance(task, ImageEdits)
):
continue
if task.task_status not in (TaskStatus.Pending, TaskStatus.Running):
+46
View File
@@ -18,6 +18,7 @@ from exo.shared.types.tasks import (
ChatCompletion,
ConnectToGroup,
ImageGeneration,
ImageEdits,
LoadModel,
Shutdown,
StartWarmup,
@@ -237,6 +238,51 @@ def main(
)
)
# Generate images using MFlux (MLX) diffusion
for response in mflux_generate(
model=model,
task=task_params,
):
match response:
case ImageGenerationResponse():
if shard_metadata.device_rank == 0:
encoded_data = base64.b64encode(
response.image_data
).decode("utf-8")
event_sender.send(
ChunkGenerated(
command_id=command_id,
chunk=ImageChunk(
idx=0,
model=shard_metadata.model_meta.model_id,
data=encoded_data,
finish_reason=response.finish_reason,
),
)
)
current_status = RunnerReady()
logger.info("runner ready")
event_sender.send(
RunnerStatusUpdated(
runner_id=runner_id, runner_status=RunnerReady()
)
)
case ImageEdits(task_params=task_params, command_id=command_id) if (
isinstance(current_status, RunnerReady)
):
assert isinstance(model, Flux1)
logger.info(
f"received image generation request: {str(task)[:500]}"
)
current_status = RunnerRunning()
logger.info("runner running")
event_sender.send(
RunnerStatusUpdated(
runner_id=runner_id, runner_status=current_status
)
)
# Generate images using MFlux (MLX) diffusion
for response in mflux_generate(
model=model,