Compare commits

..
Author SHA1 Message Date
xmarre f1519f8da3 Merge pull request #21 from xmarre/codex/disable-load-trim-highvram
Disable eager loader trim in sticky/highvram workflows
2026-04-17 10:08:00 +02:00
coderabbitai[bot] 50e8188e5a 📝 Add docstrings to codex/disable-load-trim-highvram
Docstrings generation was requested by @xmarre.

The following files were modified:

* `kj_loader.py`
2026-04-17 08:05:39 +00:00
xmarre 32dee8e647 Disable eager loader trim for sticky/highvram modes 2026-04-17 10:00:30 +02:00
xmarre 9d5d5e3b42 Merge pull request #20 from xmarre/codex/pr20-external-trim-fix
Fix external trim gate for sticky VAE
2026-04-16 08:34:26 +02:00
xmarre 026d8527f8 Make VAE preflight trim respect external opt-in 2026-04-16 08:26:30 +02:00
xmarre 84d10add70 Add second-chance external VAE trim 2026-04-16 08:17:59 +02:00
xmarre 29b0a34865 Fix external trim gate for sticky VAE 2026-04-16 08:13:41 +02:00
xmarre fd9f33f17f Merge pull request #19 from xmarre/codex/vae-preflight-model-load-fix
Budget sticky VAE preflight for model load
2026-04-16 07:37:53 +02:00
xmarre fc7c0f6946 Tighten sticky VAE load preflight accounting 2026-04-16 07:32:50 +02:00
xmarre d56a716418 Budget sticky VAE preflight for model load 2026-04-16 07:25:31 +02:00
xmarre 27aa644775 Merge pull request #18 from xmarre/codex/inpaint-vae-node-fallback
Handle sticky tiled fallback for inpaint VAE encodes
2026-04-16 06:50:48 +02:00
xmarre e345e5f9d0 Avoid caching bound tiled VAE methods 2026-04-16 06:46:49 +02:00
coderabbitai[bot]andCodeRabbit 43fce69abd fix: apply CodeRabbit auto-fixes
Fixed 1 file(s) based on 1 unresolved review comment.

Co-authored-by: CodeRabbit <noreply@coderabbit.ai>
2026-04-16 04:37:55 +00:00
xmarre bafd34d48b Route sticky inpaint VAE encodes through tiled entrypoint 2026-04-16 06:30:29 +02:00
xmarre cf6843fd70 Merge pull request #17 from xmarre/codex/face-detailer-tiled-vae-admission
Wrap tiled VAE memory admission under sticky GPU
2026-04-16 05:41:26 +02:00
xmarre 5c623b9cdb Fix tiled VAE wrapper review regressions 2026-04-16 05:35:27 +02:00
xmarre bd249781b3 Wrap tiled VAE memory admission under sticky GPU 2026-04-16 05:21:56 +02:00
xmarre 1d1fa53828 Merge pull request #16 from xmarre/codex/external-fallback-free-memory
[codex] Fix external free-memory fallback trim
2026-04-16 04:04:34 +02:00
xmarre 6342d4c3ad Protect related external fallback entries 2026-04-16 03:57:10 +02:00
xmarre f60a8f9972 Skip dynamic external fallback trim 2026-04-16 03:54:13 +02:00
coderabbitai[bot] 61ed3e16d4 📝 Add docstrings to codex/external-fallback-free-memory
Docstrings generation was requested by @xmarre.

The following files were modified:

