Implement image generation warmup
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user