The profiles were hand-written and had never been checked against the hardware. Adding a ComfyUI benchmark alongside the existing decode one made the compute side measurable for the first time, and most of what the profiles configured turned out to do nothing. Measured on this card (RTX 4080 SUPER, driver 595.84): - LLM decode is not power-bound: 73.0-73.5 tok/s flat from 222W to 370W, with the card never drawing more than 224W at any limit. The ollama profile's 370W did nothing. - Diffusion is power-bound: 5.48 it/s @222W rising to 6.71 @370W, so comfy's 370W is worth a real +2.8% over the 320W stock default. - Clock locks did nothing for either workload: 72.6 tok/s locked at 11251MHz vs 72.7 unlocked; 6.77 it/s locked at 3105MHz vs 6.73 unlocked, and 6.78 at 2400MHz. - Memory bandwidth is still the decode bottleneck (5001MHz halves throughput to 35.9 tok/s), confirming the profile's premise -- the card just gets there unaided. - Fans: 48,435 samples show 81C all-time max and zero thermal throttle events, while the ollama profile held 49.6C average by running fans at 87%. All profiles now use automatic fans and let the thermal governor escalate on demand. Code changes supporting that: - _diffusion_benchmark() queues a fixed SDXL graph via ComfyUI's API. The seed must vary per run: ComfyUI caches by node inputs, so a fixed seed returned in ~1ms without executing. Implausibly fast results are now rejected as cache hits rather than recorded as record scores. - The arbitrator's automatic profile switching is suspended during a sweep. A diffusion benchmark trips trigger_comfy_priority, which reapplies the whole profile and would silently overwrite the clock being measured. - _supported_clocks() queries the mem,gr pair; asking for a single field returned one column and reading index 1 yielded an empty list rather than an error. Graphics clocks are subsampled (the card enumerates 194 of them) and lock sweeps include an explicit unlocked control step. - offsets_supported() probes once and apply_profile skips inert offset levers with an explanation instead of pretending they applied. - Profiles carry a 'measured' field recording the evidence behind each setting. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
889 lines
36 KiB
Python
889 lines
36 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 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 used to classify how a model actually got into VRAM.
|
|
# PCIe 4.0 x16 tops out near 31.5 GB/s; this NVMe sustains well under 2 GB/s.
|
|
RAM_HIT_GBPS = 5.0
|
|
PARTIAL_HIT_GBPS = 1.5
|
|
|
|
# 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
|
|
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}
|
|
if not NVML_AVAILABLE:
|
|
return out
|
|
try:
|
|
handle = pynvml.nvmlDeviceGetHandleByIndex(0)
|
|
out["free_bytes"] = pynvml.nvmlDeviceGetMemoryInfo(handle).free
|
|
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():
|
|
kind = _PID_KIND_CACHE.get(pid)
|
|
if kind is None:
|
|
kind = _classify_pid(pid)
|
|
_PID_KIND_CACHE[pid] = kind
|
|
if kind == "ollama":
|
|
out["ollama_bytes"] += used
|
|
elif kind == "comfy":
|
|
out["comfyui_bytes"] += used
|
|
else:
|
|
out["other_bytes"] += used
|
|
except Exception as e:
|
|
logger.debug(f"get_process_vram_bytes failed: {e}")
|
|
return out
|
|
|
|
|
|
_PID_KIND_CACHE: Dict[int, str] = {}
|
|
|
|
|
|
def _classify_pid(pid: int) -> str:
|
|
try:
|
|
proc = psutil.Process(pid)
|
|
pname = proc.name().lower()
|
|
cmdline = " ".join(proc.cmdline()).lower()
|
|
except Exception:
|
|
return "other"
|
|
if "ollama" in pname or "llama-server" in cmdline:
|
|
return "ollama"
|
|
if "comfy" in cmdline or "main.py" in cmdline:
|
|
return "comfy"
|
|
return "other"
|
|
|
|
|
|
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,
|
|
"processes": []
|
|
}
|
|
|
|
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
|
|
|
|
is_ollama = "ollama" in pname.lower() or "llama-server" in cmdline.lower()
|
|
is_comfy = "comfy" in cmdline.lower() or "main.py" in cmdline.lower()
|
|
|
|
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
|
|
|
|
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,
|
|
})
|
|
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),
|
|
"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": []
|
|
}
|
|
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]:
|
|
"""Block until Ollama's VRAM has actually drained, or we give up.
|
|
|
|
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.
|
|
"""
|
|
t0 = time.perf_counter()
|
|
last = baseline_bytes
|
|
while True:
|
|
snap = get_process_vram_bytes()
|
|
last = snap["ollama_bytes"]
|
|
if last <= YIELD_RESIDUAL_BYTES:
|
|
return {
|
|
"confirmed": True,
|
|
"confirm_ms": round((time.perf_counter() - t0) * 1000, 2),
|
|
"residual_bytes": last,
|
|
"free_bytes": snap["free_bytes"],
|
|
}
|
|
if (time.perf_counter() - t0) >= timeout_s:
|
|
return {
|
|
"confirmed": False,
|
|
"confirm_ms": round((time.perf_counter() - t0) * 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",
|
|
}
|
|
await asyncio.sleep(YIELD_CONFIRM_POLL_S)
|
|
|
|
|
|
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) -> 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:
|
|
client = _client(OLLAMA_API_BASE, 5.0)
|
|
await asyncio.gather(*[
|
|
client.post("/api/generate", json={"model": t, "keep_alive": 0})
|
|
for t in targets
|
|
], return_exceptions=True)
|
|
request_ms = round((time.perf_counter() - t0) * 1000, 2)
|
|
|
|
barrier = {"confirmed": None, "confirm_ms": 0.0, "residual_bytes": baseline}
|
|
if confirm:
|
|
barrier = await _await_vram_release(baseline)
|
|
|
|
duration_ms = round((time.perf_counter() - t0) * 1000, 2)
|
|
freed_gb = round(max(baseline - barrier.get("residual_bytes", 0), 0) / (1024**3), 2)
|
|
|
|
_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",
|
|
"detail": barrier.get("error"),
|
|
})
|
|
return {
|
|
"success": True,
|
|
"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),
|
|
"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.
|
|
return {
|
|
"cache_status": "RAM Cache Hit ⚡" if load_duration_ms < 2500 else "Cold Disk Load 💾",
|
|
"load_gbps": None,
|
|
"is_ram_hit": load_duration_ms < 2500,
|
|
"detail": "size unknown, fell back to duration heuristic",
|
|
}
|
|
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") -> Dict[str, Any]:
|
|
"""High-speed hot-swap to target Ollama model, tracking swap metrics."""
|
|
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", ""),
|
|
}
|
|
return {"success": False, "error": f"HTTP {resp.status_code}: {resp.text}",
|
|
"duration_ms": total_duration_ms}
|
|
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.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
|
|
self.stats = {"yields": 0, "purges": 0, "yield_timeouts": 0, "deferred_purges": 0}
|
|
|
|
async def start(self):
|
|
if self.running:
|
|
return
|
|
self.running = True
|
|
self.ws_task = asyncio.create_task(self._ws_listener())
|
|
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):
|
|
if task:
|
|
task.cancel()
|
|
await close_clients()
|
|
logger.info("AutoArbitrator background engine stopped.")
|
|
|
|
async def trigger_comfy_priority(self, reason: str = "ComfyUI prompt detected"):
|
|
"""Yield Ollama's VRAM — and confirm it is gone — before diffusion allocates."""
|
|
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()
|
|
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)
|
|
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}")
|
|
|
|
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 _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)
|
|
|
|
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 = len(q.get("queue_running", [])) > 0 or len(q.get("queue_pending", [])) > 0
|
|
if busy:
|
|
await self.trigger_comfy_priority("Watchdog saw an active queue")
|
|
elif self.comfy_was_active:
|
|
await self.trigger_comfy_completed()
|
|
except Exception:
|
|
pass
|
|
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),
|
|
}
|
|
|
|
|
|
arbitrator = AutoArbitrator()
|
|
|
|
|