diff --git a/docs/VRAM_OPTIMIZATION_PLAN.md b/docs/VRAM_OPTIMIZATION_PLAN.md index b23ceb2..451c2ac 100644 --- a/docs/VRAM_OPTIMIZATION_PLAN.md +++ b/docs/VRAM_OPTIMIZATION_PLAN.md @@ -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 diff --git a/nodes/adapters/ollama_client.py b/nodes/adapters/ollama_client.py index 1f2a0bf..ff5c0cf 100644 --- a/nodes/adapters/ollama_client.py +++ b/nodes/adapters/ollama_client.py @@ -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 diff --git a/nodes/negative_prompt_node.py b/nodes/negative_prompt_node.py index 974f4b0..66ab2e4 100644 --- a/nodes/negative_prompt_node.py +++ b/nodes/negative_prompt_node.py @@ -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": { diff --git a/nodes/prompt_dual_stream_refiner_node.py b/nodes/prompt_dual_stream_refiner_node.py index d6e4a3f..ff80ff8 100644 --- a/nodes/prompt_dual_stream_refiner_node.py +++ b/nodes/prompt_dual_stream_refiner_node.py @@ -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": { diff --git a/nodes/prompt_refiner_node.py b/nodes/prompt_refiner_node.py index f706af5..e357f0b 100644 --- a/nodes/prompt_refiner_node.py +++ b/nodes/prompt_refiner_node.py @@ -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, + )