Files
gpu-program-swapper/vram_arbitrator.py

505 lines
21 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
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
SWITCH_HISTORY = deque(maxlen=50)
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
# 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,
"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:
async with httpx.AsyncClient(timeout=3.0) as client:
# Check running models (ps)
ps_resp = await client.get(f"{OLLAMA_API_BASE}/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(f"{OLLAMA_API_BASE}/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:
async with httpx.AsyncClient(timeout=3.0) as client:
# Check system stats
stats_resp = await client.get(f"{COMFY_API_BASE}/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(f"{COMFY_API_BASE}/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 instant_free_ollama_vram(model_name: Optional[str] = None) -> Dict[str, Any]:
"""Tell Ollama to instantly yield VRAM without evicting from OS page cache."""
t0 = time.perf_counter()
if not model_name:
ollama_state = await get_ollama_live_state()
model_name = ollama_state.get("active_model_name")
if not model_name:
return {"success": True, "message": "No active Ollama model in VRAM", "duration_ms": 0}
try:
async with httpx.AsyncClient(timeout=5.0) as client:
resp = await client.post(
f"{OLLAMA_API_BASE}/api/generate",
json={"model": model_name, "keep_alive": 0},
)
duration_ms = round((time.perf_counter() - t0) * 1000, 2)
event = {
"timestamp": time.strftime("%H:%M:%S"),
"event_type": "Ollama VRAM Yield",
"source": model_name,
"target": "VRAM 0MB (Kept in RAM)",
"duration_ms": duration_ms,
"cache_status": "RAM-Cached",
}
SWITCH_HISTORY.appendleft(event)
return {"success": True, "model": model_name, "duration_ms": duration_ms}
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:
async with httpx.AsyncClient(timeout=5.0) as client:
resp = await client.post(
f"{COMFY_API_BASE}/free",
json={"unload_models": True, "free_memory": True},
)
duration_ms = round((time.perf_counter() - t0) * 1000, 2)
event = {
"timestamp": time.strftime("%H:%M:%S"),
"event_type": "ComfyUI VRAM Purge",
"source": "ComfyUI Pipeline",
"target": "VRAM Free",
"duration_ms": duration_ms,
"cache_status": "Cleaned",
}
SWITCH_HISTORY.appendleft(event)
return {"success": True, "duration_ms": duration_ms}
except Exception as e:
return {"success": False, "error": str(e), "duration_ms": round((time.perf_counter() - t0) * 1000, 2)}
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:
async with httpx.AsyncClient(timeout=180.0) as client:
resp = await client.post(
f"{OLLAMA_API_BASE}/api/generate",
json={"model": target_model, "prompt": "Ready check", "stream": False, "keep_alive": keep_alive},
)
total_duration = time.perf_counter() - t0
total_duration_ms = round(total_duration * 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)
# Check if it was a RAM cache hit (load duration < 1500ms for large model indicates RAM hit)
is_ram_hit = load_dur_ms < 2500
event = {
"timestamp": time.strftime("%H:%M:%S"),
"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,
"cache_status": "RAM Cache Hit ⚡" if is_ram_hit else "Cold Disk Load 💾",
}
SWITCH_HISTORY.appendleft(event)
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,
"is_ram_hit": is_ram_hit,
"response": data.get("response", ""),
}
else:
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."""
def __init__(self):
self.running = False
self.ws_task: Optional[asyncio.Task] = None
self.poll_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.oc_profile = None
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())
logger.info("AutoArbitrator background engine started (Bidirectional).")
# Apply the default (balanced) overclock profile on startup.
try:
await asyncio.get_event_loop().run_in_executor(None, overclock_manager.apply_profile, "balanced")
except Exception as e:
logger.warning(f"Startup overclock apply failed: {e}")
async def stop(self):
self.running = False
if self.ws_task:
self.ws_task.cancel()
if self.poll_task:
self.poll_task.cancel()
logger.info("AutoArbitrator background engine stopped.")
async def trigger_comfy_priority(self, reason: str = "ComfyUI prompt detected"):
"""Instantly yield Ollama VRAM to 0MB when ComfyUI needs to run diffusion models."""
self.comfy_was_active = True
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}' from VRAM...")
self.last_yield_time = time.time()
res = await instant_free_ollama_vram(model)
dur = res.get("duration_ms", 0)
self.last_action = f"Auto-yielded '{model}' for ComfyUI ({dur}ms)"
logger.info(f"Ollama auto-yield completed: {res}")
async def trigger_comfy_completed(self):
"""Purge ComfyUI VRAM cache when generation finishes, keeping VRAM 100% free for Ollama."""
now = time.time()
if now - self.last_comfy_free_time < 3.0:
return
self.last_comfy_free_time = now
self.comfy_was_active = False
self._apply_oc_profile("ollama")
logger.info("⚡ ComfyUI finished generation -> Auto-purging ComfyUI VRAM cache for Ollama...")
res = await instant_free_comfyui_vram()
dur = res.get("duration_ms", 0)
self.last_action = f"Auto-purged ComfyUI VRAM ({dur}ms) - Ready for Ollama"
logger.info(f"ComfyUI auto-purge completed: {res}")
async def _ws_listener(self):
client_id = "hyperswap-arbitrator"
ws_url = f"ws://127.0.0.1:8188/ws?clientId={client_id}"
while self.running:
try:
async with websockets.connect(ws_url, ping_interval=10, ping_timeout=10) as ws:
self.connected_ws = True
logger.info("AutoArbitrator connected to ComfyUI WebSocket.")
while self.running:
msg = await ws.recv()
if isinstance(msg, str):
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:
# Prompt queue finished
await asyncio.sleep(1.5)
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 node is None and self.comfy_was_active:
# Finished executing graph
await asyncio.sleep(1.5)
await self.trigger_comfy_completed()
elif msg_type == "execution_success":
await asyncio.sleep(1.5)
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(2.0)
async def _poll_watchdog(self):
"""Watchdog polling /queue every 300ms for robust bidirectional arbitration."""
while self.running:
try:
comfy_state = await get_comfyui_live_state()
is_executing = comfy_state.get("queue_running", 0) > 0 or comfy_state.get("queue_remaining", 0) > 0 or comfy_state.get("executing", False)
if is_executing:
await self.trigger_comfy_priority("Polling detected active queue/execution")
elif self.comfy_was_active and not is_executing:
await asyncio.sleep(1.5)
await self.trigger_comfy_completed()
except Exception:
pass
await asyncio.sleep(0.3)
def _apply_oc_profile(self, profile: str):
"""Apply an overclock profile in a background thread; only fire on transition."""
if self.oc_profile == profile:
return
self.oc_profile = profile
try:
loop = asyncio.get_event_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]:
return {
"running": self.running,
"connected_ws": self.connected_ws,
"last_action": self.last_action,
"mode": "Bidirectional Hot-Swap (ComfyUI <-> Ollama)",
}
arbitrator = AutoArbitrator()