Implement image generation warmup

This commit is contained in:
ciaranbor
2026-01-06 10:50:34 +00:00
parent 66a24d59b9
commit 0ed81d8afa
2 changed files with 22 additions and 2 deletions
@@ -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,
+10 -2
View File
@@ -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(