Files
AI-Profile-Router/platform/docker/image-worker/image_worker.py
T

127 lines
4.9 KiB
Python

#!/usr/bin/env python3
"""Private Z-Image-Turbo worker used only during a GPU hot swap."""
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("Z_IMAGE_MODEL_DIR", "/models/Z-Image-Turbo")
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 ZImagePipeline
started = time.monotonic()
PIPE = ZImagePipeline.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()
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
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", 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")
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]
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"[z-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"[z-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()