#!/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()