Fit FLUX hot swap within 16 GiB VRAM

This commit is contained in:
Mikei386
2026-08-22 18:10:07 +02:00
parent 364ac6b83e
commit b00e882063
3 changed files with 25 additions and 4 deletions
+11 -1
View File
@@ -30,7 +30,15 @@ def load_pipeline() -> None:
from diffusers import DiffusionPipeline
started = time.monotonic()
PIPE = DiffusionPipeline.from_pretrained(
MODEL_DIR, torch_dtype=torch.bfloat16, device_map="cuda")
MODEL_DIR, torch_dtype=torch.bfloat16)
# The full pipeline leaves too little activation headroom on a 16 GiB
# RTX 5080. Model CPU offload keeps each active component on CUDA while
# parking inactive components in system RAM between the four steps.
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
@@ -98,6 +106,8 @@ class Handler(BaseHTTPRequestHandler):
raise ValueError("invalid request size")
self.reply(200, generate(json.loads(self.rfile.read(length))))
except Exception as exc:
print(f"[flux-worker] generation failed: "
f"{type(exc).__name__}: {str(exc)[:1000]}", flush=True)
self.reply(400, {"status": "error", "message": str(exc)})
@@ -106,7 +106,9 @@ def set_image_worker(running: bool) -> dict:
if status not in (204, 304):
raise RuntimeError(f"failed to start image worker: HTTP {status}")
else:
stop_container(item)
# CUDA/PyTorch may not react promptly to SIGTERM after an OOM.
# Bound recovery time and let Docker issue SIGKILL afterwards.
stop_container(item, timeout=20)
return {"image_worker": "running" if running else "stopped"}
+11 -2
View File
@@ -719,8 +719,17 @@ class _RemoteWorker:
headers["Content-Type"] = "application/json"
req = urllib.request.Request(IMAGE_WORKER_URL + path, data=body,
method=method, headers=headers)
with urllib.request.urlopen(req, timeout=timeout) as response:
return json.load(response)
try:
with urllib.request.urlopen(req, timeout=timeout) as response:
return json.load(response)
except urllib.error.HTTPError as exc:
try:
message = json.loads(exc.read(4096)).get("message")
except Exception:
message = None
raise RuntimeError(message or f"Bild-Worker HTTP {exc.code}") from exc
except (OSError, urllib.error.URLError, TimeoutError) as exc:
raise RuntimeError(f"Bild-Worker nicht erreichbar: {exc}") from exc
def start(self) -> None:
if not IMAGE_WORKER_TOKEN or len(IMAGE_WORKER_TOKEN) < 32: