Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f1519f8da3 | ||
|
|
50e8188e5a | ||
|
|
32dee8e647 | ||
|
|
9d5d5e3b42 | ||
|
|
026d8527f8 | ||
|
|
84d10add70 | ||
|
|
29b0a34865 | ||
|
|
fd9f33f17f | ||
|
|
fc7c0f6946 | ||
|
|
d56a716418 | ||
|
|
27aa644775 | ||
|
|
e345e5f9d0 | ||
|
|
43fce69abd | ||
|
|
bafd34d48b | ||
|
|
cf6843fd70 | ||
|
|
5c623b9cdb | ||
|
|
bd249781b3 | ||
|
|
1d1fa53828 |
+129
@@ -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:
|
||||
|
||||
+487
-7
@@ -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
|
||||
@@ -39,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,
|
||||
@@ -682,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,
|
||||
@@ -690,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,
|
||||
)
|
||||
@@ -701,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(
|
||||
@@ -711,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
|
||||
|
||||
@@ -806,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={}):
|
||||
@@ -1072,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))
|
||||
@@ -1181,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:
|
||||
@@ -1228,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(
|
||||
|
||||
Reference in New Issue
Block a user