Merge pull request #18 from xmarre/codex/inpaint-vae-node-fallback

Handle sticky tiled fallback for inpaint VAE encodes
This commit is contained in:
xmarre
2026-04-16 06:50:48 +02:00
committed by GitHub
+119
View File
@@ -994,6 +994,108 @@ def _call_tiled_vae(
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):
@@ -1441,6 +1543,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:
@@ -1496,6 +1599,22 @@ def install_patches() -> None:
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(