From 669e6014fdac42d3ab61fc7f0ab40f8d5d64d501 Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 04:14:07 +0200 Subject: [PATCH] Fix package imports and patch guards --- __init__.py | 11 +---- kj_loader.py | 24 +++++++++-- nodes.py | 8 +++- patches.py | 113 +++++++++++++++++++++++++++++++-------------------- residency.py | 17 ++++++-- startup.py | 2 +- 6 files changed, 113 insertions(+), 62 deletions(-) diff --git a/__init__.py b/__init__.py index de038e4..d01feb8 100644 --- a/__init__.py +++ b/__init__.py @@ -1,12 +1,5 @@ -import os -import sys - -CURRENT_DIR = os.path.dirname(os.path.abspath(__file__)) -if CURRENT_DIR not in sys.path: - sys.path.insert(0, CURRENT_DIR) - -from startup import install_patches -from nodes import ( +from .startup import install_patches +from .nodes import ( NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS, ) diff --git a/kj_loader.py b/kj_loader.py index faae2d2..b2c19fc 100644 --- a/kj_loader.py +++ b/kj_loader.py @@ -11,7 +11,7 @@ import comfy.utils from comfy.cli_args import PerformanceFeature, args from comfy.ldm.modules.attention import attention_pytorch, wrap_attn -from .residency import KIND_CHECKPOINT, KIND_MODEL, REGISTRY +from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, REGISTRY _LOG = logging.getLogger(__name__) @@ -44,12 +44,28 @@ def _set_cublas_linear(enabled: bool) -> None: def _set_fp16_accumulation(enabled: bool) -> None: + if not hasattr(torch.backends.cuda, "matmul"): + if enabled: + raise RuntimeError( + "Failed to enable fp16 accumulation. This requires a PyTorch build exposing " + "torch.backends.cuda.matmul.allow_fp16_accumulation." + ) + _LOG.warning( + "GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior." + ) + return + flag = getattr(torch.backends.cuda.matmul, "allow_fp16_accumulation", None) if flag is None: - raise RuntimeError( - "Failed to set fp16 accumulation. This requires a PyTorch build exposing " - "torch.backends.cuda.matmul.allow_fp16_accumulation." + if enabled: + raise RuntimeError( + "Failed to enable fp16 accumulation. This requires a PyTorch build exposing " + "torch.backends.cuda.matmul.allow_fp16_accumulation." + ) + _LOG.warning( + "GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior." ) + return torch.backends.cuda.matmul.allow_fp16_accumulation = bool(enabled) diff --git a/nodes.py b/nodes.py index 1d9f778..5289978 100644 --- a/nodes.py +++ b/nodes.py @@ -5,8 +5,8 @@ from typing import Any import comfy.model_management as model_management -from kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident -from residency import REGISTRY +from .kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident +from .residency import REGISTRY def _entry_report_json(obj: Any) -> str: @@ -133,6 +133,8 @@ class PinClipResidency: def pin(self, clip, sticky: bool, priority: int): patcher = _patcher_for_clip(clip) + if patcher is None: + raise RuntimeError("Expected a CLIP object with a patcher, but no patcher was found.") REGISTRY.set_sticky(patcher, sticky=sticky, priority=priority) return (clip,) @@ -154,6 +156,8 @@ class PinVAEResidency: def pin(self, vae, sticky: bool, priority: int): patcher = _patcher_for_vae(vae) + if patcher is None: + raise RuntimeError("Expected a VAE object with a patcher, but no patcher was found.") REGISTRY.set_sticky(patcher, sticky=sticky, priority=priority) return (vae,) diff --git a/patches.py b/patches.py index f987995..138a4ef 100644 --- a/patches.py +++ b/patches.py @@ -8,7 +8,7 @@ from typing import Any, Callable import torch from safetensors import safe_open -from residency import ( +from .residency import ( KIND_CHECKPOINT, KIND_CLIP, KIND_CLIP_VISION, @@ -155,8 +155,15 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata error=str(exc), ) return (sd, metadata) if return_metadata else sd - except Exception: - pass + except Exception as fallback_exc: + _record_generic_load( + path=ckpt, + method="safetensors_cpu_fallback_failed", + requested_device=requested_device, + actual_device="error", + error=str(fallback_exc), + ) + raise fallback_exc from exc if len(getattr(exc, "args", ())) > 0: message = exc.args[0] @@ -315,12 +322,18 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]: return wrapper +def _remember_original(key: str, value: Callable[..., Any]) -> Callable[..., Any]: + return _ORIGINALS.setdefault(key, value) + + def _patch_model_management_devices() -> None: import comfy.model_management as model_management def wrap_device_func(name: str) -> None: - original = getattr(model_management, name) - _ORIGINALS[f"model_management.{name}"] = original + key = f"model_management.{name}" + original = _remember_original(key, getattr(model_management, name)) + if getattr(model_management, name) is not original: + return @functools.wraps(original) def wrapper(*args, **kwargs): @@ -345,11 +358,13 @@ def _patch_model_management_devices() -> None: if hasattr(model_management, name): wrap_device_func(name) - _ORIGINALS["model_management.free_memory"] = model_management.free_memory - model_management.free_memory = _wrap_free_memory(model_management.free_memory) + original_free_memory = _remember_original("model_management.free_memory", model_management.free_memory) + if model_management.free_memory is original_free_memory: + model_management.free_memory = _wrap_free_memory(original_free_memory) - _ORIGINALS["model_management.load_models_gpu"] = model_management.load_models_gpu - model_management.load_models_gpu = _wrap_load_models_gpu(model_management.load_models_gpu) + original_load_models_gpu = _remember_original("model_management.load_models_gpu", model_management.load_models_gpu) + if model_management.load_models_gpu is original_load_models_gpu: + model_management.load_models_gpu = _wrap_load_models_gpu(original_load_models_gpu) def install_patches() -> None: @@ -364,49 +379,61 @@ def install_patches() -> None: import comfy.sd as comfy_sd import comfy.utils as comfy_utils - _ORIGINALS["utils.load_torch_file"] = comfy_utils.load_torch_file - comfy_utils.load_torch_file = _patched_load_torch_file + 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: + comfy_utils.load_torch_file = _patched_load_torch_file if hasattr(clip_vision, "load_torch_file"): - clip_vision.load_torch_file = comfy_utils.load_torch_file + original_clip_vision_load_torch_file = _remember_original("clip_vision.load_torch_file", clip_vision.load_torch_file) + if clip_vision.load_torch_file is original_clip_vision_load_torch_file: + clip_vision.load_torch_file = comfy_utils.load_torch_file _patch_model_management_devices() - _ORIGINALS["sd.load_checkpoint_guess_config"] = comfy_sd.load_checkpoint_guess_config - comfy_sd.load_checkpoint_guess_config = _wrap_with_load_context( - KIND_CHECKPOINT, - path_arg_index=0, - bind_output=_bind_checkpoint_outputs, - )(comfy_sd.load_checkpoint_guess_config) + original_load_checkpoint_guess_config = _remember_original( + "sd.load_checkpoint_guess_config", + comfy_sd.load_checkpoint_guess_config, + ) + if comfy_sd.load_checkpoint_guess_config is original_load_checkpoint_guess_config: + comfy_sd.load_checkpoint_guess_config = _wrap_with_load_context( + KIND_CHECKPOINT, + path_arg_index=0, + bind_output=_bind_checkpoint_outputs, + )(original_load_checkpoint_guess_config) - _ORIGINALS["sd.load_diffusion_model"] = comfy_sd.load_diffusion_model - comfy_sd.load_diffusion_model = _wrap_with_load_context( - KIND_MODEL, - path_arg_index=0, - bind_output=lambda model, source_path: model is not None - and REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL), - )(comfy_sd.load_diffusion_model) + original_load_diffusion_model = _remember_original("sd.load_diffusion_model", comfy_sd.load_diffusion_model) + if comfy_sd.load_diffusion_model is original_load_diffusion_model: + comfy_sd.load_diffusion_model = _wrap_with_load_context( + KIND_MODEL, + path_arg_index=0, + bind_output=lambda model, source_path: model is not None + and REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL), + )(original_load_diffusion_model) - _ORIGINALS["sd.load_clip"] = comfy_sd.load_clip - comfy_sd.load_clip = _wrap_load_clip(comfy_sd.load_clip) + original_load_clip = _remember_original("sd.load_clip", comfy_sd.load_clip) + if comfy_sd.load_clip is original_load_clip: + comfy_sd.load_clip = _wrap_load_clip(original_load_clip) - _ORIGINALS["clip_vision.load"] = clip_vision.load - clip_vision.load = _wrap_with_load_context( - KIND_CLIP_VISION, - path_arg_index=0, - bind_output=lambda result, source_path: result is not None - and getattr(result, "patcher", None) is not None - and REGISTRY.bind_object(result.patcher, source_path=source_path, kind=KIND_CLIP_VISION), - )(clip_vision.load) + 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( + KIND_CLIP_VISION, + path_arg_index=0, + bind_output=lambda result, source_path: result is not None + and getattr(result, "patcher", None) is not None + and REGISTRY.bind_object(result.patcher, source_path=source_path, kind=KIND_CLIP_VISION), + )(original_clip_vision_load) - _ORIGINALS["controlnet.load_controlnet"] = controlnet.load_controlnet - controlnet.load_controlnet = _wrap_with_load_context(KIND_CONTROLNET, path_arg_index=0)(controlnet.load_controlnet) + original_load_controlnet = _remember_original("controlnet.load_controlnet", controlnet.load_controlnet) + if controlnet.load_controlnet is original_load_controlnet: + controlnet.load_controlnet = _wrap_with_load_context(KIND_CONTROLNET, path_arg_index=0)(original_load_controlnet) - _ORIGINALS["diffusers_load.load_diffusers"] = diffusers_load.load_diffusers - diffusers_load.load_diffusers = _wrap_with_load_context( - KIND_CHECKPOINT, - path_arg_index=0, - bind_output=_bind_diffusers_outputs, - )(diffusers_load.load_diffusers) + original_load_diffusers = _remember_original("diffusers_load.load_diffusers", diffusers_load.load_diffusers) + if diffusers_load.load_diffusers is original_load_diffusers: + diffusers_load.load_diffusers = _wrap_with_load_context( + KIND_CHECKPOINT, + path_arg_index=0, + bind_output=_bind_diffusers_outputs, + )(original_load_diffusers) REGISTRY.refresh_runtime_state() _PATCHED = True diff --git a/residency.py b/residency.py index bd58bd9..e814b89 100644 --- a/residency.py +++ b/residency.py @@ -271,10 +271,21 @@ class ResidencyRegistry: self._path_to_entry[(kind, source_path)] = entry_id try: setattr(obj, "__gpu_resident_loader_entry_id__", entry_id) - except Exception: - pass + except (AttributeError, TypeError): + _LOG.debug( + "GPU Resident Loader: could not tag object %r with residency entry id %s", + type(obj), + entry_id, + ) - entry.object_ref = weakref.ref(obj) + try: + entry.object_ref = weakref.ref(obj) + except TypeError: + entry.object_ref = None + _LOG.debug( + "GPU Resident Loader: object %r is not weak-referenceable; tracking metadata only", + type(obj), + ) entry.sticky = entry.sticky if sticky is None else bool(sticky) entry.priority = int(priority) entry.source_path = source_path diff --git a/startup.py b/startup.py index 4bad7a7..145a6bf 100644 --- a/startup.py +++ b/startup.py @@ -1,6 +1,6 @@ import logging -from patches import install_patches as _install_patches +from .patches import install_patches as _install_patches _LOG = logging.getLogger(__name__) _PATCHES_INSTALLED = False