Add local FLUX image editing

This commit is contained in:
Mikei386 committed 2026-08-30 22:11:26 +02:00
1 parent 40e73d82a8
commit c27636ac34
13 files changed
+209 -64

No files matched your search

+42 -18
View File
@@ -1,5 +1,10 @@
#!/usr/bin/env python3
"""Private Z-Image-Turbo worker used only during a GPU hot swap."""
"""Private FLUX.2 Klein 4B worker used only during a GPU hot swap.
The same pipeline handles text-to-image and local reference-image editing.
Reference images are exchanged with the router through the shared image
volume; request bodies therefore never contain private image bytes here.
"""
from __future__ import annotations
@@ -14,7 +19,7 @@ from pathlib import Path
HOST = os.environ.get("WORKER_HOST", "0.0.0.0")
PORT = int(os.environ.get("WORKER_PORT", "8086"))
TOKEN = os.environ.get("WORKER_TOKEN", "").strip()
MODEL_DIR = os.environ.get("Z_IMAGE_MODEL_DIR", "/models/Z-Image-Turbo")
MODEL_DIR = os.environ.get("FLUX_MODEL_DIR", "/models/FLUX.2-klein-4B")
OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve()
PIPE = None
LOAD_SECONDS = 0.0
@@ -33,15 +38,13 @@ def load_pipeline() -> None:
if PIPE is not None:
return
import torch
from diffusers import ZImagePipeline
from diffusers import Flux2KleinPipeline
started = time.monotonic()
PIPE = ZImagePipeline.from_pretrained(
PIPE = Flux2KleinPipeline.from_pretrained(
MODEL_DIR, torch_dtype=torch.bfloat16, low_cpu_mem_usage=False)
# The Qwen text encoder and the DiT do not fit together in the usable
# 16 GiB of the RTX 5080. Sequential offload keeps only the active
# submodule on CUDA. This is slower than a fully resident pipeline, but
# deterministic and leaves the RTX 3060 available for XTTS.
PIPE.enable_sequential_cpu_offload()
# Officially supported low-VRAM path. It keeps the complete pipeline
# within the usable 16 GiB of the RTX 5080 and leaves the RTX 3060 alone.
PIPE.enable_model_cpu_offload()
if hasattr(PIPE, "enable_vae_slicing"):
PIPE.enable_vae_slicing()
if hasattr(PIPE, "enable_vae_tiling"):
@@ -51,6 +54,7 @@ def load_pipeline() -> None:
def generate(data: dict) -> dict:
import torch
from PIL import Image
prompt = data.get("prompt")
filename = data.get("filename")
if not isinstance(prompt, str) or not prompt.strip() or len(prompt) > 8000:
@@ -62,17 +66,37 @@ def generate(data: dict) -> dict:
if (width, height) not in {(1024, 1024), (1536, 1024), (1024, 1536),
(1920, 1088), (1088, 1920)}:
raise ValueError("unsupported image size")
steps = int(data.get("steps", 9))
guidance = float(data.get("guidance", 0.0))
if steps != 9 or guidance != 0.0:
raise ValueError("Z-Image-Turbo requires steps=9 and guidance=0.0")
steps = int(data.get("steps", 4))
guidance = float(data.get("guidance", 1.0))
if steps != 4 or guidance != 1.0:
raise ValueError("FLUX.2-klein-4B requires steps=4 and guidance=1.0")
source_files = data.get("source_files") or []
if not isinstance(source_files, list) or len(source_files) > 4:
raise ValueError("invalid source image list")
source_images = []
for source_name in source_files:
if not isinstance(source_name, str) or Path(source_name).name != source_name:
raise ValueError("invalid source image filename")
source = (OUTPUT_DIR / source_name).resolve()
if source.parent != OUTPUT_DIR or not source.is_file():
raise ValueError("source image not found")
with Image.open(source) as opened:
source_images.append(opened.convert("RGB"))
seed = data.get("seed")
generator = None if seed is None else torch.Generator(device="cuda").manual_seed(int(seed))
load_pipeline()
started = time.monotonic()
image = PIPE(prompt=prompt, height=height, width=width,
num_inference_steps=9, guidance_scale=0.0,
generator=generator).images[0]
kwargs = {
"prompt": prompt,
"height": height,
"width": width,
"num_inference_steps": 4,
"guidance_scale": 1.0,
"generator": generator,
}
if source_images:
kwargs["image"] = source_images[0] if len(source_images) == 1 else source_images
image = PIPE(**kwargs).images[0]
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
output = OUTPUT_DIR / filename
image.save(output)
@@ -84,7 +108,7 @@ def generate(data: dict) -> dict:
class Handler(BaseHTTPRequestHandler):
def log_message(self, fmt: str, *args: object) -> None:
# Never log request bodies/prompts.
print(f"[z-image-worker] {self.client_address[0]} {fmt % args}", flush=True)
print(f"[flux-image-worker] {self.client_address[0]} {fmt % args}", flush=True)
def reply(self, status: int, payload: dict) -> None:
body = json.dumps(payload, separators=(",", ":")).encode()
@@ -113,7 +137,7 @@ class Handler(BaseHTTPRequestHandler):
raise ValueError("invalid request size")
self.reply(200, generate(json.loads(self.rfile.read(length))))
except Exception as exc:
print(f"[z-image-worker] generation failed: "
print(f"[flux-image-worker] generation failed: "
f"{type(exc).__name__}: {str(exc)[:1000]}", flush=True)
self.reply(400, {"status": "error", "message": str(exc)})