Add RTX 5080 FLUX hot-swap worker

This commit is contained in:
Mikei386
2026-08-22 17:40:32 +02:00
parent 8267a85a96
commit 7ac93befc4
12 changed files with 398 additions and 41 deletions
+18
View File
@@ -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"]
+109
View File
@@ -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()
@@ -23,6 +23,8 @@ TOKEN_FILE = os.environ.get("CONTROLLER_TOKEN_FILE", "/run/secrets/controller-to
ALLOWED = tuple(x.strip() for x in os.environ.get(
"ALLOWED_PROFILES", "fast,medium,large,ultra,experimental").split(",") if x.strip())
LABEL_KEY = "com.mike-ai.llama-profile"
IMAGE_LABEL_KEY = "com.mike-ai.image-worker"
IMAGE_WORKER = os.environ.get("IMAGE_WORKER", "flux")
LOCK = threading.Lock()
log = logging.getLogger("profile-controller")
@@ -58,6 +60,56 @@ def containers() -> dict[str, dict]:
return result
def labelled_containers(label: str) -> list[dict]:
filters = urllib.parse.quote(json.dumps({"label": [label]}))
status, body = docker_request("GET", f"/containers/json?all=1&filters={filters}")
if status != 200:
raise RuntimeError(f"Docker list failed with HTTP {status}")
return json.loads(body)
def image_container() -> dict:
matches = [item for item in labelled_containers(IMAGE_LABEL_KEY)
if item.get("Labels", {}).get(IMAGE_LABEL_KEY) == IMAGE_WORKER]
if len(matches) != 1:
raise RuntimeError(
f"expected exactly one image worker {IMAGE_WORKER!r}, found {len(matches)}")
return matches[0]
def stop_container(item: dict, timeout: int = 120) -> None:
if item.get("State") != "running":
return
status, _ = docker_request("POST", f"/containers/{item['Id']}/stop?t={timeout}")
if status not in (204, 304):
raise RuntimeError(f"failed to stop container: HTTP {status}")
def stop_inference() -> dict:
with LOCK:
items = containers()
previous = active_profile(items)
for item in items.values():
stop_container(item)
return {"active_profile": None, "previous_profile": previous}
def set_image_worker(running: bool) -> dict:
with LOCK:
item = image_container()
if running:
# A FLUX worker may never overlap a llama profile on the 5080.
for profile_item in containers().values():
stop_container(profile_item)
if item.get("State") != "running":
status, _ = docker_request("POST", f"/containers/{item['Id']}/start")
if status not in (204, 304):
raise RuntimeError(f"failed to start image worker: HTTP {status}")
else:
stop_container(item)
return {"image_worker": "running" if running else "stopped"}
def active_profile(items: dict[str, dict] | None = None) -> str | None:
items = items or containers()
active = [name for name, item in items.items() if item.get("State") == "running"]
@@ -70,6 +122,8 @@ def activate(profile: str) -> dict:
if profile not in ALLOWED:
raise ValueError("profile is not allowlisted")
with LOCK:
# Defensive mutual exclusion even if a caller bypasses the router.
stop_container(image_container())
items = containers()
missing = [name for name in ALLOWED if name not in items]
if missing:
@@ -144,6 +198,20 @@ class Handler(BaseHTTPRequestHandler):
if not self.authenticated():
self.reply(401, {"error": "unauthorized"})
return
if self.path == "/inference/stop":
try:
self.reply(200, stop_inference())
except Exception as exc:
log.exception("stopping inference failed")
self.reply(503, {"error": str(exc)})
return
if self.path in ("/workers/image/start", "/workers/image/stop"):
try:
self.reply(200, set_image_worker(self.path.endswith("/start")))
except Exception as exc:
log.exception("image worker transition failed")
self.reply(503, {"error": str(exc)})
return
prefix, suffix = "/profiles/", "/activate"
if not self.path.startswith(prefix) or not self.path.endswith(suffix):
self.reply(404, {"error": "not found"})