151 lines
5.8 KiB
Python
151 lines
5.8 KiB
Python
#!/usr/bin/env python3
|
|
"""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
|
|
|
|
import gc
|
|
import json
|
|
import os
|
|
import signal
|
|
import time
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
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("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
|
|
|
|
if len(TOKEN) < 32:
|
|
raise RuntimeError("WORKER_TOKEN is missing or too short")
|
|
|
|
# The container is intentionally disposable. After the router has received a
|
|
# completed response and saved image, an immediate process exit releases the
|
|
# CUDA context much faster than Python/PyTorch interpreter teardown.
|
|
signal.signal(signal.SIGTERM, lambda *_: os._exit(0))
|
|
|
|
|
|
def load_pipeline() -> None:
|
|
global PIPE, LOAD_SECONDS
|
|
if PIPE is not None:
|
|
return
|
|
import torch
|
|
from diffusers import Flux2KleinPipeline
|
|
started = time.monotonic()
|
|
PIPE = Flux2KleinPipeline.from_pretrained(
|
|
MODEL_DIR, torch_dtype=torch.bfloat16, low_cpu_mem_usage=False)
|
|
# 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"):
|
|
PIPE.enable_vae_tiling()
|
|
LOAD_SECONDS = time.monotonic() - started
|
|
|
|
|
|
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:
|
|
raise ValueError("invalid prompt")
|
|
if (not isinstance(filename, str) or Path(filename).name != filename
|
|
or not filename.endswith(".png")):
|
|
raise ValueError("invalid filename")
|
|
width, height = int(data.get("width", 1024)), int(data.get("height", 1024))
|
|
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", 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()
|
|
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)
|
|
return {"status": "ok", "filename": filename,
|
|
"seconds": round(time.monotonic() - started, 3),
|
|
"load_seconds": round(LOAD_SECONDS, 3)}
|
|
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
def log_message(self, fmt: str, *args: object) -> None:
|
|
# Never log request bodies/prompts.
|
|
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()
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", "application/json")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def do_GET(self) -> None: # noqa: N802
|
|
if self.path == "/health":
|
|
self.reply(200, {"status": "ok", "model_loaded": PIPE is not None})
|
|
else:
|
|
self.reply(404, {"error": "not found"})
|
|
|
|
def do_POST(self) -> None: # noqa: N802
|
|
if self.headers.get("Authorization", "") != f"Bearer {TOKEN}":
|
|
self.reply(401, {"error": "unauthorized"})
|
|
return
|
|
if self.path != "/generate":
|
|
self.reply(404, {"error": "not found"})
|
|
return
|
|
try:
|
|
length = int(self.headers.get("Content-Length", "0"))
|
|
if length < 2 or length > 16384:
|
|
raise ValueError("invalid request size")
|
|
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
|
except Exception as exc:
|
|
print(f"[flux-image-worker] generation failed: "
|
|
f"{type(exc).__name__}: {str(exc)[:1000]}", flush=True)
|
|
self.reply(400, {"status": "error", "message": str(exc)})
|
|
|
|
|
|
try:
|
|
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|
|
finally:
|
|
if PIPE is not None:
|
|
del PIPE
|
|
gc.collect()
|