Merge pull request #13 from xmarre/codex/seedvr2-pr12-regression-fix

[codex] restore SeedVR2 cache reuse by default
This commit is contained in:
xmarre
2026-04-15 20:31:42 +02:00
committed by GitHub
3 changed files with 122 additions and 68 deletions
+8 -8
View File
@@ -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
+10 -8
View File
@@ -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,
}
+104 -52
View File
@@ -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