From 4b8976be514e2fbc0de20cd2790ee7d38b2ef8c3 Mon Sep 17 00:00:00 2001 From: ciaranbor Date: Thu, 8 Jan 2026 09:43:45 +0000 Subject: [PATCH] Support image editing --- src/exo/master/api.py | 104 +++++++++++++++++- src/exo/shared/types/api.py | 16 ++- src/exo/shared/types/commands.py | 4 +- src/exo/shared/types/tasks.py | 4 +- src/exo/worker/engines/image/base.py | 7 +- .../worker/engines/image/distributed_model.py | 10 +- src/exo/worker/engines/image/generate.py | 63 ++++++----- src/exo/worker/runner/runner.py | 2 +- 8 files changed, 173 insertions(+), 37 deletions(-) diff --git a/src/exo/master/api.py b/src/exo/master/api.py index e31d573f..72a05227 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -1,11 +1,12 @@ +import base64 import time from collections.abc import AsyncGenerator -from typing import cast +from typing import Literal, cast import anyio from anyio import create_task_group from anyio.abc import TaskGroup -from fastapi import FastAPI, HTTPException +from fastapi import FastAPI, File, Form, HTTPException, UploadFile from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse from fastapi.staticfiles import StaticFiles @@ -35,7 +36,7 @@ from exo.shared.types.api import ( DeleteInstanceResponse, FinishReason, ImageData, - ImageEditsTaskParams, + ImageEditsInternalParams, ImageGenerationResponse, ImageGenerationTaskParams, ModelList, @@ -192,7 +193,7 @@ class API: self.chat_completions ) self.app.post("/v1/images/generations")(self.image_generations) - # self.app.post("/v1/images/edits")(self.image_edits) + self.app.post("/v1/images/edits")(self.image_edits) self.app.get("/state")(lambda: self.state) self.app.get("/events")(lambda: self._event_log) @@ -627,6 +628,101 @@ class API: await self._send(TaskFinished(finished_command_id=command.command_id)) del self._image_generation_queues[command.command_id] + async def image_edits( + self, + image: UploadFile = File(...), + prompt: str = Form(...), + model: str = Form(...), + n: int = Form(1), + size: str = Form("1024x1024"), + response_format: Literal["url", "b64_json"] = Form("b64_json"), + input_fidelity: Literal["low", "high"] = Form("low"), + ) -> ImageGenerationResponse: + """Handle image editing requests (img2img).""" + model_meta = await resolve_model_meta(model) + resolved_model = model_meta.model_id + + if not any( + instance.shard_assignments.model_id == resolved_model + for instance in self.state.instances.values() + ): + await self._trigger_notify_user_to_download_model(resolved_model) + raise HTTPException( + status_code=404, detail=f"No instance found for model {resolved_model}" + ) + + # Read and base64 encode the uploaded image + image_content = await image.read() + image_data = base64.b64encode(image_content).decode("utf-8") + + # Map input_fidelity to image_strength + image_strength = 0.3 if input_fidelity == "high" else 0.7 + + # Create internal params + internal_params = ImageEditsInternalParams( + image_data=image_data, + prompt=prompt, + model=resolved_model, + n=n, + size=size, + response_format=response_format, + image_strength=image_strength, + ) + + command = ImageEdits( + request_params=internal_params, + ) + await self._send(command) + + num_images = n + + # Track chunks per image: {image_index: {chunk_index: data}} + image_chunks: dict[int, dict[int, str]] = {} + image_total_chunks: dict[int, int] = {} + images_complete = 0 + + try: + self._image_generation_queues[command.command_id], recv = channel[ + ImageChunk + ]() + + while images_complete < num_images: + with recv as chunks: + async for chunk in chunks: + if chunk.image_index not in image_chunks: + image_chunks[chunk.image_index] = {} + image_total_chunks[chunk.image_index] = chunk.total_chunks + + image_chunks[chunk.image_index][chunk.chunk_index] = chunk.data + + if ( + len(image_chunks[chunk.image_index]) + == image_total_chunks[chunk.image_index] + ): + images_complete += 1 + + if images_complete >= num_images: + break + + images: list[ImageData] = [] + for image_idx in range(num_images): + chunks_dict = image_chunks[image_idx] + full_data = "".join(chunks_dict[i] for i in range(len(chunks_dict))) + images.append( + ImageData( + b64_json=full_data if response_format == "b64_json" else None, + url=None, # URL format not implemented yet + ) + ) + + return ImageGenerationResponse(data=images) + except anyio.get_cancelled_exc_class(): + raise + finally: + # Send TaskFinished command + await self._send(TaskFinished(finished_command_id=command.command_id)) + del self._image_generation_queues[command.command_id] + def _calculate_total_available_memory(self) -> Memory: """Calculate total available memory across all nodes in bytes.""" total_available = Memory() diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py index c34cbb5f..f193f839 100644 --- a/src/exo/shared/types/api.py +++ b/src/exo/shared/types/api.py @@ -198,8 +198,8 @@ class ImageGenerationTaskParams(BaseModel): class ImageEditsTaskParams(BaseModel): image: UploadFile - mask: UploadFile | None prompt: str + input_fidelity: float = 0.7 model: str n: int | None = 1 quality: Literal["high", "medium", "low"] | None = "medium" @@ -209,6 +209,20 @@ class ImageEditsTaskParams(BaseModel): user: str | None = None +class ImageEditsInternalParams(BaseModel): + """Serializable version of ImageEditsTaskParams for distributed task execution.""" + + image_data: str # Base64-encoded image + prompt: str + model: str + n: int | None = 1 + quality: Literal["high", "medium", "low"] | None = "medium" + output_format: Literal["png", "jpeg", "webp"] = "png" + response_format: Literal["url", "b64_json"] | None = "b64_json" + size: str | None = "1024x1024" + image_strength: float = 0.7 + + class ImageData(BaseModel): b64_json: str | None = None url: str | None = None diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py index addd4e29..bf6fd4e3 100644 --- a/src/exo/shared/types/commands.py +++ b/src/exo/shared/types/commands.py @@ -2,7 +2,7 @@ from pydantic import Field from exo.shared.types.api import ( ChatCompletionTaskParams, - ImageEditsTaskParams, + ImageEditsInternalParams, ImageGenerationTaskParams, ) from exo.shared.types.common import CommandId, NodeId @@ -29,7 +29,7 @@ class ImageGeneration(BaseCommand): class ImageEdits(BaseCommand): - request_params: ImageEditsTaskParams + request_params: ImageEditsInternalParams class PlaceInstance(BaseCommand): diff --git a/src/exo/shared/types/tasks.py b/src/exo/shared/types/tasks.py index 7974bf9d..2c58087e 100644 --- a/src/exo/shared/types/tasks.py +++ b/src/exo/shared/types/tasks.py @@ -4,7 +4,7 @@ from pydantic import Field from exo.shared.types.api import ( ChatCompletionTaskParams, - ImageEditsTaskParams, + ImageEditsInternalParams, ImageGenerationTaskParams, ) from exo.shared.types.common import CommandId, Id @@ -70,7 +70,7 @@ class ImageGeneration(BaseTask): # emitted by Master class ImageEdits(BaseTask): # emitted by Master command_id: CommandId - task_params: ImageEditsTaskParams + task_params: ImageEditsInternalParams error_type: str | None = Field(default=None) error_message: str | None = Field(default=None) diff --git a/src/exo/worker/engines/image/base.py b/src/exo/worker/engines/image/base.py index 0de15615..e59a1d0e 100644 --- a/src/exo/worker/engines/image/base.py +++ b/src/exo/worker/engines/image/base.py @@ -1,3 +1,4 @@ +from pathlib import Path from typing import Literal, Optional, Protocol, runtime_checkable from PIL import Image @@ -18,8 +19,10 @@ class ImageGenerator(Protocol): width: int, quality: Literal["low", "medium", "high"], seed: int, + image_path: Path | None = None, + image_strength: float | None = None, ) -> Optional[Image.Image]: - """Generate an image from a text prompt. + """Generate an image from a text prompt, or edit an existing image. For distributed inference, only the first stage (rank 0) returns the image. Other stages return None after participating in the pipeline. @@ -30,6 +33,8 @@ class ImageGenerator(Protocol): width: Image width in pixels quality: Generation quality level seed: Random seed for reproducibility + image_path: Optional path to input image for img2img + image_strength: Optional strength for img2img (0.0-1.0, higher = more change) Returns: Generated PIL Image (rank 0) or None (other ranks) diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index b6e097b4..e385ddef 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -184,11 +184,19 @@ class DistributedImageModel: width: int, quality: Literal["low", "medium", "high"] = "medium", seed: int = 2, + image_path: Path | None = None, + image_strength: float | None = None, ) -> Optional[Image.Image]: # Determine number of inference steps based on quality steps = self._config.get_steps_for_quality(quality) - config = Config(num_inference_steps=steps, height=height, width=width) + config = Config( + num_inference_steps=steps, + height=height, + width=width, + image_path=image_path, + image_strength=image_strength, + ) image = self._generate_image(settings=config, prompt=prompt, seed=seed) logger.info("generated image") diff --git a/src/exo/worker/engines/image/generate.py b/src/exo/worker/engines/image/generate.py index d46ddbe4..22dee493 100644 --- a/src/exo/worker/engines/image/generate.py +++ b/src/exo/worker/engines/image/generate.py @@ -1,9 +1,12 @@ +import base64 import io +import tempfile +from pathlib import Path from typing import Generator, Literal from PIL import Image -from exo.shared.types.api import ImageGenerationTaskParams +from exo.shared.types.api import ImageEditsInternalParams, ImageGenerationTaskParams from exo.shared.types.worker.runner_response import ImageGenerationResponse from exo.worker.engines.image.base import ImageGenerator @@ -37,36 +40,46 @@ def warmup_image_generator(model: ImageGenerator) -> Image.Image | None: def generate_image( model: ImageGenerator, - task: ImageGenerationTaskParams, + task: ImageGenerationTaskParams | ImageEditsInternalParams, ) -> Generator[ImageGenerationResponse, None, None]: - # Parse parameters width, height = parse_size(task.size) quality: Literal["low", "medium", "high"] = task.quality or "medium" seed = 2 # TODO(ciaran): Randomise when not testing anymore - # Generate using the model's generate method - image = model.generate( - prompt=task.prompt, - height=height, - width=width, - quality=quality, - seed=seed, - ) + image_path: Path | None = None + image_strength: float | None = None - # Only rank 0 returns the image - if image is None: - return + with tempfile.TemporaryDirectory() as tmpdir: + if isinstance(task, ImageEditsInternalParams): + # Decode base64 image data and save to temp file + image_path = Path(tmpdir) / "input.png" + image_path.write_bytes(base64.b64decode(task.image_data)) + image_strength = task.image_strength - buffer = io.BytesIO() - image_format = task.output_format.upper() - if image_format == "JPG": - image_format = "JPEG" + image = model.generate( + prompt=task.prompt, + height=height, + width=width, + quality=quality, + seed=seed, + image_path=image_path, + image_strength=image_strength, + ) - image.save(buffer, format=image_format) - image_bytes = buffer.getvalue() + # Only final rank returns the image + if image is None: + return - # Send complete image as single response (no artificial chunking) - yield ImageGenerationResponse( - image_data=image_bytes, - format=task.output_format, - ) + buffer = io.BytesIO() + image_format = task.output_format.upper() + if image_format == "JPG": + image_format = "JPEG" + + image.save(buffer, format=image_format) + image_bytes = buffer.getvalue() + + # Send complete image as single response (no artificial chunking) + yield ImageGenerationResponse( + image_data=image_bytes, + format=task.output_format, + ) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 52979aa2..da0bae00 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -315,7 +315,7 @@ def main( ): match response: case ImageGenerationResponse(): - if shard_metadata.device_rank == 0: + if shard_metadata.device_rank == shard_metadata.world_size - 1: encoded_data = base64.b64encode( response.image_data ).decode("utf-8")