diff --git a/vram_arbitrator.py b/vram_arbitrator.py index 6473fcb..071fd59 100644 --- a/vram_arbitrator.py +++ b/vram_arbitrator.py @@ -35,7 +35,12 @@ RAM_HIT_GBPS = 5.0 PARTIAL_HIT_GBPS = 1.5 # How long Ollama's VRAM may take to actually drain before we stop waiting. -YIELD_CONFIRM_TIMEOUT_S = 3.0 +# 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" @@ -445,18 +450,29 @@ async def instant_free_ollama_vram(model_name: Optional[str] = None, the HTTP POST took. """ t0 = time.perf_counter() - if not model_name: + 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() - model_name = ollama_state.get("active_model_name") + 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 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 client.post("/api/generate", json={"model": model_name, "keep_alive": 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} @@ -468,7 +484,7 @@ async def instant_free_ollama_vram(model_name: Optional[str] = None, _record({ "event_type": "Ollama VRAM Yield", - "source": model_name, + "source": ", ".join(targets)[:200], "target": "VRAM 0MB (Kept in RAM)", "duration_ms": duration_ms, "yield_confirm_ms": barrier.get("confirm_ms"), @@ -478,6 +494,7 @@ async def instant_free_ollama_vram(model_name: Optional[str] = None, return { "success": True, "model": model_name, + "models_unloaded": targets, "duration_ms": duration_ms, "request_ms": request_ms, "confirm_ms": barrier.get("confirm_ms"),