Fix package imports and patch guards
This commit is contained in:
+2
-9
@@ -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,
|
||||
)
|
||||
|
||||
+20
-4
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
+70
-43
@@ -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
|
||||
|
||||
+14
-3
@@ -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
|
||||
|
||||
+1
-1
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user