Make Qwen Image 2.1 the production image worker
This commit is contained in:
@@ -0,0 +1,268 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Router-compatible Qwen-Image-2.1 worker backed by pinned ComfyUI."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import shutil
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
import urllib.request
|
||||
import uuid
|
||||
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()
|
||||
COMFY = "http://127.0.0.1:8188"
|
||||
COMFY_DIR = Path("/opt/ComfyUI")
|
||||
INPUT_DIR = COMFY_DIR / "input"
|
||||
OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve()
|
||||
MODEL = "Qwen-Image-2.1-int8"
|
||||
GENERATION_LOCK = threading.Lock()
|
||||
READY = threading.Event()
|
||||
ACTIVE = False
|
||||
COMFY_PROCESS: subprocess.Popen | None = None
|
||||
|
||||
if len(TOKEN) < 32:
|
||||
raise RuntimeError("WORKER_TOKEN is missing or too short")
|
||||
|
||||
|
||||
def request_json(path: str, payload: dict | None = None,
|
||||
timeout: float = 30) -> dict:
|
||||
body = None if payload is None else json.dumps(payload).encode()
|
||||
headers = {"Content-Type": "application/json"} if body is not None else {}
|
||||
method = "POST" if body is not None else "GET"
|
||||
request = urllib.request.Request(COMFY + path, data=body,
|
||||
headers=headers, method=method)
|
||||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||||
return json.load(response)
|
||||
|
||||
|
||||
def wait_for_comfy() -> None:
|
||||
for _ in range(300):
|
||||
if COMFY_PROCESS is not None and COMFY_PROCESS.poll() is not None:
|
||||
return
|
||||
try:
|
||||
request_json("/system_stats", timeout=2)
|
||||
READY.set()
|
||||
return
|
||||
except Exception:
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
def start_comfy() -> None:
|
||||
global COMFY_PROCESS
|
||||
INPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
command = [
|
||||
"python", "main.py", "--listen", "127.0.0.1", "--port", "8188",
|
||||
"--lowvram", "--preview-method", "none",
|
||||
"--input-directory", str(INPUT_DIR),
|
||||
"--output-directory", str(OUTPUT_DIR),
|
||||
]
|
||||
COMFY_PROCESS = subprocess.Popen(command, cwd=COMFY_DIR)
|
||||
threading.Thread(target=wait_for_comfy, daemon=True).start()
|
||||
|
||||
|
||||
def stop(*_: object) -> None:
|
||||
if COMFY_PROCESS is not None and COMFY_PROCESS.poll() is None:
|
||||
COMFY_PROCESS.terminate()
|
||||
try:
|
||||
COMFY_PROCESS.wait(timeout=15)
|
||||
except subprocess.TimeoutExpired:
|
||||
COMFY_PROCESS.kill()
|
||||
raise SystemExit(0)
|
||||
|
||||
|
||||
def validate_request(data: dict) -> tuple[str, str, int, int, int, float, int]:
|
||||
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 = int(data.get("width", 1024))
|
||||
height = int(data.get("height", 1024))
|
||||
if width % 32 or height % 32 or not 512 <= width <= 2048 or not 512 <= height <= 2048:
|
||||
raise ValueError("width and height must be multiples of 32 between 512 and 2048")
|
||||
steps = int(data.get("steps", 25))
|
||||
guidance = float(data.get("guidance", 1.0))
|
||||
if steps != 25 or guidance != 1.0:
|
||||
raise ValueError("Qwen-Image-2.1 requires steps=25 and guidance=1.0")
|
||||
seed = data.get("seed")
|
||||
seed = random.randrange(2**32) if seed is None else int(seed)
|
||||
if not 0 <= seed <= 2**32 - 1:
|
||||
raise ValueError("invalid seed")
|
||||
return prompt.strip(), filename, width, height, steps, guidance, seed
|
||||
|
||||
|
||||
def make_workflow(prompt: str, seed: int, steps: int, width: int, height: int,
|
||||
references: list[str], prefix: str) -> dict:
|
||||
encode_inputs: dict = {
|
||||
"clip": ["2", 0], "prompt": prompt, "negative_prompt": "",
|
||||
"resolution": max(width, height),
|
||||
}
|
||||
workflow: dict = {
|
||||
"1": {"class_type": "UNETLoader", "inputs": {
|
||||
"unet_name": "qwen_image_2.1_int8_convrot.safetensors",
|
||||
"weight_dtype": "default"}},
|
||||
"2": {"class_type": "CLIPLoader", "inputs": {
|
||||
"clip_name": "qwen3vl_8b_int8_convrot.safetensors",
|
||||
"type": "qwen_image", "device": "default"}},
|
||||
"3": {"class_type": "VAELoader", "inputs": {
|
||||
"vae_name": "qwen_image_2.1_vae_bf16.safetensors"}},
|
||||
"4": {"class_type": "TextEncodeQwenImage21", "inputs": encode_inputs},
|
||||
"6": {"class_type": "KSampler", "inputs": {
|
||||
"model": ["1", 0], "positive": ["4", 0], "negative": ["4", 1],
|
||||
"latent_image": ["4", 2] if references else ["5", 0],
|
||||
"seed": seed, "steps": steps, "cfg": 1.0,
|
||||
"sampler_name": "euler", "scheduler": "simple", "denoise": 1.0}},
|
||||
"7": {"class_type": "VAEDecode", "inputs": {
|
||||
"samples": ["6", 0], "vae": ["3", 0]}},
|
||||
"8": {"class_type": "SaveImage", "inputs": {
|
||||
"filename_prefix": prefix, "images": ["7", 0]}},
|
||||
}
|
||||
if references:
|
||||
encode_inputs["vae"] = ["3", 0]
|
||||
for index, name in enumerate(references, 1):
|
||||
node = str(8 + index)
|
||||
workflow[node] = {"class_type": "LoadImage", "inputs": {"image": name}}
|
||||
encode_inputs[f"image_{index}"] = [node, 0]
|
||||
else:
|
||||
workflow["5"] = {"class_type": "EmptyLatentImage", "inputs": {
|
||||
"width": width, "height": height, "batch_size": 1}}
|
||||
return workflow
|
||||
|
||||
|
||||
def wait_result(prompt_id: str, timeout: int = 1800) -> dict:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
history = request_json(f"/history/{prompt_id}", timeout=10)
|
||||
if prompt_id in history:
|
||||
result = history[prompt_id]
|
||||
status = result.get("status", {})
|
||||
if status.get("status_str") == "error" or not status.get("completed", True):
|
||||
raise RuntimeError("ComfyUI generation failed: " + json.dumps(status))
|
||||
return result
|
||||
time.sleep(1)
|
||||
raise TimeoutError("Qwen image generation timed out")
|
||||
|
||||
|
||||
def prepare_references(data: dict, job: str) -> list[str]:
|
||||
source_files = data.get("source_files") or []
|
||||
if not isinstance(source_files, list) or len(source_files) > 4:
|
||||
raise ValueError("invalid source image list")
|
||||
copied: list[str] = []
|
||||
for index, name in enumerate(source_files, 1):
|
||||
if not isinstance(name, str) or Path(name).name != name:
|
||||
raise ValueError("invalid source image filename")
|
||||
source = (OUTPUT_DIR / name).resolve()
|
||||
if source.parent != OUTPUT_DIR or not source.is_file():
|
||||
raise ValueError("source image not found")
|
||||
with source.open("rb") as stream:
|
||||
header = stream.read(16)
|
||||
if header.startswith(b"\x89PNG\r\n\x1a\n"):
|
||||
suffix = ".png"
|
||||
elif header.startswith(b"\xff\xd8\xff"):
|
||||
suffix = ".jpg"
|
||||
elif header.startswith((b"RIFF",)) and header[8:12] == b"WEBP":
|
||||
suffix = ".webp"
|
||||
else:
|
||||
raise ValueError("unsupported source image format")
|
||||
target_name = f"{job}-ref-{index}{suffix}"
|
||||
shutil.copyfile(source, INPUT_DIR / target_name)
|
||||
copied.append(target_name)
|
||||
return copied
|
||||
|
||||
|
||||
def generate(data: dict) -> dict:
|
||||
global ACTIVE
|
||||
with GENERATION_LOCK:
|
||||
ACTIVE = True
|
||||
started = time.monotonic()
|
||||
copied: list[str] = []
|
||||
try:
|
||||
prompt, filename, width, height, steps, _, seed = validate_request(data)
|
||||
job = "router-" + uuid.uuid4().hex
|
||||
copied = prepare_references(data, job)
|
||||
queued = request_json("/prompt", {"prompt": make_workflow(
|
||||
prompt, seed, steps, width, height, copied, job),
|
||||
"client_id": "athena-image-router"})
|
||||
result = wait_result(queued["prompt_id"])
|
||||
images = [image for output in result.get("outputs", {}).values()
|
||||
for image in output.get("images", [])]
|
||||
if not images:
|
||||
raise RuntimeError("generation completed without an image")
|
||||
image = images[0]
|
||||
generated = (OUTPUT_DIR / image.get("subfolder", "") /
|
||||
image["filename"]).resolve()
|
||||
if OUTPUT_DIR not in generated.parents or not generated.is_file():
|
||||
raise RuntimeError("ComfyUI returned an invalid output path")
|
||||
target = (OUTPUT_DIR / filename).resolve()
|
||||
if target.parent != OUTPUT_DIR:
|
||||
raise RuntimeError("invalid output path")
|
||||
generated.replace(target)
|
||||
return {"status": "ok", "filename": filename, "seed": seed,
|
||||
"seconds": round(time.monotonic() - started, 3),
|
||||
"model": MODEL}
|
||||
finally:
|
||||
for name in copied:
|
||||
try:
|
||||
(INPUT_DIR / name).unlink()
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
ACTIVE = False
|
||||
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, fmt: str, *args: object) -> None:
|
||||
print(f"[qwen-image-2.1] {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(404, {"error": "not found"})
|
||||
elif not READY.is_set():
|
||||
self.reply(503, {"status": "starting", "model": MODEL})
|
||||
else:
|
||||
self.reply(200, {"status": "ok", "model_loaded": ACTIVE,
|
||||
"model": MODEL})
|
||||
|
||||
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 > 32768:
|
||||
raise ValueError("invalid request size")
|
||||
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
||||
except Exception as exc:
|
||||
print(f"[qwen-image-2.1] generation failed: {type(exc).__name__}: "
|
||||
f"{str(exc)[:1000]}", flush=True)
|
||||
self.reply(500, {"status": "error", "message": str(exc)})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
signal.signal(signal.SIGTERM, stop)
|
||||
signal.signal(signal.SIGINT, stop)
|
||||
start_comfy()
|
||||
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|
||||
Reference in New Issue
Block a user