Files
gpu-program-swapper/vram_arbitrator.py
drjones ca97f18be6 Let any tenant declare its GPU profile and event source; fix priority semantics
Two remaining pieces of the two-application coupling are gone.

Overclock profiles were switched by naming 'comfy' and 'ollama' directly, so a third
application could never get tuned clocks. A tenant declares overclock_profile and the
arbitrator applies whichever the highest-priority *working* tenant asks for, falling
back to the idle profile when nothing is running.

The websocket listener parsed ComfyUI's message schema -- status, execution_start,
executing, execution_success -- which tied the fast path to one application. An event
source is now declarative and the messages are not parsed at all: any message means
"look now", and the tenant's own busy probe decides what is true. That gives the same
sub-second reaction to any application that emits anything on state change, with no
knowledge of what it emits.

Generalising this exposed a design error in the priority rule I had introduced.
plan_release excluded candidates ranking above the demander, which broke both
directions in turn. With the LLM at priority 60 and diffusion at 50, ComfyUI could
never reclaim from Ollama -- the premise the whole service is built on, and preserved
until now only by the ComfyUI-specific trigger that was about to be removed. Swapping
the ranks then broke the reverse: a starved Ollama could no longer reclaim from an
idle ComfyUI.

Priority now orders rather than vetoes. Any idle reclaimable tenant is a candidate,
because an idle tenant is not using its VRAM; priority decides who is asked first, and
busy tenants are never interrupted whatever their rank. Diffusion outranks the LLM,
whose weights reload from page cache in seconds. All three cases are pinned by tests,
including that busy work is never interrupted even by a far higher-priority demander.

Tests: 244 (was 242).

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-07 15:23:48 -07:00

1465 lines
65 KiB
Python