* `cleanup.py`
* `patches.py`
2026-04-16 01:47:19 +00:00
xmarre d827213bb2 Fix external free-memory fallback trim 2026-04-16 03:41:44 +02:00
xmarre 8998be1d78 Merge pull request #15 from xmarre/codex/fix-vae-preload-boundary
Fix sticky VAE preflight at the preload boundary
2026-04-16 03:12:08 +02:00
xmarre 9d34f008ab Fix VAE preload boundary preflight 2026-04-16 03:06:11 +02:00
xmarre 1f282ea6f0 Merge pull request #14 from xmarre/codex/auto-vae-preflight
[codex] Add sticky VAE preflight tiling
2026-04-15 21:47:18 +02:00
4 changed files with 797 additions and 26 deletions
+55 -2
View File
@@ -130,7 +130,27 @@ def _trim_candidates(
respect_sticky: bool,
sticky_floor_priority: int,
keep_models: tuple[Any, ...],
include_external: bool | None = None,
) -> list[tuple[Any, Any, bool, bool]]:
"""
Collects and returns eviction/unload candidates from in-memory models and, optionally, external integrations.
Filters currently loaded models by the given device, excludes dead or missing models, and skips models listed in `keep_models`. If `respect_sticky` is true, entries whose registry metadata mark them as sticky and whose priority meets or exceeds `sticky_floor_priority` are flagged so they are treated as higher-priority to keep. When `include_external` is true (or when `include_external` is None and external trimming is enabled), candidates from the external registry are included.
Parameters:
device: Device filter for candidates; if not None only candidates matching this device are considered.
respect_sticky (bool): Whether to respect sticky registry entries when computing candidate priority.
sticky_floor_priority (int): Minimum priority value for a registry entry to be considered sticky.
keep_models (tuple[Any, ...]): Objects that must not be selected as candidates.
include_external (bool | None): If True include external-registry candidates; if False exclude them; if None defer to runtime external_trim_enabled().
Returns:
list[tuple[Any, Any, bool, bool]]: A sorted list of tuples (candidate_obj, registry_entry_or_None, sticky_respected, is_external_candidate).
- candidate_obj: The loaded model object (internal) or the external object.
- registry_entry_or_None: Registry metadata for the candidate, or None if unavailable.
- sticky_respected: `True` when the candidate is marked sticky and meets `sticky_floor_priority`.
- is_external_candidate: `True` for candidates originating from the external registry.
"""
candidates: list[tuple[Any, Any, bool, bool]] = []
ensure_external_integrations_installed()
for loaded in list(model_management.current_loaded_models):
@@ -154,7 +174,8 @@ def _trim_candidates(
)
candidates.append((loaded, entry, sticky_respected, False))
if external_trim_enabled():
should_include_external = external_trim_enabled() if include_external is None else bool(include_external)
if should_include_external:
for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates(
device=device,
respect_sticky=respect_sticky,
@@ -175,7 +196,38 @@ def trim_resident_vram(
sticky_floor_priority: int,
allow_partial_unload: bool,
keep_models: tuple[Any, ...] = (),
include_external: bool | None = None,
) -> dict[str, Any]:
"""
Trim resident models (and optional external integration objects) until the requested amount of free VRAM is available or a stopping condition occurs.
Tries to free GPU memory on `device` by evicting or unloading loaded models (and, optionally, external candidates). Respects sticky/priority hints, can perform partial unloads of pinned RAM when supported, and records each attempted action in the returned report.
Parameters:
device (str | torch.device | None): Target device to free (defaults to model_management.get_torch_device()).
target_free_vram_bytes (int): Desired amount of free VRAM, in bytes.
respect_sticky (bool): If true, prefer protecting entries marked as sticky with sufficient priority.
sticky_floor_priority (int): Minimum priority required for a sticky entry to be respected.
allow_partial_unload (bool): If true, allow partial unloads and attempts to free pinned host RAM before full eviction.
keep_models (tuple[Any, ...]): Sequence of model objects that must not be unloaded (exact matches or recognized clones).
include_external (bool | None): If None, use external_trim_enabled() at runtime; otherwise force inclusion/exclusion of external candidates.
Returns:
dict[str, Any]: Report of the trimming operation containing:
- status: "met_target", "partial", or "error".
- stopped_reason: reason the loop stopped (e.g., "target_met", "no_candidates", "no_progress", "error").
- target_met (bool): whether the target free VRAM was reached.
- device (str): string form of the device used.
- target_free_vram_bytes (int), free_before_bytes (int), free_after_bytes (int).
- freed_vram_bytes (int): total freed VRAM during this call.
- respect_sticky (bool), sticky_floor_priority (int), allow_partial_unload (bool).
- external_trim_enabled (bool): computed flag indicating whether external candidates were considered.
- actions (list[dict]): ordered per-candidate action records; each entry includes metadata such as
entry_id, basename, tracked, external_candidate, sticky_respected, priority,
need_before_bytes, loaded_before_bytes, freed_pinned_ram_bytes,
mode (e.g., "full_unload", "partial_unload", "external_evict", "error"),
loaded_after_bytes, freed_vram_bytes, free_after_bytes, and any warnings/errors.
"""
ensure_external_integrations_installed()
cleanup_models_gc = getattr(model_management, "cleanup_models_gc", None)
if callable(cleanup_models_gc):
@@ -203,6 +255,7 @@ def trim_resident_vram(
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
keep_models=keep_models,
include_external=include_external,
)
if not candidates:
stopped_reason = "no_candidates"
@@ -306,7 +359,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()),
"external_trim_enabled": bool(external_trim_enabled() if include_external is None else include_external),
"actions": actions,
}
+31 -1
View File
@@ -403,6 +403,36 @@ def external_trim_enabled() -> bool:
return value in {"1", "true", "yes", "on"}
def external_objects_for_models(models: tuple[Any, ...] | list[Any]) -> tuple[Any, ...]:
"""
Returns registered external objects that are part of any supplied model wrapper chain.
This lets callers preserve external cache entries when they are associated with a kept
model through wrapper indirection rather than exact object identity.
"""
related_ids: set[int] = set()
for model in models:
for related in _iter_seedvr2_wrapper_chain(model):
related_ids.add(id(related))
if not related_ids:
return ()
EXTERNAL_REGISTRY.refresh_runtime_state()
matches: list[Any] = []
seen_ids: set[int] = set()
with EXTERNAL_REGISTRY._lock:
for entry in EXTERNAL_REGISTRY._entries.values():
obj = entry.object()
if obj is None:
continue
obj_id = id(obj)
if obj_id in related_ids and obj_id not in seen_ids:
matches.append(obj)
seen_ids.add(obj_id)
return tuple(matches)
def _call_seedvr2_method_with_optional_expected_model(
method: Callable[..., Any],
*args: Any,
@@ -693,4 +723,4 @@ def ensure_external_integrations_installed() -> None:
"GPU Resident Loader: failed to install SeedVR2 external integration from %s: %s",
normalized_file,
exc,
)
)
+129
View File
@@ -395,6 +395,21 @@ def _estimate_model_load_bytes(
def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> int:
"""
Estimate the byte size of a specific checkpoint auxiliary component (e.g., CLIP or VAE).
If `ckpt_path` is a safetensors file this attempts a component-specific estimate via
`estimate_checkpoint_component_bytes`. If that returns `None` or the file is not
safetensors, the function falls back to the file size on disk.
Parameters:
ckpt_path (str): Path to the checkpoint file.
kind (str): Component kind to estimate (for example `"clip"`, `"vae"`, or other
checkpoint component identifiers accepted by `estimate_checkpoint_component_bytes`).
Returns:
int: Estimated number of bytes required by the requested component.
"""
estimated = None
if _is_safetensors_path(ckpt_path):
estimated = estimate_checkpoint_component_bytes(ckpt_path, kind)
@@ -403,6 +418,46 @@ def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> in
return int(estimated)
def _env_bool(name: str) -> bool | None:
"""
Parse a boolean-like environment variable value.
Parameters:
name (str): Environment variable name to read.
Returns:
bool | None: `True` if the variable value is one of "1", "true", "yes", or "on";
`False` if it is one of "0", "false", "no", or "off"; `None` if the variable is unset
or contains an unrecognized value.
"""
value = os.environ.get(name, "").strip().lower()
if value in {"1", "true", "yes", "on"}:
return True
if value in {"0", "false", "no", "off"}:
return False
return None
def _should_trim_before_load(*, effective_policy: str) -> bool:
"""
Decides whether to perform adaptive VRAM trimming before loading based on environment overrides, the effective residency policy, and runtime flags.
Parameters:
effective_policy (str): The resolved residency policy name to evaluate (e.g., "sticky_gpu").
Returns:
True if trimming should run before load, False otherwise.
"""
env_override = _env_bool("COMFYUI_GPU_RESIDENT_LOAD_TRIM")
if env_override is not None:
return env_override
if effective_policy == "sticky_gpu":
return False
if getattr(args, "highvram", False) or getattr(args, "gpu_only", False):
return False
return True
def _maybe_trim_before_load(
*,
loader_name: str,
@@ -410,7 +465,26 @@ def _maybe_trim_before_load(
explicit_device: torch.device | None,
required_bytes: int,
keep_models: tuple[Any, ...] = (),
enabled: bool = True,
) -> None:
"""
Attempt to free GPU VRAM proactively to create headroom for a forthcoming load.
If trimming is enabled and an explicit CUDA device is provided and `required_bytes` > 0,
calls the adaptive VRAM trimmer to free memory while preserving any `keep_models`.
Logs an informational message when memory was freed and a warning if the trimmer
could not reach the target headroom.
Parameters:
loader_name (str): Short name used in log messages for the entity requesting the trim.
reason (str): Human-readable reason for the trim (included in logs).
explicit_device (torch.device | None): The explicit device to trim on; trimming is skipped if `None` or not CUDA.
required_bytes (int): Estimated number of bytes needed for the upcoming load; trimming is skipped if <= 0.
keep_models (tuple[Any, ...], optional): Objects to preserve from eviction while trimming (defaults to ()).
enabled (bool, optional): If `False`, the function is a no-op (defaults to True).
"""
if not enabled:
return
if explicit_device is None or explicit_device.type != "cuda":
return
if required_bytes <= 0:
@@ -601,6 +675,26 @@ def _load_resident_diffusion_model(
policy_override: str | None = None,
keep_models: tuple[Any, ...] = (),
) -> Any:
"""
Load a diffusion model state dict from `source_path`, applying residency-aware VRAM trimming, optional extra UNet state merging, backend flags, and post-load model patches, and bind or reuse the resulting model for future loads.
Parameters:
loader_name (str): Human-readable loader identifier used for logging.
cache_scope (str): Logical cache scope used when constructing the loader key.
source_path (str): Filesystem path to the diffusion model checkpoint.
note (str): Short note describing the load context (used when binding).
weight_dtype (str): Weight dtype selection used to build model options.
compute_dtype (str): Compute dtype selection applied to the loaded model.
patch_cublaslinear (bool): Whether to enable the cublas linear optimization during load.
sage_attention (str): SageAttention mode to apply to the model after loading.
enable_fp16_accumulation (bool): If true, enable fp16 accumulation backend flag during load.
extra_state_dict (str | None): Optional path to an additional state dict whose matching UNet keys will be merged into the main state dict.
policy_override (str | None): Optional residency policy override name (affects trimming decision).
keep_models (tuple[Any, ...]): Iterable of already-loaded objects whose residency should be preserved during trimming.
Returns:
The loaded diffusion model object.
"""
model_options = _build_model_options(weight_dtype)
effective_policy = _effective_policy_name(policy_override)
loader_key = _make_loader_key(
@@ -632,6 +726,7 @@ def _load_resident_diffusion_model(
extra_state_dict=extra_state_dict,
),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with _temporary_backend_flags(
@@ -697,6 +792,20 @@ def _load_checkpoint_clip_only(
loader_name: str,
keep_models: tuple[Any, ...] = (),
):
"""
Load or reuse the CLIP (text encoder) component from a checkpoint file.
This resolves a loader cache key, attempts to reuse a previously loaded CLIP, and if absent loads the checkpoint (safetensors or torch format), optionally applies model-config-based extraction or older-quant conversion, constructs a CLIP object when weights are present, and binds the result for reuse.
Parameters:
ckpt_path (str): Filesystem path to the checkpoint file.
policy_override (str | None): Optional residency policy override used to form the loader key.
loader_name (str): Human-readable name used in log messages and trimming decisions.
keep_models (tuple[Any, ...]): Sequence of model-like objects to preserve during any pre-load VRAM trimming.
Returns:
clip (comfy.sd.CLIP | None): A constructed CLIP/text-encoder instance when weights are available, or `None` if no CLIP weights were found.
"""
loader_key = _checkpoint_component_loader_key("clip", policy_override)
reused_clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=loader_key)
if reused_clip is not None:
@@ -707,6 +816,7 @@ def _load_checkpoint_clip_only(
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
effective_policy = _effective_policy_name(policy_override)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=ckpt_path)
_maybe_trim_before_load(
loader_name=loader_name,
@@ -714,6 +824,7 @@ def _load_checkpoint_clip_only(
explicit_device=explicit_device,
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_CLIP),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with REGISTRY.load_context(kind=KIND_CLIP, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
if is_safetensors and model_config is None:
@@ -783,6 +894,22 @@ def _load_checkpoint_vae_only(
loader_name: str,
keep_models: tuple[Any, ...] = (),
):
"""
Load only the VAE component from a checkpoint, reusing a live VAE when available and binding the loaded VAE for reuse.
Parameters:
ckpt_path (str): Filesystem path to the checkpoint (safetensors or torch file).
policy_override (str | None): Optional residency policy override used to resolve trimming/loading behavior.
loader_name (str): Human-readable name used in logging messages.
keep_models (tuple[Any, ...]): Objects to preserve from eviction when freeing VRAM before loading.
Returns:
vae (comfy.sd.VAE | None): The loaded VAE object, or `None` if no VAE weights were found in the checkpoint.
Side effects:
- May trigger adaptive VRAM trimming before loading depending on the effective policy and runtime flags.
- Binds the resulting VAE (if any) into the live-object registry for reuse.
"""
loader_key = _checkpoint_component_loader_key("vae", policy_override)
reused_vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=loader_key)
if reused_vae is not None:
@@ -793,6 +920,7 @@ def _load_checkpoint_vae_only(
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
effective_policy = _effective_policy_name(policy_override)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_VAE, source_path=ckpt_path)
_maybe_trim_before_load(
loader_name=loader_name,
@@ -800,6 +928,7 @@ def _load_checkpoint_vae_only(
explicit_device=explicit_device,
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_VAE),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with REGISTRY.load_context(kind=KIND_VAE, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
if is_safetensors and model_config is None:
+582 -23
View File
@@ -2,10 +2,12 @@ from __future__ import annotations
import contextlib
import functools
import inspect
import json
import logging
import os
import struct
import threading
from typing import Any, Callable
import torch
@@ -18,6 +20,7 @@ from .cleanup import (
trim_resident_vram,
unload_loaded_model,
)
from .external_residency import EXTERNAL_REGISTRY, external_objects_for_models, external_trim_enabled
from .residency import (
KIND_CHECKPOINT,
KIND_CLIP,
@@ -38,6 +41,8 @@ _WARNED_PICKLE_GPU_PATHS: set[str] = set()
_SAFE_TENSORS_COMPONENT_CACHE_MAX = 32
_STICKY_PROTECTION_VRAM_FLOOR_RATIO = 0.125
_STICKY_PROTECTION_VRAM_FLOOR_CEIL_BYTES = 16 * 1024 ** 3
_TILED_VAE_MEMORY_LOCK_ATTR = "_gpu_resident_loader_tiled_memory_lock"
_TILED_VAE_LOCK_INIT = threading.Lock()
_SAFETENSORS_DTYPE_MAP = {
"BOOL": torch.bool,
"U8": torch.uint8,
@@ -668,26 +673,96 @@ def _scaled_batch_memory(total_memory_used: int, total_batch_count: int, batch_n
return max(1, (total_memory * current_batch + total_batches - 1) // total_batches)
def _sticky_vae_free_memory(*, device: Any, patcher: Any) -> int:
import comfy.model_management as model_management
get_free_memory = getattr(model_management, "get_free_memory", None)
if callable(get_free_memory):
try:
return max(0, int(get_free_memory(device)))
except Exception:
pass
return max(0, int(patcher.get_free_memory(device)))
def _sticky_model_load_requirement(*, device: Any, models: tuple[Any, ...]) -> int | None:
import comfy.model_management as model_management
loaded_model_cls = getattr(model_management, "LoadedModel", None)
is_device_cpu = getattr(model_management, "is_device_cpu", None)
if device is None or loaded_model_cls is None:
return None
if callable(is_device_cpu) and is_device_cpu(device):
return 0
requested_models: set[Any] = set()
for model in models:
if model is None:
continue
requested_models.add(model)
additional_models = getattr(model, "model_patches_models", None)
if callable(additional_models):
try:
for additional in additional_models():
if additional is not None:
requested_models.add(additional)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: failed to enumerate patched models for sticky VAE preflight: %s",
exc,
)
return None
total_required = 0
for model in requested_models:
try:
loaded_model = loaded_model_cls(model)
loaded_device = getattr(loaded_model, "device", None)
if loaded_device != device:
continue
if callable(is_device_cpu) and is_device_cpu(loaded_device):
continue
total_required += max(0, int(loaded_model.model_memory_required(loaded_model.device)))
except Exception:
return None
return total_required
def _sticky_model_load_target(model_load_required: int | None) -> int | None:
if model_load_required is None:
return None
required = max(0, int(model_load_required))
return (required * 11 + 9) // 10
def _prepare_sticky_vae_batch(
*,
device: Any,
patcher: Any,
total_memory_used: int,
total_batch_count: int,
) -> tuple[int, bool]:
free_memory = int(patcher.get_free_memory(device))
) -> tuple[int, int, bool]:
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
model_load_required = _sticky_model_load_requirement(device=device, models=(patcher,))
model_load_target = _sticky_model_load_target(model_load_required)
batch_budget = 0 if model_load_target is None else max(0, free_memory - model_load_target)
batch_number = _sticky_safe_batch_number(
batch_count=total_batch_count,
free_memory=free_memory,
free_memory=batch_budget,
memory_used=total_memory_used,
device=device,
)
batch_memory_used = _scaled_batch_memory(total_memory_used, total_batch_count, batch_number)
if REGISTRY.get_policy() != "sticky_gpu" or device is None:
return batch_number, False
return batch_number, batch_memory_used, False
target_free = _sticky_protection_target(batch_memory_used, device)
target_free = (
free_memory + 1
if model_load_target is None
else model_load_target + _sticky_protection_target(batch_memory_used, device)
)
if free_memory < target_free:
try:
trim_resident_vram(
@@ -697,29 +772,65 @@ def _prepare_sticky_vae_batch(
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=(patcher,),
include_external=False,
)
except Exception as exc:
_LOG.debug("GPU Resident Loader: proactive VAE trim failed for %s bytes: %s", batch_memory_used, exc)
_LOG.debug(
"GPU Resident Loader: proactive VAE trim failed for batch=%s load=%s bytes: %s",
batch_memory_used,
model_load_required,
exc,
)
free_memory = int(patcher.get_free_memory(device))
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
if free_memory < target_free and external_trim_enabled():
try:
trim_resident_vram(
device=device,
target_free_vram_bytes=target_free,
respect_sticky=True,
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=(patcher,),
include_external=True,
)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: second-chance external VAE trim failed for batch=%s load=%s bytes: %s",
batch_memory_used,
model_load_required,
exc,
)
else:
_LOG.info(
"GPU Resident Loader: native sticky VAE trim was insufficient; retried with external candidates enabled"
)
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
batch_budget = 0 if model_load_target is None else max(0, free_memory - model_load_target)
batch_number = _sticky_safe_batch_number(
batch_count=total_batch_count,
free_memory=free_memory,
free_memory=batch_budget,
memory_used=total_memory_used,
device=device,
)
batch_memory_used = _scaled_batch_memory(total_memory_used, total_batch_count, batch_number)
target_free = _sticky_protection_target(batch_memory_used, device)
target_free = (
free_memory + 1
if model_load_target is None
else model_load_target + _sticky_protection_target(batch_memory_used, device)
)
should_tile = free_memory < target_free and batch_number <= 1
if should_tile:
_LOG.info(
"GPU Resident Loader: skipping regular VAE pass and switching directly to tiled mode; free=%s target=%s batch_memory=%s",
"GPU Resident Loader: skipping regular VAE pass and switching directly to tiled mode; free=%s target=%s batch_memory=%s model_load=%s",
free_memory,
target_free,
batch_memory_used,
model_load_required,
)
return batch_number, should_tile
return batch_number, batch_memory_used, should_tile
def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
@@ -741,8 +852,7 @@ def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
pixel_samples = pixel_samples.unsqueeze(2)
try:
memory_used = self.memory_used_encode(pixel_samples.shape, self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload)
batch_number, should_tile = _prepare_sticky_vae_batch(
batch_number, batch_memory_used, should_tile = _prepare_sticky_vae_batch(
device=self.device,
patcher=self.patcher,
total_memory_used=memory_used,
@@ -751,6 +861,11 @@ def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
if should_tile:
do_tile = True
else:
model_management.load_models_gpu(
[self.patcher],
memory_required=batch_memory_used,
force_full_load=self.disable_offload,
)
samples = None
for x in range(0, pixel_samples.shape[0], batch_number):
pixels_in = self.process_input(pixel_samples[x:x + batch_number]).to(self.vae_dtype)
@@ -788,6 +903,364 @@ def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
return wrapper
def _default_tiled_vae_axes(
*,
latent_dim: int,
extra_1d_channel: Any,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
decode: bool,
) -> tuple[int | None, int | None, int | None]:
if latent_dim == 3:
default_tile_x = 32 if decode else 512
default_tile_y = 32 if decode else 512
default_tile_t = 999 if decode else 9999
elif latent_dim == 1 or extra_1d_channel is not None:
default_tile_x = 256 * 2048
default_tile_y = None
default_tile_t = None
else:
default_tile_x = 64 if decode else 512
default_tile_y = 64 if decode else 512
default_tile_t = None
resolved_tile_x = default_tile_x if tile_x is None else max(1, int(tile_x))
resolved_tile_y = default_tile_y if tile_y is None else max(1, int(tile_y))
resolved_tile_t = default_tile_t if tile_t is None else max(1, int(tile_t))
return resolved_tile_x, resolved_tile_y, resolved_tile_t
def _shape_with_capped_tail(shape: tuple[int, ...], tail_caps: dict[int, int | None]) -> tuple[int, ...]:
capped = list(shape)
for index, cap in tail_caps.items():
if cap is None:
continue
capped[index] = min(int(capped[index]), max(1, int(cap)))
return tuple(capped)
def _tiled_vae_memory_shapes(
*,
shape: tuple[int, ...],
latent_dim: int,
extra_1d_channel: Any,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
decode: bool,
) -> list[tuple[int, ...]]:
resolved_tile_x, resolved_tile_y, resolved_tile_t = _default_tiled_vae_axes(
latent_dim=latent_dim,
extra_1d_channel=extra_1d_channel,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
decode=decode,
)
if latent_dim == 3:
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 3: resolved_tile_t,
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
)
]
if latent_dim == 1 or extra_1d_channel is not None:
return [_shape_with_capped_tail(shape, {len(shape) - 1: resolved_tile_x})]
if decode:
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
)
]
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
),
_shape_with_capped_tail(
shape,
{
len(shape) - 2: max(1, resolved_tile_y // 2),
len(shape) - 1: max(1, resolved_tile_x * 2),
},
),
_shape_with_capped_tail(
shape,
{
len(shape) - 2: max(1, resolved_tile_y * 2),
len(shape) - 1: max(1, resolved_tile_x // 2),
},
),
]
@contextlib.contextmanager
def _temporary_tiled_vae_memory_estimate(
self,
*,
decode: bool,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
) -> Any:
memory_attr = "memory_used_decode" if decode else "memory_used_encode"
original = getattr(self, memory_attr, None)
if not callable(original):
yield
return
had_instance_attr = memory_attr in getattr(self, "__dict__", {})
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
with _TILED_VAE_LOCK_INIT:
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
lock = threading.RLock()
setattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, lock)
def estimated(shape, dtype, *args, **kwargs):
shapes = _tiled_vae_memory_shapes(
shape=tuple(int(dim) for dim in shape),
latent_dim=int(getattr(self, "latent_dim", 2)),
extra_1d_channel=getattr(self, "extra_1d_channel", None),
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
decode=decode,
)
return max(int(original(candidate, dtype, *args, **kwargs)) for candidate in shapes)
with lock:
setattr(self, memory_attr, estimated)
try:
yield
finally:
if had_instance_attr:
setattr(self, memory_attr, original)
else:
delattr(self, memory_attr)
@functools.lru_cache(maxsize=None)
def _tiled_vae_supported_kwargs(func: Callable[..., Any]) -> frozenset[str]:
return frozenset(inspect.signature(func).parameters)
def _call_tiled_vae(
func: Callable[..., Any],
self,
data,
*,
tile_x=None,
tile_y=None,
overlap=None,
tile_t=None,
overlap_t=None,
):
kwargs = {}
supported_kwargs = _tiled_vae_supported_kwargs(func)
if "tile_x" in supported_kwargs:
kwargs["tile_x"] = tile_x
if "tile_y" in supported_kwargs:
kwargs["tile_y"] = tile_y
if "overlap" in supported_kwargs:
kwargs["overlap"] = overlap
if "tile_t" in supported_kwargs:
kwargs["tile_t"] = tile_t
if "overlap_t" in supported_kwargs:
kwargs["overlap_t"] = overlap_t
return func(self, data, **kwargs)
def _should_prefer_tiled_vae_encode(vae: Any, pixel_samples: Any) -> bool:
if REGISTRY.get_policy() != "sticky_gpu":
return False
if vae is None or pixel_samples is None:
return False
try:
vae.throw_exception_if_invalid()
prepared = vae.vae_encode_crop_pixels(pixel_samples)
prepared = prepared.movedim(-1, 1)
if int(getattr(vae, "latent_dim", 2)) == 3 and prepared.ndim < 5:
if not getattr(vae, "not_video", False):
prepared = prepared.movedim(1, 0).unsqueeze(0)
else:
prepared = prepared.unsqueeze(2)
memory_used = vae.memory_used_encode(prepared.shape, vae.vae_dtype)
_, _, should_tile = _prepare_sticky_vae_batch(
device=getattr(vae, "device", None),
patcher=getattr(vae, "patcher", None),
total_memory_used=memory_used,
total_batch_count=prepared.shape[0],
)
return bool(should_tile)
except Exception as exc:
_LOG.debug("GPU Resident Loader: failed to preflight sticky VAE encode preference: %s", exc)
return False
def _call_bound_tiled_vae(func: Callable[..., Any], pixel_samples: Any, *args: Any, **kwargs: Any) -> Any:
supported_kwargs = _tiled_vae_supported_kwargs(getattr(func, "__func__", func))
filtered_kwargs = {key: value for key, value in kwargs.items() if key in supported_kwargs}
return func(pixel_samples, *args, **filtered_kwargs)
@contextlib.contextmanager
def _temporary_prefer_tiled_vae_encode(vae: Any):
original_encode = getattr(vae, "encode", None)
encode_tiled = getattr(vae, "encode_tiled", None)
if not callable(original_encode) or not callable(encode_tiled):
yield
return
had_instance_attr = "encode" in getattr(vae, "__dict__", {})
lock = getattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
with _TILED_VAE_LOCK_INIT:
lock = getattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
lock = threading.RLock()
setattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, lock)
def prefer_encode(pixel_samples, *args, **kwargs):
if _should_prefer_tiled_vae_encode(vae, pixel_samples):
return _call_bound_tiled_vae(encode_tiled, pixel_samples, *args, **kwargs)
return original_encode(pixel_samples, *args, **kwargs)
with lock:
setattr(vae, "encode", prefer_encode)
try:
yield
finally:
if had_instance_attr:
setattr(vae, "encode", original_encode)
else:
delattr(vae, "encode")
def _wrap_vae_encode_for_inpaint_node(func: Callable[..., Any]) -> Callable[..., Any]:
supported_kwargs = frozenset(inspect.signature(func).parameters)
@functools.wraps(func)
def wrapper(self, *args, **kwargs):
# Filter kwargs to only include those supported by the wrapped function
filtered_kwargs = {key: value for key, value in kwargs.items() if key in supported_kwargs}
# Supply default grow_mask_by if not present and supported
if "grow_mask_by" in supported_kwargs and "grow_mask_by" not in filtered_kwargs:
filtered_kwargs["grow_mask_by"] = 6
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, *args, **filtered_kwargs)
# Extract vae from args for the context manager
vae = args[0] if args else kwargs.get("vae")
with _temporary_prefer_tiled_vae_encode(vae):
return func(self, *args, **filtered_kwargs)
return wrapper
def _wrap_inpaint_model_conditioning_node(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, positive, negative, pixels, vae, mask, noise_mask=True):
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, positive, negative, pixels, vae, mask, noise_mask=noise_mask)
with _temporary_prefer_tiled_vae_encode(vae):
return func(self, positive, negative, pixels, vae, mask, noise_mask=noise_mask)
return wrapper
def _wrap_vae_encode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
decode=False,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
):
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper
def _wrap_vae_decode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
decode=True,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
):
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper
def _wrap_vae_decode(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, samples_in, vae_options={}):
@@ -803,8 +1276,7 @@ def _wrap_vae_decode(func: Callable[..., Any]) -> Callable[..., Any]:
samples_in = samples_in[:, :, 0]
try:
memory_used = self.memory_used_decode(samples_in.shape, self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload)
batch_number, should_tile = _prepare_sticky_vae_batch(
batch_number, batch_memory_used, should_tile = _prepare_sticky_vae_batch(
device=self.device,
patcher=self.patcher,
total_memory_used=memory_used,
@@ -812,17 +1284,21 @@ def _wrap_vae_decode(func: Callable[..., Any]) -> Callable[..., Any]:
)
preallocated = False
if not should_tile and getattr(self.first_stage_model, "comfy_has_chunked_io", False):
pixel_samples = torch.empty(
self.first_stage_model.decode_output_shape(samples_in.shape),
device=self.output_device,
dtype=self.vae_output_dtype(),
)
preallocated = True
if should_tile:
do_tile = True
else:
model_management.load_models_gpu(
[self.patcher],
memory_required=batch_memory_used,
force_full_load=self.disable_offload,
)
if getattr(self.first_stage_model, "comfy_has_chunked_io", False):
pixel_samples = torch.empty(
self.first_stage_model.decode_output_shape(samples_in.shape),
device=self.output_device,
dtype=self.vae_output_dtype(),
)
preallocated = True
for x in range(0, samples_in.shape[0], batch_number):
samples = samples_in[x:x + batch_number].to(device=self.device, dtype=self.vae_dtype)
if preallocated:
@@ -972,6 +1448,28 @@ def _wrap_model_patcher_detach(func: Callable[..., Any]) -> Callable[..., Any]:
def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
"""
Wraps a free-memory function to enforce sticky-GPU protection and external fallback trimming.
When the registry policy is "sticky_gpu" and a device is provided, the wrapper:
- Reserves VRAM for sticky-loaded models by attempting a pre-trim to a computed protection target.
- Protects a subset of sticky-loaded wrappers from unloading when calling the original function by adding them to `keep_loaded`.
- After the original free-memory call, refreshes both REGISTRY and EXTERNAL_REGISTRY runtime state.
- If external trimming is available and still needed, attempts a fallback trim that includes external residency.
Dynamic free-memory calls are excluded because Comfy reduces their effective target internally.
Parameters:
memory_required: Number of bytes the caller needs to free.
device: Target device for which memory is being freed (may be None).
keep_loaded: Iterable of loaded-wrapper objects that must be kept; the wrapper may extend this list with additional protected wrappers.
Returns:
The value returned by the wrapped `func`.
Notes:
- The wrapper may call `trim_resident_vram` and `model_management.get_free_memory`; exceptions from trimming or free-memory queries are caught and logged, not propagated.
- Side effects include invoking trims and refreshing runtime state on REGISTRY and EXTERNAL_REGISTRY.
"""
@functools.wraps(func)
def wrapper(memory_required, device, keep_loaded=None, *args, **kwargs):
import comfy.model_management as model_management
@@ -1026,6 +1524,42 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
unloaded = func(memory_required, device, keep_loaded + protected_wrappers, *args, **kwargs)
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
for_dynamic = bool(kwargs.get("for_dynamic", args[0] if args else False))
if device is not None and external_trim_enabled() and not for_dynamic:
fallback_target = memory_required
if REGISTRY.get_policy() == "sticky_gpu":
fallback_target = max(fallback_target, _sticky_protection_target(memory_required, device))
try:
free_now = model_management.get_free_memory(device)
except Exception:
free_now = None
if free_now is not None and int(free_now) < int(fallback_target):
protected_models = tuple(
model
for model in (getattr(loaded_wrapper, "model", None) for loaded_wrapper in keep_loaded + protected_wrappers)
if model is not None
)
keep_models = protected_models + external_objects_for_models(protected_models)
try:
trim_resident_vram(
device=device,
target_free_vram_bytes=int(fallback_target),
respect_sticky=True,
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=keep_models,
include_external=True,
)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: external fallback trim failed for free_memory(%s): %s",
memory_required,
exc,
)
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
return unloaded
return wrapper
@@ -1102,6 +1636,7 @@ def install_patches() -> None:
import comfy.model_patcher as model_patcher
import comfy.sd as comfy_sd
import comfy.utils as comfy_utils
import nodes as comfy_nodes
original_load_torch_file = _remember_original("utils.load_torch_file", comfy_utils.load_torch_file)
if comfy_utils.load_torch_file is original_load_torch_file:
@@ -1149,6 +1684,30 @@ def install_patches() -> None:
if comfy_sd.VAE.decode is original_vae_decode:
comfy_sd.VAE.decode = _wrap_vae_decode(original_vae_decode)
original_vae_encode_tiled = _remember_original("sd.VAE.encode_tiled", comfy_sd.VAE.encode_tiled)
if comfy_sd.VAE.encode_tiled is original_vae_encode_tiled:
comfy_sd.VAE.encode_tiled = _wrap_vae_encode_tiled(original_vae_encode_tiled)
original_vae_decode_tiled = _remember_original("sd.VAE.decode_tiled", comfy_sd.VAE.decode_tiled)
if comfy_sd.VAE.decode_tiled is original_vae_decode_tiled:
comfy_sd.VAE.decode_tiled = _wrap_vae_decode_tiled(original_vae_decode_tiled)
if hasattr(comfy_nodes, "VAEEncodeForInpaint") and hasattr(comfy_nodes.VAEEncodeForInpaint, "encode"):
original_vae_encode_for_inpaint = _remember_original(
"nodes.VAEEncodeForInpaint.encode",
comfy_nodes.VAEEncodeForInpaint.encode,
)
if comfy_nodes.VAEEncodeForInpaint.encode is original_vae_encode_for_inpaint:
comfy_nodes.VAEEncodeForInpaint.encode = _wrap_vae_encode_for_inpaint_node(original_vae_encode_for_inpaint)
if hasattr(comfy_nodes, "InpaintModelConditioning") and hasattr(comfy_nodes.InpaintModelConditioning, "encode"):
original_inpaint_model_conditioning = _remember_original(
"nodes.InpaintModelConditioning.encode",
comfy_nodes.InpaintModelConditioning.encode,
)
if comfy_nodes.InpaintModelConditioning.encode is original_inpaint_model_conditioning:
comfy_nodes.InpaintModelConditioning.encode = _wrap_inpaint_model_conditioning_node(original_inpaint_model_conditioning)
original_clip_vision_load = _remember_original("clip_vision.load", clip_vision.load)
if clip_vision.load is original_clip_vision_load:
clip_vision.load = _wrap_with_load_context(