diff --git a/vram_arbitrator.py b/vram_arbitrator.py
index 9610080..968e771 100644
--- a/vram_arbitrator.py
+++ b/vram_arbitrator.py
@@ -44,13 +44,23 @@ SWITCH_HISTORY = deque(maxlen=50)
RAM_HIT_GBPS = 2.0
PARTIAL_HIT_GBPS = 0.8
-# How long Ollama's VRAM may take to actually drain before we stop waiting.
-# Ollama will not unload a model while a generation is in flight, so a short ceiling
-# reports a timeout for what is really just a busy model finishing its request. Observed
-# here: two yields hit the old 3 s limit with 8.2 GB still held while ComfyUI was starting.
-# Waiting longer is the safer failure mode -- the alternative is diffusion allocating into
-# VRAM that is still occupied.
-YIELD_CONFIRM_TIMEOUT_S = 10.0
+# How long to wait for Ollama's VRAM to actually drain.
+#
+# Ollama will not unload a model mid-generation. With OLLAMA_NUM_PARALLEL=1 our
+# keep_alive:0 request queues behind the running one and takes effect the moment it
+# finishes, so a model that is busy is not failing -- it is finishing, and it will
+# release on its own. Blocking the arbitrator for the length of someone's inference
+# helps nobody: ComfyUI is not gated on our return value, and every second spent
+# blocked is a second the watchdog and profile switching are stalled.
+#
+# So: wait briefly for the common case (an idle model releases in 40-110 ms here),
+# then classify. A caller who genuinely wants to block can ask for a longer wait.
+YIELD_CONFIRM_TIMEOUT_S = 2.0
+YIELD_CONFIRM_TIMEOUT_BLOCKING_S = 30.0
+
+# A model still holding VRAM while the GPU is pinned is generating, not wedged.
+BUSY_UTIL_PCT = 50
+BUSY_PROBE_S = 0.6
YIELD_CONFIRM_POLL_S = 0.02
YIELD_RESIDUAL_BYTES = 256 * 1024 ** 2 # treat <256 MB as "released"
@@ -105,12 +115,17 @@ def get_process_vram_bytes() -> Dict[str, int]:
Deliberately avoids psutil lookups: this runs every 20 ms while we wait for VRAM
to actually drain.
"""
- out = {"ollama_bytes": 0, "comfyui_bytes": 0, "other_bytes": 0, "free_bytes": 0}
+ out = {"ollama_bytes": 0, "comfyui_bytes": 0, "other_bytes": 0, "free_bytes": 0,
+ "gpu_util_pct": 0}
if not NVML_AVAILABLE:
return out
try:
handle = pynvml.nvmlDeviceGetHandleByIndex(0)
out["free_bytes"] = pynvml.nvmlDeviceGetMemoryInfo(handle).free
+ try:
+ out["gpu_util_pct"] = pynvml.nvmlDeviceGetUtilizationRates(handle).gpu
+ except Exception:
+ pass
procs = list(pynvml.nvmlDeviceGetComputeRunningProcesses(handle))
try:
procs += list(pynvml.nvmlDeviceGetGraphicsRunningProcesses(handle))
@@ -415,35 +430,106 @@ async def get_comfyui_live_state() -> Dict[str, Any]:
async def _await_vram_release(baseline_bytes: int,
timeout_s: float = YIELD_CONFIRM_TIMEOUT_S) -> Dict[str, Any]:
- """Block until Ollama's VRAM has actually drained, or we give up.
+ """Wait for Ollama's VRAM to drain, distinguishing "busy" from "stuck".
Posting keep_alive:0 only *asks* Ollama to unload; the driver frees the allocation
- some milliseconds later. Returning before that happens is how ComfyUI ends up
- allocating into VRAM that is still occupied, which shows up as a CUDA OOM mid-graph.
+ some milliseconds later, and returning before that happens is how ComfyUI ends up
+ allocating into VRAM that is still occupied.
+
+ But there is a second case the first version of this got wrong. If the model is
+ mid-generation it cannot unload at all, and reporting that as a timeout made a
+ perfectly healthy cron job look like a 95% failure rate. When the VRAM has not
+ moved and the GPU is pinned, the model is working; the queued unload will fire when
+ it finishes. That is `busy`, not a failure.
+
+ Returns an `outcome` of "released", "busy" or "stuck".
"""
t0 = time.perf_counter()
- last = baseline_bytes
+ peak_util = 0
while True:
snap = get_process_vram_bytes()
last = snap["ollama_bytes"]
+ peak_util = max(peak_util, snap.get("gpu_util_pct", 0))
+ elapsed = time.perf_counter() - t0
+
if last <= YIELD_RESIDUAL_BYTES:
return {
+ "outcome": "released",
"confirmed": True,
- "confirm_ms": round((time.perf_counter() - t0) * 1000, 2),
+ "confirm_ms": round(elapsed * 1000, 2),
"residual_bytes": last,
"free_bytes": snap["free_bytes"],
+ "gpu_util_pct": snap.get("gpu_util_pct", 0),
}
- if (time.perf_counter() - t0) >= timeout_s:
+
+ # Unmoved VRAM plus a pinned GPU means a generation is in flight.
+ busy = (elapsed >= BUSY_PROBE_S
+ and last >= baseline_bytes - YIELD_RESIDUAL_BYTES
+ and peak_util >= BUSY_UTIL_PCT)
+
+ if busy or elapsed >= timeout_s:
+ outcome = "busy" if busy else "stuck"
return {
+ "outcome": outcome,
"confirmed": False,
- "confirm_ms": round((time.perf_counter() - t0) * 1000, 2),
+ "confirm_ms": round(elapsed * 1000, 2),
"residual_bytes": last,
"free_bytes": snap["free_bytes"],
- "error": f"Ollama still holding {round(last / (1024**3), 2)} GB after {timeout_s}s",
+ "gpu_util_pct": snap.get("gpu_util_pct", 0),
+ "peak_util_pct": peak_util,
+ "error": (
+ f"Ollama is mid-generation ({peak_util}% GPU, "
+ f"{round(last / (1024**3), 2)} GB held); the queued unload will apply "
+ f"when it finishes"
+ if outcome == "busy" else
+ f"Ollama still holding {round(last / (1024**3), 2)} GB after "
+ f"{timeout_s}s with the GPU idle"
+ ),
}
await asyncio.sleep(YIELD_CONFIRM_POLL_S)
+# Detached tasks need a strong reference or the loop may garbage-collect them mid-flight.
+_DETACHED: set = set()
+
+
+def _spawn_detached(coro) -> None:
+ task = asyncio.ensure_future(coro)
+ _DETACHED.add(task)
+ task.add_done_callback(_DETACHED.discard)
+
+
+async def _confirm_release_later(targets: List[str], baseline_bytes: int,
+ max_wait_s: float = 900.0) -> None:
+ """Watch for a queued unload to land after the in-flight generation finishes.
+
+ Runs detached so the caller is never held for the length of an inference. Logs the
+ eventual release so the event log tells the whole story rather than stopping at
+ "deferred".
+ """
+ t0 = time.perf_counter()
+ while (time.perf_counter() - t0) < max_wait_s:
+ await asyncio.sleep(0.5)
+ snap = get_process_vram_bytes()
+ if snap["ollama_bytes"] <= YIELD_RESIDUAL_BYTES:
+ waited_ms = round((time.perf_counter() - t0) * 1000, 2)
+ _record({
+ "event_type": "Ollama VRAM Yield",
+ "source": ", ".join(targets)[:200],
+ "target": "VRAM 0MB (Kept in RAM)",
+ "duration_ms": waited_ms,
+ "yield_confirm_ms": waited_ms,
+ "cache_status": "RAM-Cached",
+ "detail": "released after the in-flight generation completed",
+ })
+ arbitrator.note_deferred_release(waited_ms)
+ logger.info(f"Deferred VRAM yield completed after {round(waited_ms / 1000, 1)}s "
+ f"({round(snap['free_bytes'] / (1024**3), 2)} GB free)")
+ return
+ logger.warning("Deferred VRAM yield never landed within "
+ f"{max_wait_s}s for {', '.join(targets)}")
+
+
def _record(event: Dict[str, Any]) -> None:
"""Push an event to both the in-memory ring and the durable store."""
event.setdefault("ts", time.time())
@@ -453,7 +539,8 @@ def _record(event: Dict[str, Any]) -> None:
async def instant_free_ollama_vram(model_name: Optional[str] = None,
- confirm: bool = True) -> Dict[str, Any]:
+ confirm: bool = True,
+ timeout_s: Optional[float] = None) -> Dict[str, Any]:
"""Yield Ollama's VRAM and wait for the driver to actually release it.
The returned duration_ms is now the real end-to-end release time, not just how long
@@ -478,31 +565,51 @@ async def instant_free_ollama_vram(model_name: Optional[str] = None,
model_name = targets[0] if len(targets) == 1 else f"{len(targets)} models"
baseline = get_process_vram_bytes()["ollama_bytes"]
try:
- client = _client(OLLAMA_API_BASE, 5.0)
- await asyncio.gather(*[
+ # Generous client timeout: with OLLAMA_NUM_PARALLEL=1 this request queues behind
+ # any running generation, and a short timeout would drop the connection before
+ # Ollama ever processed the unload -- losing it entirely.
+ client = _client(OLLAMA_API_BASE, 120.0)
+ unload_calls = [
client.post("/api/generate", json={"model": t, "keep_alive": 0})
for t in targets
- ], return_exceptions=True)
+ ]
+ # Do not await the queued unloads; a busy model would block us for the length of
+ # its inference. They are fire-and-confirm: the barrier below watches the VRAM.
+ _spawn_detached(asyncio.gather(*unload_calls, return_exceptions=True))
request_ms = round((time.perf_counter() - t0) * 1000, 2)
- barrier = {"confirmed": None, "confirm_ms": 0.0, "residual_bytes": baseline}
+ barrier: Dict[str, Any] = {"outcome": "unconfirmed", "confirmed": None,
+ "confirm_ms": 0.0, "residual_bytes": baseline}
if confirm:
- barrier = await _await_vram_release(baseline)
+ barrier = await _await_vram_release(
+ baseline, timeout_s if timeout_s is not None else YIELD_CONFIRM_TIMEOUT_S)
+
+ if barrier.get("outcome") == "busy":
+ # The unload is queued and will fire when the generation ends. Keep watching
+ # in the background so the release is still logged and the counters stay true,
+ # without holding the caller here for the length of someone's inference.
+ _spawn_detached(_confirm_release_later(targets, baseline))
duration_ms = round((time.perf_counter() - t0) * 1000, 2)
freed_gb = round(max(baseline - barrier.get("residual_bytes", 0), 0) / (1024**3), 2)
+ outcome = barrier.get("outcome", "unconfirmed")
_record({
"event_type": "Ollama VRAM Yield",
"source": ", ".join(targets)[:200],
"target": "VRAM 0MB (Kept in RAM)",
"duration_ms": duration_ms,
"yield_confirm_ms": barrier.get("confirm_ms"),
- "cache_status": "RAM-Cached" if barrier.get("confirmed") else "Yield Timeout",
+ "cache_status": {
+ "released": "RAM-Cached",
+ "busy": "Deferred — LLM generating",
+ "stuck": "Yield Stalled",
+ }.get(outcome, "Yield Unconfirmed"),
"detail": barrier.get("error"),
})
return {
"success": True,
+ "outcome": outcome,
"model": model_name,
"models_unloaded": targets,
"duration_ms": duration_ms,
@@ -512,6 +619,7 @@ async def instant_free_ollama_vram(model_name: Optional[str] = None,
"freed_gb": freed_gb,
"residual_gb": round(barrier.get("residual_bytes", 0) / (1024**3), 2),
"free_vram_gb": round(barrier.get("free_bytes", 0) / (1024**3), 2),
+ "gpu_util_pct": barrier.get("gpu_util_pct"),
"error": barrier.get("error"),
}
except Exception as e:
@@ -683,7 +791,18 @@ class AutoArbitrator:
# benchmark would otherwise trip trigger_comfy_priority, which reapplies the whole
# 'comfy' profile and silently overwrites the clock the sweep is measuring.
self.oc_suspended = False
- self.stats = {"yields": 0, "purges": 0, "yield_timeouts": 0, "deferred_purges": 0}
+ # Per-model backoff. A model that is mid-generation cannot yield, and asking it
+ # again every second just blocks the loop repeatedly for no benefit.
+ self._yield_backoff_until: Dict[str, float] = {}
+ self._yield_busy_streak: Dict[str, int] = {}
+ self.stats = {
+ "yields": 0, # release confirmed
+ "yield_deferred_busy": 0, # model mid-generation; unload queued behind it
+ "yield_stalled": 0, # VRAM held with an idle GPU -- the real failure
+ "deferred_releases": 0, # queued unloads that later landed
+ "purges": 0,
+ "deferred_purges": 0,
+ }
async def start(self):
if self.running:
@@ -708,8 +827,19 @@ class AutoArbitrator:
await close_clients()
logger.info("AutoArbitrator background engine stopped.")
+ # Backoff schedule for a model that keeps reporting busy, in seconds.
+ BUSY_BACKOFF_S = (5.0, 15.0, 30.0, 60.0)
+
+ def note_deferred_release(self, waited_ms: float) -> None:
+ """Called when a queued unload finally lands after a generation finished."""
+ self.stats["deferred_releases"] += 1
+ self._yield_backoff_until.clear()
+ self._yield_busy_streak.clear()
+ self.last_action = (f"VRAM released after the LLM finished "
+ f"({round(waited_ms / 1000, 1)}s) — ComfyUI can proceed")
+
async def trigger_comfy_priority(self, reason: str = "ComfyUI prompt detected"):
- """Yield Ollama's VRAM — and confirm it is gone — before diffusion allocates."""
+ """Yield Ollama's VRAM before diffusion allocates, without fighting a busy model."""
self.comfy_was_active = True
self.comfy_idle_since = None
self._apply_oc_profile("comfy")
@@ -718,21 +848,39 @@ class AutoArbitrator:
return
ollama_state = await get_ollama_live_state()
- if ollama_state.get("active_model_name"):
- model = ollama_state["active_model_name"]
- logger.info(f"⚡ ComfyUI active ({reason}) -> Auto-yielding Ollama model '{model}'...")
- self.last_yield_time = time.time()
- res = await instant_free_ollama_vram(model, confirm=True)
+ model = ollama_state.get("active_model_name")
+ if not model:
+ return
+
+ # Still finishing an inference we already asked to unload: leave it alone.
+ until = self._yield_backoff_until.get(model, 0.0)
+ if now < until:
+ return
+
+ logger.info(f"⚡ ComfyUI active ({reason}) -> Auto-yielding Ollama model '{model}'...")
+ self.last_yield_time = time.time()
+ res = await instant_free_ollama_vram(model, confirm=True)
+ outcome = res.get("outcome")
+
+ if outcome == "released":
self.stats["yields"] += 1
- if res.get("confirmed"):
- self.last_action = (f"Yielded '{model}' for ComfyUI in "
- f"{res.get('confirm_ms')}ms (confirmed {res.get('freed_gb')}GB free)")
- else:
- self.stats["yield_timeouts"] += 1
- self.last_action = (f"⚠ Yield of '{model}' NOT confirmed: "
- f"{res.get('residual_gb')}GB still held")
- logger.warning(f"VRAM yield barrier timed out: {res.get('error')}")
- logger.info(f"Ollama auto-yield completed: {res}")
+ self._yield_backoff_until.pop(model, None)
+ self._yield_busy_streak.pop(model, None)
+ self.last_action = (f"Yielded '{model}' for ComfyUI in "
+ f"{res.get('confirm_ms')}ms ({res.get('freed_gb')}GB freed)")
+ elif outcome == "busy":
+ streak = self._yield_busy_streak.get(model, 0)
+ delay = self.BUSY_BACKOFF_S[min(streak, len(self.BUSY_BACKOFF_S) - 1)]
+ self._yield_busy_streak[model] = streak + 1
+ self._yield_backoff_until[model] = time.time() + delay
+ self.stats["yield_deferred_busy"] += 1
+ self.last_action = (f"'{model}' is mid-generation ({res.get('gpu_util_pct')}% GPU); "
+ f"unload is queued and will apply when it finishes")
+ logger.info(f"Yield deferred: {res.get('error')} — backing off {delay}s")
+ else:
+ self.stats["yield_stalled"] += 1
+ self.last_action = (f"⚠ '{model}' holding {res.get('residual_gb')}GB with an idle GPU")
+ logger.warning(f"VRAM yield stalled: {res.get('error')}")
async def trigger_comfy_completed(self, immediate: bool = False):
"""Mark the end of a generation. The actual purge is deferred unless forced."""
@@ -893,6 +1041,9 @@ class AutoArbitrator:
"idle_purge_after_s": self.COMFY_IDLE_PURGE_S,
"oc_profile": self.oc_profile,
"counters": dict(self.stats),
+ "yield_backoff": {m: round(max(t - time.time(), 0), 1)
+ for m, t in self._yield_backoff_until.items()
+ if t > time.time()},
}