diff --git a/README.md b/README.md index 0187bd1..a022059 100644 --- a/README.md +++ b/README.md @@ -6,7 +6,7 @@ This repo does three related jobs: 1. **Installs startup-time monkey patches** before any workflow nodes run. 2. Ships **KJ-style resident loader nodes** for diffusion models and checkpoints. -3. Maintains a **live residency system** for native ComfyUI objects and compatible external caches, with preload / pin / evict / report controls for native tracked objects and automatic snapshot / trim / eviction support for compatible external entries. +3. Maintains a **live residency system** for native ComfyUI objects and compatible external caches, with preload / pin / evict / report controls for native tracked objects and automatic snapshot support plus provider-specific eviction for compatible external entries. It is not just a “clean RAM” addon. The main target is the path from model file -> tensors -> live ComfyUI object -> VRAM retention, including GPU-resident caches that live outside ComfyUI’s normal loaded-model list. @@ -24,7 +24,7 @@ This repo focuses on both: - For **`.safetensors`**, it tries to keep eligible loads on the narrowest, most GPU-friendly path it can. - For **resident diffusion and checkpoint-model loads**, it avoids broad checkpoint materialization by selecting only the detected UNet keys where possible. - For **runtime VRAM pressure**, it adds a sticky-priority registry and teaches ComfyUI’s unload path to protect higher-value resident entries until enough VRAM must be reclaimed. -- For **compatible external GPU caches** that bypass `comfy.model_management.current_loaded_models`, it can discover supported providers at runtime and include their entries in snapshot and trim decisions. +- For **compatible external GPU caches** that bypass `comfy.model_management.current_loaded_models`, it can discover supported providers at runtime and include their entries in snapshot output and provider-specific eviction decisions, with automatic trim kept opt-in. - For **manual control**, it exposes nodes that let you preload, pin, evict, and inspect tracked native models, CLIPs, and VAEs. ## What changes at startup @@ -115,8 +115,8 @@ Those entries: - are refreshed from the live cached object at runtime - appear in **Registry Snapshot** output under `external_entries` -- are considered by load-scoped VRAM trimming even though they are not part of `current_loaded_models` - are evicted through SeedVR2’s own cache-removal methods rather than the normal Comfy unload path +- remain outside automatic trim unless `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1` is set ### 4) Device/offload policy is overridden @@ -341,7 +341,7 @@ The trim path prefers to: - then lower-priority sticky entries - preserve explicitly kept models - use partial unload where available -- include compatible external cache entries in the same candidate search when they are visible, on the same device, and not currently claimed/in use +- include compatible external cache entries in the same candidate search only when `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1`, they are visible, on the same device, and not currently claimed/in use This logic lives in the resident loader path. You do not need a separate “target free VRAM” node for it. @@ -408,7 +408,7 @@ Typical external-entry fields include: - `alive` - `external` -For SeedVR2-backed entries, `claimed: true` means the cache object is currently marked in use and is skipped by the automatic trim candidate search. +For SeedVR2-backed entries, `claimed: true` means the cache object is currently marked in use and is skipped by the external trim candidate search when that opt-in path is enabled. ## What gets tracked @@ -432,7 +432,7 @@ Current external integration coverage is: - SeedVR2 global cached **DiT** entries - SeedVR2 global cached **VAE** entries -Those entries are discovered lazily from compatible SeedVR2 cache modules at runtime. They are tracked separately from the native registry and participate in snapshot, load-scoped trim, and provider-specific eviction decisions. +Those entries are discovered lazily from compatible SeedVR2 cache modules at runtime. They are tracked separately from the native registry and participate in snapshot plus provider-specific eviction decisions, with load-scoped trim available only through the explicit external-trim opt-in. ## Important limits and non-goals @@ -477,7 +477,7 @@ The native registry only sees objects that pass through the patched ComfyUI load If another custom node loads models through its own private code path and bypasses those patched entry points, that object may never become a tracked native registry entry. In that case, the preload / pin / evict / report nodes from this repo cannot manage it until that external loader is integrated or patched. -Likewise, even for supported external providers such as SeedVR2, the current external integration is about **observation + automatic trim/eviction**. This repo does **not** yet expose dedicated external preload / pin / report / evict nodes for provider-owned cache entries. +Likewise, even for supported external providers such as SeedVR2, the current external integration is about **observation + provider-specific eviction**, with automatic trim kept opt-in. This repo does **not** yet expose dedicated external preload / pin / report / evict nodes for provider-owned cache entries. ## Installation @@ -541,7 +541,7 @@ If SeedVR2 keeps DiT or VAE models in its own global cache, those objects can no That means: - you can see that those bytes exist even though they are outside `current_loaded_models` -- resident loader trim can reclaim them automatically when they are visible, on the same device, and not currently claimed/in use +- resident loader trim only reclaims them when `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1`, they are visible, on the same device, and not currently claimed/in use - eviction goes through SeedVR2’s own cache-removal path instead of a normal Comfy wrapper unload ## Notes on compatibility and migration diff --git a/cleanup.py b/cleanup.py index 648f559..7861482 100644 --- a/cleanup.py +++ b/cleanup.py @@ -7,7 +7,7 @@ from typing import Any import comfy.model_management as model_management import torch -from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed +from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed, external_trim_enabled from .residency import REGISTRY _ADAPTIVE_HEADROOM_RATIO = 0.125 @@ -154,13 +154,14 @@ def _trim_candidates( ) candidates.append((loaded, entry, sticky_respected, False)) - for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates( - device=device, - respect_sticky=respect_sticky, - sticky_floor_priority=sticky_floor_priority, - keep_models=keep_models, - ): - candidates.append((external_obj, entry, sticky_respected, True)) + if external_trim_enabled(): + for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates( + device=device, + respect_sticky=respect_sticky, + sticky_floor_priority=sticky_floor_priority, + keep_models=keep_models, + ): + candidates.append((external_obj, entry, sticky_respected, True)) candidates.sort(key=lambda item: _sort_key_for_candidate(item[1], sticky_respected=item[2])) return candidates @@ -305,6 +306,7 @@ def trim_resident_vram( "respect_sticky": bool(respect_sticky), "sticky_floor_priority": int(sticky_floor_priority), "allow_partial_unload": bool(allow_partial_unload), + "external_trim_enabled": bool(external_trim_enabled()), "actions": actions, } diff --git a/external_residency.py b/external_residency.py index ed90306..3f3ba98 100644 --- a/external_residency.py +++ b/external_residency.py @@ -280,7 +280,7 @@ class ExternalResidencyRegistry: entry.state_provider = state_provider entry.evict_callback = evict_callback entry.last_touched = _now() - if note: + if note and note not in entry.notes: entry.notes.append(note) self.refresh_runtime_state() @@ -398,6 +398,11 @@ class ExternalResidencyRegistry: EXTERNAL_REGISTRY = ExternalResidencyRegistry() +def external_trim_enabled() -> bool: + value = os.environ.get("COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL", "").strip().lower() + return value in {"1", "true", "yes", "on"} + + def _call_seedvr2_method_with_optional_expected_model( method: Callable[..., Any], *args: Any, @@ -418,18 +423,18 @@ def _call_seedvr2_method_with_optional_expected_model( return method(*args, **kwargs) -def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, model: Any) -> None: +def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, model: Any) -> ExternalResidencyEntry | None: if not isinstance(config, Mapping): _LOG.debug( "GPU Resident Loader: skipping SeedVR2 %s cache entry with unexpected config type: %s", kind, type(config).__name__, ) - return + return None node_id = config.get("node_id") if node_id is None or model is None: - return + return None try: model_ref = weakref.ref(model) @@ -464,7 +469,7 @@ def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, ) note = f"SeedVR2 cached VAE node {node_id}" - EXTERNAL_REGISTRY.bind( + return EXTERNAL_REGISTRY.bind( cache_key=cache_key, obj=model, kind=registry_kind, @@ -485,24 +490,33 @@ def _install_seedvr2_integration_for_module(module: Any) -> bool: class_id = id(model_cache_cls) thread_id = threading.get_ident() - with _SEEDVR2_PATCH_CONDITION: - while ( - class_id in _SEEDVR2_PATCHING_CLASS_IDS - and _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) != thread_id - ): - _SEEDVR2_PATCH_CONDITION.wait() - if class_id in _SEEDVR2_PATCHED_CLASS_IDS or _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) == thread_id: - return True - _SEEDVR2_PATCHING_CLASS_IDS.add(class_id) - _SEEDVR2_PATCHING_THREAD_IDS[class_id] = thread_id + original_set_dit = None + original_set_vae = None + original_replace_dit = None + original_replace_vae = None + original_remove_dit = None + original_remove_vae = None + provisional_cache_keys: set[str] = set() + provisional_cache_ownership: dict[str, str] = {} try: - original_set_dit = model_cache_cls.set_dit - original_set_vae = model_cache_cls.set_vae - original_replace_dit = getattr(model_cache_cls, "replace_dit", None) - original_replace_vae = getattr(model_cache_cls, "replace_vae", None) - original_remove_dit = model_cache_cls.remove_dit - original_remove_vae = model_cache_cls.remove_vae + with _SEEDVR2_PATCH_CONDITION: + while ( + class_id in _SEEDVR2_PATCHING_CLASS_IDS + and _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) != thread_id + ): + _SEEDVR2_PATCH_CONDITION.wait() + if class_id in _SEEDVR2_PATCHED_CLASS_IDS or _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) == thread_id: + return True + _SEEDVR2_PATCHING_CLASS_IDS.add(class_id) + _SEEDVR2_PATCHING_THREAD_IDS[class_id] = thread_id + + original_set_dit = model_cache_cls.set_dit + original_set_vae = model_cache_cls.set_vae + original_replace_dit = getattr(model_cache_cls, "replace_dit", None) + original_replace_vae = getattr(model_cache_cls, "replace_vae", None) + original_remove_dit = model_cache_cls.remove_dit + original_remove_vae = model_cache_cls.remove_vae def set_dit_wrapper(self, dit_config, model, model_name, debug=None): result = original_set_dit(self, dit_config, model, model_name, debug) @@ -570,54 +584,92 @@ def _install_seedvr2_integration_for_module(module: Any) -> bool: EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("vae", vae_config.get("node_id"))) return result - model_cache_cls.set_dit = set_dit_wrapper - model_cache_cls.set_vae = set_vae_wrapper - if original_replace_dit is not None: - model_cache_cls.replace_dit = replace_dit_wrapper - if original_replace_vae is not None: - model_cache_cls.replace_vae = replace_vae_wrapper - model_cache_cls.remove_dit = remove_dit_wrapper - model_cache_cls.remove_vae = remove_vae_wrapper global_cache = get_global_cache() model_cache_lock = getattr(global_cache, "_model_cache_lock", None) lock_context = model_cache_lock if model_cache_lock is not None else contextlib.nullcontext() with lock_context: + model_cache_cls.set_dit = set_dit_wrapper + model_cache_cls.set_vae = set_vae_wrapper + if original_replace_dit is not None: + model_cache_cls.replace_dit = replace_dit_wrapper + if original_replace_vae is not None: + model_cache_cls.replace_vae = replace_vae_wrapper + model_cache_cls.remove_dit = remove_dit_wrapper + model_cache_cls.remove_vae = remove_vae_wrapper dit_items = list(getattr(global_cache, "_dit_models", {}).items()) vae_items = list(getattr(global_cache, "_vae_models", {}).items()) - for node_id, entry in dit_items: - if not isinstance(entry, tuple) or len(entry) != 2: - continue - model, config = entry - if model is not None: - _register_seedvr2_cached_model(global_cache, kind="dit", config=config, model=model) - for node_id, entry in vae_items: - if not isinstance(entry, tuple) or len(entry) != 2: - continue - model, config = entry - if model is not None: - _register_seedvr2_cached_model(global_cache, kind="vae", config=config, model=model) - except Exception: + for _node_id, entry in dit_items: + if not isinstance(entry, tuple) or len(entry) != 2: + continue + model, config = entry + if model is not None: + if isinstance(config, Mapping) and config.get("node_id") is not None: + cache_key = _seedvr2_entry_key("dit", config.get("node_id")) + provisional_cache_keys.add(cache_key) + prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key) + else: + cache_key = None + prior_entry_id = None + registered_entry = _register_seedvr2_cached_model(global_cache, kind="dit", config=config, model=model) + if cache_key is not None and registered_entry is not None: + if prior_entry_id != registered_entry.entry_id: + provisional_cache_ownership[cache_key] = registered_entry.entry_id + for _node_id, entry in vae_items: + if not isinstance(entry, tuple) or len(entry) != 2: + continue + model, config = entry + if model is not None: + if isinstance(config, Mapping) and config.get("node_id") is not None: + cache_key = _seedvr2_entry_key("vae", config.get("node_id")) + provisional_cache_keys.add(cache_key) + prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key) + else: + cache_key = None + prior_entry_id = None + registered_entry = _register_seedvr2_cached_model(global_cache, kind="vae", config=config, model=model) + if cache_key is not None and registered_entry is not None: + if prior_entry_id != registered_entry.entry_id: + provisional_cache_ownership[cache_key] = registered_entry.entry_id with _SEEDVR2_PATCH_CONDITION: - model_cache_cls.set_dit = original_set_dit - model_cache_cls.set_vae = original_set_vae + _SEEDVR2_PATCHING_CLASS_IDS.discard(class_id) + _SEEDVR2_PATCHED_CLASS_IDS.add(class_id) + _SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None) + _SEEDVR2_PATCH_CONDITION.notify_all() + except Exception: + _LOG.debug( + "GPU Resident Loader: rolling back SeedVR2 integration for class_id=%s provisional_cache_keys=%s", + class_id, + sorted(provisional_cache_keys), + exc_info=True, + ) + for cache_key in provisional_cache_keys: + created_entry_id = provisional_cache_ownership.get(cache_key) + if created_entry_id is not None: + with EXTERNAL_REGISTRY._lock: + current_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key) + if current_entry_id == created_entry_id: + EXTERNAL_REGISTRY.remove(cache_key=cache_key) + with lock_context: + if original_set_dit is not None: + model_cache_cls.set_dit = original_set_dit + if original_set_vae is not None: + model_cache_cls.set_vae = original_set_vae if original_replace_dit is not None: model_cache_cls.replace_dit = original_replace_dit if original_replace_vae is not None: model_cache_cls.replace_vae = original_replace_vae - model_cache_cls.remove_dit = original_remove_dit - model_cache_cls.remove_vae = original_remove_vae + if original_remove_dit is not None: + model_cache_cls.remove_dit = original_remove_dit + if original_remove_vae is not None: + model_cache_cls.remove_vae = original_remove_vae + with _SEEDVR2_PATCH_CONDITION: _SEEDVR2_PATCHING_CLASS_IDS.discard(class_id) _SEEDVR2_PATCHED_CLASS_IDS.discard(class_id) _SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None) _SEEDVR2_PATCH_CONDITION.notify_all() raise - with _SEEDVR2_PATCH_CONDITION: - _SEEDVR2_PATCHING_CLASS_IDS.discard(class_id) - _SEEDVR2_PATCHED_CLASS_IDS.add(class_id) - _SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None) - _SEEDVR2_PATCH_CONDITION.notify_all() - _LOG.info("GPU Resident Loader: integrated external SeedVR2 cache eviction hooks") + _LOG.info("GPU Resident Loader: integrated external SeedVR2 cache visibility hooks") return True