"""VRAM Arbitrator and High-Speed Switch Manager for Ollama and ComfyUI."""
import time
import httpx
import psutil
import logging
from typing import Dict, List, Any, Optional
from collections import deque
import asyncio
import json
import websockets
import overclock_manager
import ram_optimizer
import tenants as tenants_mod
import telemetry_store
try:
import pynvml
pynvml.nvmlInit()
NVML_AVAILABLE = True
except Exception as e:
NVML_AVAILABLE = False
logger = logging.getLogger("vram_arbitrator")
OLLAMA_API_BASE = "http://localhost:11434"
COMFY_API_BASE = "http://127.0.0.1:8188"
# Circular buffer for transition events (the durable log lives in telemetry_store)
SWITCH_HISTORY = deque(maxlen=50)
# Bandwidth thresholds for classifying how a model reached VRAM, calibrated by measuring
# the same 12.87 GB model loaded cold and warm on this box (2026-08-28):
#
# 3.1% resident -> 34.3 s -> 0.38 GB/s
# 100% resident -> 4.9 s -> 2.63 GB/s
#
# The first cut at these numbers assumed a page-cache-fed load would approach the bus
# rate and set the cache-hit bar at 5 GB/s. It does not: Ollama's load_duration covers
# host-to-device transfer and model initialisation as well as the file read, so a fully
# resident model still reports ~2.6 GB/s while the page cache itself reads at 6.4 GB/s.
# A 5 GB/s bar could therefore never be met, and every warm load was being reported as
# a partial hit. Thresholds now sit either side of the measured 6.9x separation.
RAM_HIT_GBPS = 2.0
PARTIAL_HIT_GBPS = 0.8
# 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
# Fraction of a model that may sit outside VRAM before we call it starved. A little
# slack absorbs rounding and KV-cache accounting; beyond it, layers are on the CPU.
CPU_OFFLOAD_TOLERANCE = 0.02
# Only intervene when ComfyUI is actually holding enough VRAM to be the cause.
RECLAIM_MIN_COMFY_BYTES = 512 * 1024 ** 2
# Ollama's response when a model will not fit. Which of the two failure modes you get
# depends on configuration: with n_gpu_layers left to Ollama it spills layers to the CPU
# and reports size_vram < size; with n_gpu_layers pinned (99 on this box) it refuses and
# returns a hard CUDA OOM instead. Both are handled -- the spill by
# AutoArbitrator._arbitrate (generically, from the tenant registry), the hard failure by
# the retry below.
OOM_SIGNATURES = ("out of memory", "cudamalloc", "unable to allocate",
"failed to allocate", "cuda error")
def describe_unmanaged() -> Dict[str, Any]:
"""VRAM held by processes this service cannot reclaim, named explicitly."""
stats = get_gpu_hardware_stats()
bd = stats.get("breakdown", {}) if stats.get("available") else {}
entries = bd.get("unmanaged", [])
return {
"unmanaged_gb": bd.get("unmanaged_gb", 0.0),
"processes": entries,
"note": ("VRAM held by processes outside HyperSwap's control; it cannot be "
"reclaimed automatically" if entries else
"no third-party GPU processes are holding VRAM"),
}
def looks_like_vram_oom(text: str) -> bool:
low = (text or "").lower()
return any(sig in low for sig in OOM_SIGNATURES)
YIELD_CONFIRM_POLL_S = 0.02
YIELD_RESIDUAL_BYTES = 256 * 1024 ** 2 # treat <256 MB as "released"
# Connection-pooled clients. Re-creating an AsyncClient per call meant a fresh TCP
# handshake on every one of the watchdog's polls.
_clients: Dict[str, httpx.AsyncClient] = {}
def _client(base_url: str, timeout: float) -> httpx.AsyncClient:
key = f"{base_url}|{timeout}"
c = _clients.get(key)
if c is None or c.is_closed:
c = httpx.AsyncClient(
base_url=base_url,
timeout=timeout,
limits=httpx.Limits(max_keepalive_connections=4, max_connections=8),
)
_clients[key] = c
return c
async def close_clients() -> None:
for c in list(_clients.values()):
try:
await c.aclose()
except Exception:
pass
_clients.clear()
# Bit flags from nvmlDeviceGetCurrentClocksThrottleReasons, decoded for the governor.
THROTTLE_REASONS = {
0x0000000000000001: "gpu_idle",
0x0000000000000002: "applications_clocks_setting",
0x0000000000000004: "sw_power_cap",
0x0000000000000008: "hw_slowdown",
0x0000000000000010: "sync_boost",
0x0000000000000020: "sw_thermal_slowdown",
0x0000000000000040: "hw_thermal_slowdown",
0x0000000000000080: "hw_power_brake_slowdown",
0x0000000000000100: "display_clock_setting",
}
def decode_throttle_reasons(bits: int) -> List[str]:
return [name for mask, name in THROTTLE_REASONS.items() if bits & mask]
def get_process_vram_bytes() -> Dict[str, int]:
"""Fast NVML-only VRAM attribution, used by the yield barrier's tight poll loop.
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,
"desktop_bytes": 0, "unmanaged_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))
except Exception:
pass
merged: Dict[int, int] = {}
for p in procs:
merged[p.pid] = max(merged.get(p.pid, 0), p.usedGpuMemory or 0)
for pid, used in merged.items():
key = _pid_key(pid)
kind = _PID_KIND_CACHE.get(key) if key else None
if kind is None:
kind = _classify_pid(pid)
if key:
if len(_PID_KIND_CACHE) >= _PID_KIND_CACHE_MAX:
_PID_KIND_CACHE.clear()
_PID_KIND_CACHE[key] = kind
if kind == "ollama":
out["ollama_bytes"] += used
elif kind == "comfy":
out["comfyui_bytes"] += used
elif kind == "desktop":
out["desktop_bytes"] += used
out["other_bytes"] += used
else:
out["unmanaged_bytes"] += used
out["other_bytes"] += used
except Exception as e:
logger.debug(f"get_process_vram_bytes failed: {e}")
return out
# Keyed by (pid, process start time) rather than pid alone. Linux recycles PIDs, and a
# stale entry would attribute a new process's VRAM to Ollama or ComfyUI -- in the same
# snapshot the yield barrier uses to decide whether VRAM was released.
_PID_KIND_CACHE: Dict[tuple, str] = {}
_PID_KIND_CACHE_MAX = 512
def _pid_key(pid: int) -> Optional[tuple]:
try:
return (pid, psutil.Process(pid).create_time())
except Exception:
return None
# Tenant names as used by this module's buckets. The tenant registry is the source of
# truth for *which* application a process belongs to; these two names are kept because
# the REST payloads and the dashboard have used them since the beginning.
_BUCKET_ALIASES = {"comfyui": "comfy"}
def _pid_key(pid: int) -> Optional[tuple]:
try:
return (pid, psutil.Process(pid).create_time())
except Exception:
return None
# Compositors and display servers. Their VRAM is small, permanent and not ours to
# reclaim, so it should not be confused with a real workload.
DESKTOP_PROCESS_HINTS = (
"gnome-shell", "xorg", "gnome-remote-desktop", "mutter", "kwin", "plasmashell",
"gnome-session", "wayland", "weston", "sddm", "gdm", "picom", "compiz",
)
def _classify_pid(pid: int) -> str:
"""Which tenant owns this GPU process.
The matching rules used to be substrings compiled into this function, which made the
two applications on this box part of the arbitrator rather than input to it. They now
come from the tenant registry, so a third application is a config entry.
"unmanaged" still means something specific and useful: VRAM held by something with no
declared way to release it, and therefore headroom this service can never offer.
"""
return _BUCKET_ALIASES.get(tenants_mod.classify_pid(pid), tenants_mod.classify_pid(pid))
def get_gpu_hardware_stats() -> Dict[str, Any]:
"""Retrieve comprehensive GPU hardware and process metrics via NVML."""
if not NVML_AVAILABLE:
return {"available": False, "error": "NVML not initialized"}
try:
handle = pynvml.nvmlDeviceGetHandleByIndex(0)
name = pynvml.nvmlDeviceGetName(handle)
if isinstance(name, bytes):
name = name.decode("utf-8")
mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle)
util_rates = pynvml.nvmlDeviceGetUtilizationRates(handle)
temp_c = pynvml.nvmlDeviceGetTemperature(handle, pynvml.NVML_TEMPERATURE_GPU)
try:
power_mw = pynvml.nvmlDeviceGetPowerUsage(handle)
power_w = round(power_mw / 1000.0, 1)
except Exception:
power_w = 0.0
fan_pct = 0
fans = []
try:
num_fans = pynvml.nvmlDeviceGetNumFans(handle)
for i in range(num_fans):
try:
fans.append(pynvml.nvmlDeviceGetFanSpeed_v2(handle, i))
except Exception:
pass
if fans:
fan_pct = max(fans)
else:
fan_pct = pynvml.nvmlDeviceGetFanSpeed(handle)
fans = [fan_pct]
except Exception:
try:
fan_pct = pynvml.nvmlDeviceGetFanSpeed(handle)
fans = [fan_pct]
except Exception:
fan_pct = 0
fans = []
try:
clock_graphics = pynvml.nvmlDeviceGetClockInfo(handle, pynvml.NVML_CLOCK_GRAPHICS)
clock_mem = pynvml.nvmlDeviceGetClockInfo(handle, pynvml.NVML_CLOCK_MEM)
except Exception:
clock_graphics = 0
clock_mem = 0
# PCIe throughput (KB/s) — TX + RX. Key metric for the RAM-cache
# PCIe-speed swap thesis (assimilated from pmady/gpu-mcp-server).
pcie_tx_kbps = 0
pcie_rx_kbps = 0
try:
pcie_tx_kbps = pynvml.nvmlDeviceGetPcieThroughput(handle, pynvml.NVML_PCIE_UTIL_TX_BYTES)
except Exception:
pass
try:
pcie_rx_kbps = pynvml.nvmlDeviceGetPcieThroughput(handle, pynvml.NVML_PCIE_UTIL_RX_BYTES)
except Exception:
pass
# Why the GPU is not running at full clocks — the thermal governor reads this.
throttle_bits = 0
throttle_reasons: List[str] = []
try:
throttle_bits = pynvml.nvmlDeviceGetCurrentClocksThrottleReasons(handle)
throttle_reasons = decode_throttle_reasons(throttle_bits)
except Exception:
pass
# Power management limit (watts) — the OC ceiling.
power_limit_w = 0.0
try:
power_limit_w = round(pynvml.nvmlDeviceGetPowerManagementLimit(handle) / 1000.0, 1)
except Exception:
pass
# Driver + CUDA version (completeness).
driver_version = ""
cuda_version = ""
try:
dv = pynvml.nvmlSystemGetDriverVersion()
driver_version = dv.decode("utf-8") if isinstance(dv, bytes) else str(dv)
except Exception:
pass
try:
cv = pynvml.nvmlSystemGetCudaDriverVersion()
cuda_version = cv # int like 12030 == CUDA 12.3
except Exception:
pass
# Discover processes on GPU
proc_breakdown = {
"ollama_bytes": 0,
"comfyui_bytes": 0,
"system_bytes": 0,
"desktop_bytes": 0,
"unmanaged_bytes": 0,
"unmanaged": [],
"processes": [],
# Generic attribution: one entry per tenant, so an application added to the
# registry is reported without any change here.
"by_tenant": {},
}
try:
procs = pynvml.nvmlDeviceGetComputeRunningProcesses(handle)
graphics_procs = pynvml.nvmlDeviceGetGraphicsRunningProcesses(handle)
all_procs = {p.pid: p.usedGpuMemory for p in procs}
for p in graphics_procs:
all_procs[p.pid] = max(all_procs.get(p.pid, 0), p.usedGpuMemory or 0)
for pid, used_mem in all_procs.items():
pname = "Unknown"
cmdline = ""
try:
proc = psutil.Process(pid)
pname = proc.name()
cmdline = " ".join(proc.cmdline())
except Exception:
pass
kind = _classify_pid(pid)
is_ollama = kind == "ollama"
is_comfy = kind == "comfy"
if is_ollama:
proc_breakdown["ollama_bytes"] += used_mem
elif is_comfy:
proc_breakdown["comfyui_bytes"] += used_mem
else:
proc_breakdown["system_bytes"] += used_mem
if kind == "desktop":
proc_breakdown["desktop_bytes"] += used_mem
else:
proc_breakdown["unmanaged_bytes"] += used_mem
proc_breakdown["unmanaged"].append({
"pid": pid, "name": pname,
"cmdline": cmdline[:120],
"vram_mb": round(used_mem / (1024**2), 1),
})
proc_breakdown["by_tenant"][kind] = (
proc_breakdown["by_tenant"].get(kind, 0) + used_mem)
proc_breakdown["processes"].append({
"pid": pid,
"name": pname,
"cmdline": cmdline[:60],
"vram_bytes": used_mem,
"vram_mb": round(used_mem / (1024**2), 1),
"is_ollama": is_ollama,
"is_comfy": is_comfy,
"kind": kind,
})
except Exception as e:
logger.error(f"Error enumerating GPU processes: {e}")
total_vram = mem_info.total
used_vram = mem_info.used
free_vram = mem_info.free
return {
"available": True,
"device_name": name,
"vram_total_bytes": total_vram,
"vram_total_gb": round(total_vram / (1024**3), 2),
"vram_used_bytes": used_vram,
"vram_used_gb": round(used_vram / (1024**3), 2),
"vram_free_bytes": free_vram,
"vram_free_gb": round(free_vram / (1024**3), 2),
"vram_used_pct": round((used_vram / total_vram * 100) if total_vram > 0 else 0, 1),
"gpu_util_pct": util_rates.gpu,
"mem_util_pct": util_rates.memory,
"temperature_c": temp_c,
"power_w": power_w,
"power_limit_w": power_limit_w,
"pcie_tx_kbps": pcie_tx_kbps,
"pcie_rx_kbps": pcie_rx_kbps,
"throttle_bits": throttle_bits,
"throttle_reasons": throttle_reasons,
"driver_version": driver_version,
"cuda_version": cuda_version,
"fan_pct": fan_pct,
"fans": fans,
"num_fans": len(fans),
"clock_graphics_mhz": clock_graphics,
"clock_mem_mhz": clock_mem,
"breakdown": {
"ollama_mb": round(proc_breakdown["ollama_bytes"] / (1024**2), 1),
"ollama_gb": round(proc_breakdown["ollama_bytes"] / (1024**3), 2),
"comfyui_mb": round(proc_breakdown["comfyui_bytes"] / (1024**2), 1),
"comfyui_gb": round(proc_breakdown["comfyui_bytes"] / (1024**3), 2),
"system_mb": round(proc_breakdown["system_bytes"] / (1024**2), 1),
"system_gb": round(proc_breakdown["system_bytes"] / (1024**3), 2),
"desktop_gb": round(proc_breakdown["desktop_bytes"] / (1024**3), 2),
# VRAM held by workloads this service has no control over. It cannot be
# reclaimed, so it is permanently unavailable headroom.
"unmanaged_gb": round(proc_breakdown["unmanaged_bytes"] / (1024**3), 2),
"unmanaged": proc_breakdown["unmanaged"],
"by_tenant_gb": {k: round(b / (1024**3), 2)
for k, b in proc_breakdown["by_tenant"].items()},
"free_mb": round(free_vram / (1024**2), 1),
"free_gb": round(free_vram / (1024**3), 2),
"processes": proc_breakdown["processes"],
}
}
except Exception as e:
return {"available": False, "error": str(e)}
async def get_ollama_live_state() -> Dict[str, Any]:
"""Get active models, running status, and VRAM expiration from Ollama."""
state = {
"online": False,
"loaded_models": [],
"active_model_name": None,
"active_model_vram_gb": 0.0,
"active_context": 0,
"expires_at": None,
"installed_models": [],
# Ollama silently spills layers to CPU when VRAM is short. size_vram < size is the
# only externally visible sign, and the cost is roughly an order of magnitude in
# decode speed, so it is worth surfacing loudly.
"gpu_fraction": 1.0,
"cpu_offload_pct": 0.0,
"partially_offloaded": False,
}
try:
client = _client(OLLAMA_API_BASE, 3.0)
# Check running models (ps)
ps_resp = await client.get("/api/ps")
if ps_resp.status_code == 200:
state["online"] = True
models = ps_resp.json().get("models", [])
state["loaded_models"] = models
if models:
first = models[0]
state["active_model_name"] = first.get("name")
vram_bytes = first.get("size_vram", first.get("size", 0))
state["active_model_vram_gb"] = round(vram_bytes / (1024**3), 2)
state["active_context"] = first.get("context_length", 0)
state["expires_at"] = first.get("expires_at")
# Check all tags
tags_resp = await client.get("/api/tags")
if tags_resp.status_code == 200:
state["installed_models"] = tags_resp.json().get("models", [])
except Exception as e:
logger.debug(f"Ollama check error: {e}")
return state
async def get_comfyui_live_state() -> Dict[str, Any]:
"""Get prompt queue, device status, and active execution from ComfyUI."""
state = {
"online": False,
"executing": False,
"queue_remaining": 0,
"queue_running": 0,
"current_node": None,
"current_prompt_id": None,
"vram_free_mb": 0,
"vram_total_mb": 0,
}
try:
client = _client(COMFY_API_BASE, 3.0)
# Check system stats
stats_resp = await client.get("/system_stats")
if stats_resp.status_code == 200:
state["online"] = True
data = stats_resp.json()
devices = data.get("devices", [])
if devices:
dev = devices[0]
state["vram_free_mb"] = round(dev.get("vram_free", 0) / (1024**2), 1)
state["vram_total_mb"] = round(dev.get("vram_total", 0) / (1024**2), 1)
# Check queue
queue_resp = await client.get("/queue")
if queue_resp.status_code == 200:
qdata = queue_resp.json()
running = qdata.get("queue_running", [])
pending = qdata.get("queue_pending", [])
state["queue_running"] = len(running)
state["queue_remaining"] = len(pending)
state["executing"] = len(running) > 0
if running:
state["current_prompt_id"] = running[0][1] if len(running[0]) > 1 else str(running[0])
except Exception as e:
logger.debug(f"ComfyUI check error: {e}")
return state
async def _await_vram_release(baseline_bytes: int,
timeout_s: float = YIELD_CONFIRM_TIMEOUT_S) -> Dict[str, Any]:
"""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, 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()
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(elapsed * 1000, 2),
"residual_bytes": last,
"free_bytes": snap["free_bytes"],
"gpu_util_pct": snap.get("gpu_util_pct", 0),
}
# 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(elapsed * 1000, 2),
"residual_bytes": last,
"free_bytes": snap["free_bytes"],
"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())
event.setdefault("timestamp", time.strftime("%H:%M:%S"))
SWITCH_HISTORY.appendleft(event)
telemetry_store.record_event(event, profile=overclock_manager.ACTIVE_PROFILE)
async def instant_free_ollama_vram(model_name: Optional[str] = None,
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
the HTTP POST took.
"""
t0 = time.perf_counter()
if model_name:
targets = [model_name]
else:
# Unload *every* resident model, not just loaded_models[0]. Ollama will happily
# keep several models in VRAM at once; releasing only the first left the rest
# allocated, which the confirm barrier caught as "still holding 8.2 GB after 3s".
ollama_state = await get_ollama_live_state()
targets = [m.get("name") for m in ollama_state.get("loaded_models", []) if m.get("name")]
if not targets and ollama_state.get("active_model_name"):
targets = [ollama_state["active_model_name"]]
if not targets:
return {"success": True, "message": "No active Ollama model in VRAM",
"duration_ms": 0, "confirmed": True}
model_name = targets[0] if len(targets) == 1 else f"{len(targets)} models"
baseline = get_process_vram_bytes()["ollama_bytes"]
try:
# 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
]
# 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: Dict[str, Any] = {"outcome": "unconfirmed", "confirmed": None,
"confirm_ms": 0.0, "residual_bytes": baseline}
if confirm:
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": {
"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,
"request_ms": request_ms,
"confirm_ms": barrier.get("confirm_ms"),
"confirmed": barrier.get("confirmed"),
"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:
return {"success": False, "error": str(e),
"duration_ms": round((time.perf_counter() - t0) * 1000, 2)}
async def instant_free_comfyui_vram() -> Dict[str, Any]:
"""Tell ComfyUI to purge loaded diffusion models from VRAM."""
t0 = time.perf_counter()
try:
client = _client(COMFY_API_BASE, 5.0)
await client.post("/free", json={"unload_models": True, "free_memory": True})
duration_ms = round((time.perf_counter() - t0) * 1000, 2)
snap = get_process_vram_bytes()
_record({
"event_type": "ComfyUI VRAM Purge",
"source": "ComfyUI Pipeline",
"target": "VRAM Free",
"duration_ms": duration_ms,
"cache_status": "Cleaned",
})
return {"success": True, "duration_ms": duration_ms,
"free_vram_gb": round(snap["free_bytes"] / (1024**3), 2)}
except Exception as e:
return {"success": False, "error": str(e),
"duration_ms": round((time.perf_counter() - t0) * 1000, 2)}
_MODEL_SIZE_CACHE: Dict[str, int] = {}
def _model_size_bytes(model_name: str) -> int:
"""On-disk weight size for an Ollama model, used to turn load time into bandwidth."""
if model_name in _MODEL_SIZE_CACHE:
return _MODEL_SIZE_CACHE[model_name]
try:
for f in ram_optimizer.find_ollama_model_files():
_MODEL_SIZE_CACHE[f["model"]] = f["size_bytes"]
except Exception as e:
logger.debug(f"model size lookup failed: {e}")
return _MODEL_SIZE_CACHE.get(model_name, 0)
def classify_load(size_bytes: int, load_duration_ms: float) -> Dict[str, Any]:
"""Classify how a model reached VRAM, from achieved bandwidth rather than a constant.
The old rule was `load_duration_ms < 2500`, which called a 27B Q2_K read from NVMe a
cache hit and a small model read from RAM a cold load. Bandwidth separates them
cleanly: page cache feeds PCIe at many GB/s, this NVMe does not.
"""
if load_duration_ms <= 1.0:
return {"cache_status": "Already in VRAM", "load_gbps": None, "is_ram_hit": True}
if not size_bytes:
# No size on record — fall back to the old heuristic, but say so.
# Without a size we cannot compute bandwidth at all; this is a guess and is
# labelled as one. 8s roughly splits the measured warm (4.9s) and cold (34.3s)
# loads for a mid-size model, but it is meaningless for very small or large ones.
return {
"cache_status": "RAM Cache Hit ⚡" if load_duration_ms < 8000 else "Cold Disk Load 💾",
"load_gbps": None,
"is_ram_hit": load_duration_ms < 8000,
"detail": "size unknown, fell back to a duration guess",
}
gbps = (size_bytes / (1024**3)) / (load_duration_ms / 1000.0)
if gbps >= RAM_HIT_GBPS:
status = "RAM Cache Hit ⚡"
elif gbps >= PARTIAL_HIT_GBPS:
status = "Partial Cache 🌤"
else:
status = "Cold Disk Load 💾"
return {"cache_status": status, "load_gbps": round(gbps, 2), "is_ram_hit": gbps >= RAM_HIT_GBPS}
async def switch_ollama_model(target_model: str, keep_alive: str = "30m",
_retrying: bool = False) -> Dict[str, Any]:
"""High-speed hot-swap to target Ollama model, tracking swap metrics.
If the load fails because the model will not fit, reclaims VRAM from an idle ComfyUI
and retries once. `_retrying` guards against recursing more than one level.
"""
t0 = time.perf_counter()
cur_state = await get_ollama_live_state()
prev_model = cur_state.get("active_model_name") or "None"
try:
client = _client(OLLAMA_API_BASE, 180.0)
resp = await client.post(
"/api/generate",
json={"model": target_model, "prompt": "Ready check", "stream": False,
"keep_alive": keep_alive},
)
total_duration_ms = round((time.perf_counter() - t0) * 1000, 2)
if resp.status_code == 200:
data = resp.json()
load_dur_ms = round(data.get("load_duration", 0) / 1e6, 2)
eval_dur_ms = round(data.get("eval_duration", 0) / 1e6, 2)
eval_count = data.get("eval_count", 0)
tokens_per_sec = round((eval_count / (eval_dur_ms / 1000)) if eval_dur_ms > 0 else 0, 1)
size_bytes = _model_size_bytes(target_model)
cls = classify_load(size_bytes, load_dur_ms)
_record({
"event_type": "LLM Model Switch",
"source": prev_model,
"target": target_model,
"duration_ms": total_duration_ms,
"load_duration_ms": load_dur_ms,
"tokens_per_sec": tokens_per_sec,
"bytes_loaded": size_bytes,
"load_gbps": cls["load_gbps"],
"cache_status": cls["cache_status"],
"detail": cls.get("detail"),
})
return {
"success": True,
"prev_model": prev_model,
"target_model": target_model,
"total_duration_ms": total_duration_ms,
"load_duration_ms": load_dur_ms,
"tokens_per_sec": tokens_per_sec,
"model_size_gb": round(size_bytes / (1024**3), 2) if size_bytes else None,
"load_gbps": cls["load_gbps"],
"cache_status": cls["cache_status"],
"is_ram_hit": cls["is_ram_hit"],
"response": data.get("response", ""),
}
# A model that will not fit is the exact contention this service exists to
# resolve. Rather than handing the caller a CUDA OOM, take the VRAM back from an
# idle ComfyUI and try once more.
body = resp.text
if looks_like_vram_oom(body) and not _retrying:
# Which application should give up memory is a question for the registry,
# not something to answer by purging ComfyUI by name. Any reclaimable idle
# tenant below Ollama in priority is a candidate.
state = await arbitrator._tenant_state()
free_gb = arbitrator._last_tenant_state["free_gb"]
size_gb = _model_size_bytes(target_model) / (1024**3)
needed = size_gb * 1.16 if size_gb else free_gb + 1.0
plan = tenants_mod.plan_release("ollama", state, free_gb, needed)
if plan["release"]:
logger.warning(
f"Ollama could not fit '{target_model}' — {plan['reason']}")
freed_before = free_gb
for victim in plan["release"]:
await arbitrator._release_tenant(
victim, f"Ollama could not load '{target_model}'")
arbitrator.stats["reclaims_for_ollama"] += 1
arbitrator.last_action = (
f"Released {', '.join(plan['release'])} so '{target_model}' could load")
_record({
"event_type": "VRAM Reclaim for Ollama",
"source": ", ".join(plan["release"]),
"target": target_model,
"cache_status": "Reclaimed",
"detail": f"Ollama OOM: {body[:160]}",
})
await asyncio.sleep(0.3)
retry = await switch_ollama_model(target_model, keep_alive, _retrying=True)
retry["released_tenants"] = plan["release"]
retry["would_free_gb"] = plan.get("would_free_gb")
retry["first_attempt_error"] = "CUDA OOM; retried after reclaiming VRAM"
if not retry.get("success"):
# Be specific about why the reclaim was not enough. Blaming a tenant
# when a process nobody can release is holding the memory sends the
# user looking in the wrong place.
retry["blockers"] = plan.get("blockers")
retry["unmanaged_blockers"] = describe_unmanaged()
return retry
return {"success": False, "error": f"HTTP {resp.status_code}: {body}",
"duration_ms": total_duration_ms,
"upstream_status": resp.status_code,
"vram_oom": looks_like_vram_oom(body)}
except Exception as e:
return {"success": False, "error": str(e),
"duration_ms": round((time.perf_counter() - t0) * 1000, 2)}
def get_switch_history() -> List[Dict[str, Any]]:
return list(SWITCH_HISTORY)
class AutoArbitrator:
"""Real-time bidirectional background arbitrator for seamless Ollama <-> ComfyUI hot-swapping.
Two behavioural changes worth knowing about:
* ComfyUI's VRAM is no longer purged 1.5 s after every finished prompt. Iterating on
a workflow is the common case, and purging between runs forced a full checkpoint
reload each time. The purge now waits for COMFY_IDLE_PURGE_S of genuinely empty
queue, and happens immediately only when Ollama actually needs the VRAM.
* The watchdog no longer polls two ComfyUI endpoints every 300 ms. The WebSocket is
the primary signal; polling is a fallback that runs at 1 Hz and only hits /queue,
backing off further while the socket is healthy.
"""
COMFY_IDLE_PURGE_S = 30.0
WATCHDOG_INTERVAL_S = 1.0
WATCHDOG_INTERVAL_WS_OK_S = 3.0
def __init__(self):
self.running = False
self.ws_task: Optional[asyncio.Task] = None
self.event_tasks: List[asyncio.Task] = []
self.poll_task: Optional[asyncio.Task] = None
self.idle_task: Optional[asyncio.Task] = None
self.last_yield_time = 0.0
self.last_comfy_free_time = 0.0
self.connected_ws = False
self.last_action = "Idle"
self.comfy_was_active = False
self.comfy_idle_since: Optional[float] = None
self.oc_profile = None
self.pending_purge = False
# While a tuning sweep is running, the arbitrator must not fight it: a ComfyUI
# 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
# 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.last_reclaim_time = 0.0
self.watchdog_branches = {"busy": 0, "completed": 0, "idle_check": 0,
"bad_status": 0, "error": 0}
self._running_id: Optional[str] = None
self._running_since: Optional[float] = None
self._peak_comfy_bytes = 0
self.comfy_stale_job: Optional[str] = None
self._idle_since: Dict[str, float] = {}
self._last_event_wake = 0.0
self.event_sources: Dict[str, str] = {}
self._last_tenant_state: Optional[Dict[str, Any]] = None
self.last_arbitration: Optional[Dict[str, Any]] = None
self.last_watchdog_error: Optional[str] = None
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,
"reclaims_for_ollama": 0, # ComfyUI purged because the LLM was spilling to CPU
}
async def start(self):
if self.running:
return
self.running = True
self.ws_task = asyncio.create_task(self._ws_listener())
for t in tenants_mod.load_tenants():
if t.enabled and t.events.type == "websocket" and t.events.url:
self.event_tasks.append(asyncio.create_task(
self._event_listener(t.name, t.events.url,
t.events.reconnect_backoff_s,
t.events.max_backoff_s)))
self.poll_task = asyncio.create_task(self._poll_watchdog())
self.idle_task = asyncio.create_task(self._idle_purge_loop())
logger.info("AutoArbitrator background engine started (Bidirectional).")
try:
await asyncio.get_running_loop().run_in_executor(
None, overclock_manager.apply_profile, "balanced")
self.oc_profile = "balanced"
except Exception as e:
logger.warning(f"Startup overclock apply failed: {e}")
async def stop(self):
self.running = False
for task in [self.ws_task, self.poll_task, self.idle_task, *self.event_tasks]:
if task:
task.cancel()
self.event_tasks.clear()
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 before diffusion allocates, without fighting a busy model."""
self.comfy_was_active = True
self.comfy_idle_since = None
self._apply_oc_profile("comfy")
now = time.time()
if now - self.last_yield_time < 1.0:
return
ollama_state = await get_ollama_live_state()
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
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."""
self.comfy_was_active = False
if self.comfy_idle_since is None:
self.comfy_idle_since = time.time()
self._apply_oc_profile("ollama")
if immediate:
await self._purge_comfy_now("Ollama needs VRAM")
else:
self.pending_purge = True
self.stats["deferred_purges"] += 1
self.last_action = (f"ComfyUI idle — holding its checkpoints for "
f"{int(self.COMFY_IDLE_PURGE_S)}s in case you iterate")
async def _purge_comfy_now(self, reason: str):
now = time.time()
if now - self.last_comfy_free_time < 3.0:
return
self.last_comfy_free_time = now
self.pending_purge = False
logger.info(f"⚡ Purging ComfyUI VRAM cache ({reason})...")
res = await instant_free_comfyui_vram()
self.stats["purges"] += 1
self.last_action = f"Purged ComfyUI VRAM ({res.get('duration_ms')}ms) — {reason}"
logger.info(f"ComfyUI purge completed: {res}")
async def _idle_purge_loop(self):
"""Purge ComfyUI's VRAM only after a real idle gap, not between iterations."""
while self.running:
try:
if self.pending_purge and self.comfy_idle_since and not self.comfy_was_active:
idle_for = time.time() - self.comfy_idle_since
if idle_for >= self.COMFY_IDLE_PURGE_S:
await self._purge_comfy_now(
f"idle {int(idle_for)}s")
except Exception as e:
logger.debug(f"idle purge loop error: {e}")
await asyncio.sleep(2.0)
async def request_vram_for_ollama(self, needed_gb: float = 0.0) -> Dict[str, Any]:
"""Called when Ollama needs VRAM now: purge ComfyUI immediately rather than waiting."""
snap = get_process_vram_bytes()
free_gb = snap["free_bytes"] / (1024**3)
if needed_gb and free_gb >= needed_gb:
return {"purged": False, "free_gb": round(free_gb, 2), "reason": "enough free VRAM"}
if snap["comfyui_bytes"] > YIELD_RESIDUAL_BYTES:
await self._purge_comfy_now(f"Ollama requested {needed_gb or '?'}GB")
snap = get_process_vram_bytes()
return {"purged": True, "free_gb": round(snap["free_bytes"] / (1024**3), 2)}
return {"purged": False, "free_gb": round(free_gb, 2), "reason": "ComfyUI holds no VRAM"}
async def _event_listener(self, tenant_name: str, url: str,
backoff_s: float, max_backoff_s: float) -> None:
"""Wake on a tenant's event stream instead of waiting for the next poll.
Deliberately does not parse the messages. The previous listener understood
ComfyUI's schema -- status/execution_start/executing/execution_success -- which
tied the fast path to one application. Treating any message as "look now" and
letting the tenant's own busy probe decide gives the same sub-second reaction
for any application that emits anything on state change.
"""
backoff = backoff_s
while self.running:
try:
async with websockets.connect(url, ping_interval=10, ping_timeout=10) as ws:
self.event_sources[tenant_name] = "connected"
self.connected_ws = True
backoff = backoff_s
logger.info(f"Event source connected for '{tenant_name}': {url}")
while self.running:
await ws.recv()
# Coalesce bursts: a single graph emits many messages, and one
# arbitration pass per burst is enough.
now = time.time()
if now - self._last_event_wake < 0.25:
continue
self._last_event_wake = now
self.stats["event_wakeups"] = self.stats.get("event_wakeups", 0) + 1
try:
await self._arbitrate()
except Exception as e:
logger.debug(f"arbitration from event failed: {e}")
except (websockets.exceptions.ConnectionClosed, OSError, asyncio.CancelledError):
self.event_sources[tenant_name] = "disconnected"
self.connected_ws = False
except Exception as e:
self.event_sources[tenant_name] = f"error: {str(e)[:60]}"
self.connected_ws = False
logger.debug(f"event source error for '{tenant_name}': {e}")
await asyncio.sleep(backoff)
backoff = min(backoff * 1.5, max_backoff_s)
async def _ws_listener(self):
client_id = "hyperswap-arbitrator"
ws_url = f"ws://127.0.0.1:8188/ws?clientId={client_id}"
backoff = 2.0
while self.running:
try:
async with websockets.connect(ws_url, ping_interval=10, ping_timeout=10) as ws:
self.connected_ws = True
backoff = 2.0
logger.info("AutoArbitrator connected to ComfyUI WebSocket.")
while self.running:
msg = await ws.recv()
if not isinstance(msg, str):
continue
try:
data = json.loads(msg)
msg_type = data.get("type")
msg_data = data.get("data", {})
if msg_type == "status":
queue_rem = (msg_data.get("status", {})
.get("exec_info", {}).get("queue_remaining", 0))
if queue_rem > 0:
await self.trigger_comfy_priority(f"Queue remaining: {queue_rem}")
elif queue_rem == 0 and self.comfy_was_active:
await self.trigger_comfy_completed()
elif msg_type in ("execution_start", "execution_cached"):
await self.trigger_comfy_priority(f"Event: {msg_type}")
elif msg_type == "executing":
node = msg_data.get("node")
if node is not None:
await self.trigger_comfy_priority(f"Executing node: {node}")
elif self.comfy_was_active:
await self.trigger_comfy_completed()
elif msg_type == "execution_success":
await self.trigger_comfy_completed()
elif msg_type == "execution_error":
logger.warning(f"ComfyUI execution error: {msg_data}")
await self.trigger_comfy_completed()
except Exception as e:
logger.debug(f"WS parse error: {e}")
except (websockets.exceptions.ConnectionClosed, OSError, asyncio.CancelledError):
self.connected_ws = False
except Exception as e:
self.connected_ws = False
logger.debug(f"WS connection error: {e}")
await asyncio.sleep(backoff)
backoff = min(backoff * 1.5, 15.0)
RECLAIM_COOLDOWN_S = 30.0
# A queue entry that has claimed to be running this long without the GPU ever going
# busy is stale, not slow.
STALE_RUNNING_S = 90.0
# ComfyUI's own VRAM, not GPU utilisation, is what distinguishes a real job from a
# stale row. Utilisation is shared: Ollama and any third-party process drive it too,
# so peak utilisation stayed above any sensible threshold and a stuck entry never
# looked stale. A real diffusion job loads gigabytes of checkpoint; a dead one holds
# only the CUDA context.
STALE_COMFY_BYTES = 1.5 * 1024 ** 3
def _comfy_genuinely_busy(self, queue: Dict[str, Any]) -> bool:
"""Decide whether ComfyUI is really working, not just claiming to be.
ComfyUI can leave an entry in queue_running after a job dies -- observed here as
a WAN 2.1 i2v entry that sat there with the GPU at 0% and ComfyUI holding 0.56 GB.
Trusting that flag alone made this service believe ComfyUI was permanently busy,
which meant it evicted the LLM on every poll, never ran the idle purge, and never
checked whether the LLM had been squeezed onto the CPU. Half the arbitration was
disabled by one stale row.
A running entry is corroborated against GPU utilisation before it is believed.
"""
running = queue.get("queue_running") or []
pending = queue.get("queue_pending") or []
if pending:
self._running_since = None
self._running_id = None
return True
if not running:
self._running_since = None
self._running_id = None
self.comfy_stale_job = None
return False
entry = running[0]
prompt_id = entry[1] if isinstance(entry, (list, tuple)) and len(entry) > 1 else str(entry)
now = time.time()
if prompt_id != self._running_id:
self._running_id = prompt_id
self._running_since = now
self._peak_comfy_bytes = 0
snap = get_process_vram_bytes()
self._peak_comfy_bytes = max(self._peak_comfy_bytes, snap.get("comfyui_bytes", 0))
elapsed = now - (self._running_since or now)
if elapsed > self.STALE_RUNNING_S and self._peak_comfy_bytes < self.STALE_COMFY_BYTES:
if self.comfy_stale_job != prompt_id:
logger.warning(
f"ComfyUI reports prompt {prompt_id} running for {int(elapsed)}s while "
f"holding only {self._peak_comfy_bytes / (1024**3):.2f} GB — no checkpoint "
f"is loaded, so the queue entry is stale. Ignoring it; otherwise ComfyUI "
f"looks permanently busy and arbitration stops working.")
self.comfy_stale_job = prompt_id
return False
return True
async def _tenant_state(self) -> List[Dict[str, Any]]:
"""Current VRAM and busy state for every configured tenant."""
snap = get_process_vram_bytes()
stats = get_gpu_hardware_stats()
by_tenant = (stats.get("breakdown", {}) or {}).get("by_tenant_gb", {})
out = []
for t in tenants_mod.load_tenants():
if not t.enabled:
continue
bucket = _BUCKET_ALIASES.get(t.name, t.name)
vram_gb = by_tenant.get(bucket, 0.0)
probe = await tenants_mod.probe_busy(t, vram_gb=vram_gb)
out.append({
"name": t.name,
"priority": t.priority,
"vram_gb": vram_gb,
"busy": bool(probe.get("busy")),
"below_floor": bool(probe.get("below_floor")),
"reclaimable": t.reclaimable,
"needs_vram_gb": t.needs_vram_gb,
"overclock_profile": t.overclock_profile,
"idle_release_after_s": t.idle_release_after_s,
"reason": probe.get("reason"),
})
self._last_tenant_state = {"ts": time.time(), "free_gb":
round(snap["free_bytes"] / (1024**3), 2),
"tenants": out}
return out
async def _release_tenant(self, name: str, reason: str) -> Dict[str, Any]:
"""Release one tenant's VRAM by whatever mechanism it declares."""
t = tenants_mod.get_tenant(name)
if not t or not t.reclaimable:
return {"success": False, "reason": "not reclaimable"}
models = None
if t.release.per_model:
state = await get_ollama_live_state()
models = [m.get("name") for m in state.get("loaded_models", []) if m.get("name")]
logger.info(f"Releasing VRAM from '{name}': {reason}")
res = await tenants_mod.release_vram(t, models=models)
self.stats["tenant_releases"] = self.stats.get("tenant_releases", 0) + 1
return res
IDLE_PROFILE = "balanced"
def _apply_profile_for_active(self, state: List[Dict[str, Any]]) -> None:
"""Apply the GPU profile declared by whichever tenant is currently working.
This used to be two calls naming 'comfy' and 'ollama' directly, so a third
application could never get tuned clocks. The highest-priority busy tenant wins;
with nothing working the card returns to the idle profile.
"""
busy = [s for s in state if s["busy"] and s.get("overclock_profile")]
if busy:
busy.sort(key=lambda s: -s["priority"])
self._apply_oc_profile(busy[0]["overclock_profile"])
else:
self._apply_oc_profile(self.IDLE_PROFILE)
async def _arbitrate(self) -> None:
"""Generic arbitration over any number of tenants.
The two-application version was a pair of hardcoded rules -- yield Ollama when
ComfyUI is busy, purge ComfyUI when Ollama is starved -- which could not express
a third participant at all. This works from the registry instead: a busy tenant
that lacks the VRAM it declares it needs is starved, and the memory comes from
idle reclaimable tenants below it in priority, lowest first.
"""
state = await self._tenant_state()
free_gb = self._last_tenant_state["free_gb"]
self._apply_profile_for_active(state)
# 1. Starvation: highest-priority demanding tenant first.
for s in sorted(state, key=lambda x: -x["priority"]):
if not s["busy"] or not s["needs_vram_gb"]:
continue
# Starved means it cannot reach what it needs even counting what it already
# holds. Comparing free VRAM alone flagged a tenant that was working
# perfectly well on 13 GB as demanding, purely because little was left over
# -- which is the normal state of a busy GPU, and would have caused
# pointless releases from everyone else.
if s["vram_gb"] + free_gb >= s["needs_vram_gb"]:
continue
plan = tenants_mod.plan_release(s["name"], state, free_gb, s["needs_vram_gb"])
self.last_arbitration = {"ts": time.time(), "demanding": s["name"],
"free_gb": free_gb, **plan}
if not plan["release"]:
logger.debug(f"'{s['name']}' is short of VRAM but {plan['reason']}")
return
if time.time() - self.last_reclaim_time < self.RECLAIM_COOLDOWN_S:
return
self.last_reclaim_time = time.time()
for victim in plan["release"]:
await self._release_tenant(
victim, f"{s['name']} needs {s['needs_vram_gb']} GB, {free_gb} GB free")
self.last_action = (f"Released {', '.join(plan['release'])} so "
f"'{s['name']}' could work")
return
# 2. Idle release: a tenant holding VRAM it is not using, after a grace period.
now = time.time()
for s in state:
if not s["reclaimable"] or s["vram_gb"] <= 0.25:
self._idle_since.pop(s["name"], None)
continue
if s["busy"]:
self._idle_since.pop(s["name"], None)
continue
since = self._idle_since.setdefault(s["name"], now)
grace = s["idle_release_after_s"]
if grace and (now - since) >= grace:
self._idle_since.pop(s["name"], None)
await self._release_tenant(
s["name"], f"idle {int(now - since)}s holding {s['vram_gb']} GB")
self.last_action = (f"Released idle '{s['name']}' after "
f"{int(now - since)}s")
return
async def _poll_watchdog(self):
"""Fallback for when the WebSocket is down. One cheap /queue call, 1 Hz.
The previous version hit /system_stats and /queue every 300 ms on fresh TCP
connections — roughly 6.6 requests/second against ComfyUI, forever.
"""
while self.running:
interval = self.WATCHDOG_INTERVAL_WS_OK_S if self.connected_ws else self.WATCHDOG_INTERVAL_S
try:
client = _client(COMFY_API_BASE, 3.0)
resp = await client.get("/queue")
if resp.status_code == 200:
q = resp.json()
busy = self._comfy_genuinely_busy(q)
if busy:
self.watchdog_branches["busy"] += 1
await self.trigger_comfy_priority("Watchdog saw an active queue")
elif self.comfy_was_active:
self.watchdog_branches["completed"] += 1
await self.trigger_comfy_completed()
else:
self.watchdog_branches["idle_check"] += 1
await self._arbitrate()
else:
self.watchdog_branches["bad_status"] += 1
except Exception as e:
# This used to swallow everything silently, including anything raised by
# the starvation check, which is why that check could appear to run and
# do nothing.
self.watchdog_branches["error"] += 1
self.last_watchdog_error = str(e)[:200]
logger.debug(f"watchdog poll error: {e}")
await asyncio.sleep(interval)
def suspend_oc(self, reason: str = "tuning sweep") -> None:
self.oc_suspended = True
logger.info(f"Overclock auto-switching suspended ({reason})")
def resume_oc(self, profile: Optional[str] = None) -> None:
self.oc_suspended = False
# Forget the cached profile so the next transition actually reapplies.
self.oc_profile = profile
logger.info("Overclock auto-switching resumed")
def _apply_oc_profile(self, profile: str):
"""Apply an overclock profile in a background thread; only fire on transition."""
if self.oc_suspended or self.oc_profile == profile:
return
self.oc_profile = profile
try:
loop = asyncio.get_running_loop()
loop.run_in_executor(None, overclock_manager.apply_profile, profile)
logger.info(f"🎛️ Overclock profile switched -> '{profile}'")
except Exception as e:
logger.warning(f"Overclock profile switch failed ({profile}): {e}")
def get_status(self) -> Dict[str, Any]:
idle_for = (time.time() - self.comfy_idle_since) if self.comfy_idle_since else None
return {
"running": self.running,
"connected_ws": self.connected_ws,
"last_action": self.last_action,
"mode": "Bidirectional Hot-Swap (ComfyUI <-> Ollama)",
"comfy_active": self.comfy_was_active,
"pending_purge": self.pending_purge,
"comfy_idle_s": round(idle_for, 1) if idle_for is not None else None,
"idle_purge_after_s": self.COMFY_IDLE_PURGE_S,
"oc_profile": self.oc_profile,
"counters": dict(self.stats),
"comfy_stale_job": self.comfy_stale_job,
"event_sources": dict(self.event_sources),
"last_arbitration": self.last_arbitration,
"tenant_state": self._last_tenant_state,
"watchdog_branches": dict(self.watchdog_branches),
"last_watchdog_error": self.last_watchdog_error,
"yield_backoff": {m: round(max(t - time.time(), 0), 1)
for m, t in self._yield_backoff_until.items()
if t > time.time()},
}
arbitrator = AutoArbitrator()