From 032159e5009ad38122b405f1398f73b0653c110b Mon Sep 17 00:00:00 2001 From: Level Pixel Dev Date: Tue, 4 Feb 2025 20:38:06 +0600 Subject: [PATCH] Added override device nodes for clip and vae --- nodes/unloaders/override_device_LP.py | 74 +++++++++++++++++++++++++++ nodes/utils/utils_LP.py | 2 +- 2 files changed, 75 insertions(+), 1 deletion(-) create mode 100644 nodes/unloaders/override_device_LP.py diff --git a/nodes/unloaders/override_device_LP.py b/nodes/unloaders/override_device_LP.py new file mode 100644 index 0000000..af827b0 --- /dev/null +++ b/nodes/unloaders/override_device_LP.py @@ -0,0 +1,74 @@ +import types +import torch +import comfy.model_management + +class OverrideDevice: + @classmethod + def INPUT_TYPES(s): + devices = ["cpu",] + for k in range(0, torch.cuda.device_count()): + devices.append(f"cuda:{k}") + + return { + "required": { + "device": (devices, {"default":"cpu"}), + } + } + + FUNCTION = "patch" + CATEGORY = "LevelPixel/Unloaders" + + def override(self, model, model_attr, device): + model.device = device + patcher = getattr(model, "patcher", model) #.clone() + for name in ["device", "load_device", "offload_device", "current_device", "output_device"]: + setattr(patcher, name, device) + + py_model = getattr(model, model_attr) + py_model.to = types.MethodType(torch.nn.Module.to, py_model) + py_model.to(device) + + def to(*args, **kwargs): + pass + py_model.to = types.MethodType(to, py_model) + return (model,) + + def patch(self, *args, **kwargs): + raise NotImplementedError + +class OverrideCLIPDevice(OverrideDevice): + @classmethod + def INPUT_TYPES(s): + k = super().INPUT_TYPES() + k["required"]["clip"] = ("CLIP",) + return k + + RETURN_TYPES = ("CLIP",) + CATEGORY = "LevelPixel/Unloaders" + + def patch(self, clip, device): + return self.override(clip, "cond_stage_model", torch.device(device)) + +class OverrideVAEDevice(OverrideDevice): + @classmethod + def INPUT_TYPES(s): + k = super().INPUT_TYPES() + k["required"]["vae"] = ("VAE",) + return k + + RETURN_TYPES = ("VAE",) + CATEGORY = "LevelPixel/Unloaders" + + def patch(self, vae, device): + return self.override(vae, "first_stage_model", torch.device(device)) + + +NODE_CLASS_MAPPINGS = { + "OverrideCLIPDevice|LP": OverrideCLIPDevice, + "OverrideVAEDevice|LP": OverrideVAEDevice, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "OverrideCLIPDevice|LP": "Override CLIP Device [LP]", + "OverrideVAEDevice|LP": "Override VAE Device [LP]", +} diff --git a/nodes/utils/utils_LP.py b/nodes/utils/utils_LP.py index d00dcfd..808b1bb 100644 --- a/nodes/utils/utils_LP.py +++ b/nodes/utils/utils_LP.py @@ -23,7 +23,7 @@ class Delay: RETURN_TYPES = (any,) RETURN_NAMES = ("output",) FUNCTION = "add_delay" - CATEGORY = "LevelPixel/Unloaders" + CATEGORY = "LevelPixel/Utils" def add_delay(self, input, delay_seconds): delay_text = f"{delay_seconds:.1f} second{'s' if delay_seconds != 1 else ''}"