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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user