fix: address review feedback on VRAM autounload

- PromptRefiner unload_model toggle now actually controls eviction:
  keep_alive is "0s" when on, None (server default) when off, so the
  UI no longer unloads against the user's choice.
- Cleanup always runs; unload fires only when unload_model is on AND the
  subprocess fallback left a model loaded, instead of skipping cleanup
  entirely when the toggle is off.
- release_vram() clears empty_cache() on every CUDA device, not just
  device 0, for multi-GPU rigs.
- Add docstrings to INPUT_TYPES methods and language ids to plan doc
  code fences.
This commit is contained in:
limbicnation
2026-06-23 04:52:43 +02:00
parent 3bfc4f7db7
commit 77b2c48fde
5 changed files with 27 additions and 19 deletions
+3 -3
View File
@@ -11,7 +11,7 @@
The ComfyUI `PromptRefinerNode` (and sibling nodes: `PromptGeneratorNode`, `NegativePromptNode`, `PromptDualStreamRefinerNode`) calls local LLMs via Ollama for text generation. After each node completes, Ollama's default behavior retains the loaded model in GPU VRAM for **5 minutes** (`keep_alive=5m`). When the ComfyUI pipeline subsequently attempts to load VRAM-heavy diffusion models (e.g., LTXAVTEModel_ text encoder at 11.2 GB staged), CUDA OOM occurs:
```
```text
RuntimeError: VRAM grow failed: 2013757440 bytes
```
@@ -30,7 +30,7 @@ RuntimeError: VRAM grow failed: 2013757440 bytes
## Current State (Before Fix)
```
```text
Timeline (from ComfyUI error logs):
03:38:21 — PromptRefiner calls qwen3:8b → Ollama loads ~4.7 GB into VRAM
03:38:27 — PromptGenerator calls qwen3-4b-deforum → Ollama loads ~2.6 GB (qwen3:8b evicted, swapped)
@@ -53,7 +53,7 @@ Timeline (from ComfyUI error logs):
## Target State (After Fix)
```
```text
Timeline (with autounload):
03:38:21 — PromptRefiner calls qwen3:8b with keep_alive="0s"
03:38:27 — Model evicted from VRAM immediately, torch.cuda.empty_cache() runs async
+5 -1
View File
@@ -433,7 +433,11 @@ class OllamaClient:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
# empty_cache() only frees the current device; clear every GPU so
# multi-GPU rigs don't leave cached blocks on non-default devices.
for device in range(torch.cuda.device_count()):
with torch.cuda.device(device):
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
except ImportError:
pass
+1
View File
@@ -47,6 +47,7 @@ Negative prompt:"""
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
"""Define input parameters for the node."""
styles = list(cls.STYLE_HINTS.keys())
return {
"required": {
+1
View File
@@ -80,6 +80,7 @@ Description: {prompt}"""
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
"""Define input parameters for the node."""
available_models = cls._get_available_models()
return {
"required": {
+17 -15
View File
@@ -41,6 +41,7 @@ Refined prompt:"""
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
"""Define input parameters for the node."""
return {
"required": {
"prompt": (
@@ -164,6 +165,10 @@ Refined prompt:"""
# Determine effective seed
effective_seed: int | None = None if seed == -1 else seed
# Honour the unload_model toggle: "0s" evicts immediately, None keeps the
# model loaded for the Ollama server default (so the UI toggle is truthful).
keep_alive = "0s" if unload_model else None
used_subprocess = False
try:
for i in range(passes):
@@ -179,8 +184,7 @@ Refined prompt:"""
# Derive per-pass seed so multi-pass refinement isn't a no-op
pass_seed = None if effective_seed is None else effective_seed + i
# Generate refined version with keep_alive="0s" to unload
# the model from Ollama VRAM immediately after each pass.
# keep_alive evicts (or retains) the model per the unload_model toggle.
result = client.generate_streaming(
model=model,
prompt=refinement,
@@ -189,7 +193,7 @@ Refined prompt:"""
timeout=timeout,
pbar=pbar,
seed=pass_seed,
keep_alive="0s",
keep_alive=keep_alive,
)
if result.kind == "ok" and result.text is not None:
@@ -218,15 +222,13 @@ Refined prompt:"""
return (current_prompt,)
finally:
# Always release VRAM after node execution, regardless of success/failure.
# keep_alive="0s" already handles Ollama-side unloading; this covers
# the PyTorch CUDA allocator cache.
if unload_model:
OllamaClient.cleanup_async(
model=model,
logger_prefix="PromptRefiner",
# streaming path already evicted via keep_alive="0s"; only the
# subprocess fallback leaves a model loaded at the 5m default.
unload=used_subprocess,
release_cuda=True,
)
# Always release the PyTorch CUDA cache, regardless of success/failure.
# Evict the Ollama model only when the user asked to unload AND the
# subprocess fallback left one loaded at the 5m default (the streaming
# path already evicted it via keep_alive="0s").
OllamaClient.cleanup_async(
model=model,
logger_prefix="PromptRefiner",
unload=used_subprocess and unload_model,
release_cuda=True,
)