From bafd34d48bd4f994186906575ae3ba96bb0e721c Mon Sep 17 00:00:00 2001 From: xmarre Date: Thu, 16 Apr 2026 06:30:29 +0200 Subject: [PATCH 1/3] Route sticky inpaint VAE encodes through tiled entrypoint --- patches.py | 107 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 107 insertions(+) diff --git a/patches.py b/patches.py index 1ae91c1..77b8be0 100644 --- a/patches.py +++ b/patches.py @@ -994,6 +994,96 @@ 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(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]: + @functools.wraps(func) + def wrapper(self, vae, pixels, mask, grow_mask_by=6): + if REGISTRY.get_policy() != "sticky_gpu": + return func(self, vae, pixels, mask, grow_mask_by=grow_mask_by) + with _temporary_prefer_tiled_vae_encode(vae): + return func(self, vae, pixels, mask, grow_mask_by=grow_mask_by) + + 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 +1531,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 +1587,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( From 43fce69abde693c9a66325a024ea308d34aa406b Mon Sep 17 00:00:00 2001 From: "coderabbitai[bot]" <136622811+coderabbitai[bot]@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:37:55 +0000 Subject: [PATCH 2/3] fix: apply CodeRabbit auto-fixes Fixed 1 file(s) based on 1 unresolved review comment. Co-authored-by: CodeRabbit --- patches.py | 20 ++++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/patches.py b/patches.py index 77b8be0..2c48810 100644 --- a/patches.py +++ b/patches.py @@ -1063,12 +1063,24 @@ def _temporary_prefer_tiled_vae_encode(vae: Any): 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, vae, pixels, mask, grow_mask_by=6): + 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, vae, pixels, mask, grow_mask_by=grow_mask_by) + 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, vae, pixels, mask, grow_mask_by=grow_mask_by) + return func(self, *args, **filtered_kwargs) return wrapper @@ -1627,4 +1639,4 @@ def install_patches() -> None: REGISTRY.refresh_runtime_state() _PATCHED = True - _LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths") + _LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths") \ No newline at end of file From e345e5f9d050a12c9c94156d1821d3376aee958e Mon Sep 17 00:00:00 2001 From: xmarre Date: Thu, 16 Apr 2026 06:46:49 +0200 Subject: [PATCH 3/3] Avoid caching bound tiled VAE methods --- patches.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/patches.py b/patches.py index 2c48810..dfbbb68 100644 --- a/patches.py +++ b/patches.py @@ -1024,7 +1024,7 @@ def _should_prefer_tiled_vae_encode(vae: Any, pixel_samples: Any) -> bool: def _call_bound_tiled_vae(func: Callable[..., Any], pixel_samples: Any, *args: Any, **kwargs: Any) -> Any: - supported_kwargs = _tiled_vae_supported_kwargs(func) + 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) @@ -1639,4 +1639,4 @@ def install_patches() -> None: REGISTRY.refresh_runtime_state() _PATCHED = True - _LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths") \ No newline at end of file + _LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths")