Add RTX 5080 FLUX hot-swap worker
This commit is contained in:
@@ -0,0 +1,18 @@
|
||||
FROM pytorch/pytorch:2.11.0-cuda12.8-cudnn9-runtime
|
||||
|
||||
ARG DIFFUSERS_VERSION=0.40.0
|
||||
ARG TRANSFORMERS_VERSION=5.15.1
|
||||
ARG ACCELERATE_VERSION=1.14.0
|
||||
ARG HF_HUB_VERSION=1.28.0
|
||||
|
||||
RUN pip install --no-cache-dir \
|
||||
"diffusers==${DIFFUSERS_VERSION}" \
|
||||
"transformers==${TRANSFORMERS_VERSION}" \
|
||||
"accelerate==${ACCELERATE_VERSION}" \
|
||||
"huggingface-hub==${HF_HUB_VERSION}" \
|
||||
sentencepiece protobuf safetensors pillow && \
|
||||
useradd --system --uid 10002 --home /nonexistent --shell /usr/sbin/nologin flux
|
||||
|
||||
COPY flux_worker.py /app/flux_worker.py
|
||||
USER 10002:10002
|
||||
ENTRYPOINT ["python", "/app/flux_worker.py"]
|
||||
@@ -0,0 +1,109 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Private FLUX.2 Klein Distilled worker used only during a GPU hot swap."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
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")
|
||||
|
||||
|
||||
def load_pipeline() -> None:
|
||||
global PIPE, LOAD_SECONDS
|
||||
if PIPE is not None:
|
||||
return
|
||||
import torch
|
||||
from diffusers import DiffusionPipeline
|
||||
started = time.monotonic()
|
||||
PIPE = DiffusionPipeline.from_pretrained(
|
||||
MODEL_DIR, torch_dtype=torch.bfloat16, device_map="cuda")
|
||||
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", 4))
|
||||
guidance = float(data.get("guidance", 1.0))
|
||||
if steps != 4 or guidance != 1.0:
|
||||
raise ValueError("distilled FLUX.2 Klein requires steps=4 and guidance=1.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=4, guidance_scale=1.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"[flux-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:
|
||||
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()
|
||||
Reference in New Issue
Block a user