Fit FLUX hot swap within 16 GiB VRAM
This commit is contained in:
@@ -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"}
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user