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
2 changed files with 229 additions and 7 deletions
+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:
+100 -7
View File
@@ -686,6 +686,56 @@ def _sticky_vae_free_memory(*, device: Any, patcher: Any) -> int:
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,
@@ -694,9 +744,12 @@ def _prepare_sticky_vae_batch(
total_batch_count: int,
) -> 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,
)
@@ -705,7 +758,11 @@ def _prepare_sticky_vae_batch(
if REGISTRY.get_policy() != "sticky_gpu" or device is None:
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(
@@ -715,27 +772,63 @@ 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 = _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, batch_memory_used, should_tile
@@ -1434,7 +1527,7 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
EXTERNAL_REGISTRY.refresh_runtime_state()
for_dynamic = bool(kwargs.get("for_dynamic", args[0] if args else False))
if device is not None and not external_trim_enabled() and not for_dynamic:
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))