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
|
from diffusers import DiffusionPipeline
|
||||||
started = time.monotonic()
|
started = time.monotonic()
|
||||||
PIPE = DiffusionPipeline.from_pretrained(
|
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
|
LOAD_SECONDS = time.monotonic() - started
|
||||||
|
|
||||||
|
|
||||||
@@ -98,6 +106,8 @@ class Handler(BaseHTTPRequestHandler):
|
|||||||
raise ValueError("invalid request size")
|
raise ValueError("invalid request size")
|
||||||
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
||||||
except Exception as exc:
|
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)})
|
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):
|
if status not in (204, 304):
|
||||||
raise RuntimeError(f"failed to start image worker: HTTP {status}")
|
raise RuntimeError(f"failed to start image worker: HTTP {status}")
|
||||||
else:
|
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"}
|
return {"image_worker": "running" if running else "stopped"}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -719,8 +719,17 @@ class _RemoteWorker:
|
|||||||
headers["Content-Type"] = "application/json"
|
headers["Content-Type"] = "application/json"
|
||||||
req = urllib.request.Request(IMAGE_WORKER_URL + path, data=body,
|
req = urllib.request.Request(IMAGE_WORKER_URL + path, data=body,
|
||||||
method=method, headers=headers)
|
method=method, headers=headers)
|
||||||
with urllib.request.urlopen(req, timeout=timeout) as response:
|
try:
|
||||||
return json.load(response)
|
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:
|
def start(self) -> None:
|
||||||
if not IMAGE_WORKER_TOKEN or len(IMAGE_WORKER_TOKEN) < 32:
|
if not IMAGE_WORKER_TOKEN or len(IMAGE_WORKER_TOKEN) < 32:
|
||||||
|
|||||||
Reference in New Issue
Block a user