From 0ed81d8afa917de505eb2fde9c005754bab83707 Mon Sep 17 00:00:00 2001 From: ciaranbor Date: Tue, 2 Dec 2025 15:30:36 +0000 Subject: [PATCH] Implement image generation warmup --- src/exo/worker/engines/mflux/generator/generate.py | 12 ++++++++++++ src/exo/worker/runner/runner.py | 12 ++++++++++-- 2 files changed, 22 insertions(+), 2 deletions(-) diff --git a/src/exo/worker/engines/mflux/generator/generate.py b/src/exo/worker/engines/mflux/generator/generate.py index e3bdcc8a..872a15ec 100644 --- a/src/exo/worker/engines/mflux/generator/generate.py +++ b/src/exo/worker/engines/mflux/generator/generate.py @@ -1,6 +1,7 @@ import io from typing import Generator +from PIL import Image import mlx.core as mx from mflux.config.config import Config from mflux.models.flux.variants.txt2img.flux import Flux1 @@ -30,6 +31,17 @@ def parse_size(size_str: str | None) -> tuple[int, int]: return (1024, 1024) +def warmup_mflux(model: Flux1) -> Image.Image: + prompt = "Warmup" + image = model.generate_image( + seed=2, + prompt=prompt, + config=Config(num_inference_steps=1, height=256, width=256), + ) + + return image.image + + def mflux_generate( model: Flux1, task: ImageGenerationTaskParams, diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 7f60cfc4..703cbcc4 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -17,8 +17,8 @@ from exo.shared.types.models import ModelTask from exo.shared.types.tasks import ( ChatCompletion, ConnectToGroup, - ImageGeneration, ImageEdits, + ImageGeneration, LoadModel, Shutdown, StartWarmup, @@ -45,7 +45,7 @@ from exo.shared.types.worker.runners import ( RunnerWarmingUp, ) from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender -from exo.worker.engines.mflux.generator.generate import mflux_generate +from exo.worker.engines.mflux.generator.generate import mflux_generate, warmup_mflux from exo.worker.engines.mflux.utils_mflux import initialize_mflux from exo.worker.engines.mlx.generator.generate import mlx_generate, warmup_inference from exo.worker.engines.mlx.utils_mlx import ( @@ -171,6 +171,14 @@ def main( logger.info( f"runner initialized in {time.time() - setup_start_time} seconds" ) + elif ( + ModelTask.TextToImage in model_tasks + or ModelTask.ImageToImage in model_tasks + ): + assert isinstance(model, Flux1) + image = warmup_mflux(model=model) + logger.info(f"warmed up by generating {image.size} image") + current_status = RunnerReady() logger.info("runner ready") case ChatCompletion(