From 7e96bf676c3b0d65f66ce744eefcf7665174b149 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Fri, 2 Aug 2024 15:27:46 +0200 Subject: [PATCH] Add force offload node --- __init__.py | 4 +++ utils/nodes.py | 11 +++++++ utils/offload.py | 77 ++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 92 insertions(+) create mode 100644 utils/nodes.py create mode 100644 utils/offload.py diff --git a/__init__.py b/__init__.py index b348dc6..b4260bb 100644 --- a/__init__.py +++ b/__init__.py @@ -34,5 +34,9 @@ else: from .MiaoBi.nodes import NODE_CLASS_MAPPINGS as MiaoBi_Nodes NODE_CLASS_MAPPINGS.update(MiaoBi_Nodes) + # Extra + from .utils.nodes import NODE_CLASS_MAPPINGS as Extra_Nodes + NODE_CLASS_MAPPINGS.update(Extra_Nodes) + NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/utils/nodes.py b/utils/nodes.py new file mode 100644 index 0000000..c7e2b63 --- /dev/null +++ b/utils/nodes.py @@ -0,0 +1,11 @@ +NODE_CLASS_MAPPINGS = {} + +from .offload import NODE_CLASS_MAPPINGS as Offload_Nodes +NODE_CLASS_MAPPINGS.update(Offload_Nodes) + +for name, node in NODE_CLASS_MAPPINGS.items(): + cat = node.CATEGORY + if not cat.startswith("ExtraModels/"): + node.CATEGORY = f"ExtraModels/{cat}" + +__all__ = ["NODE_CLASS_MAPPINGS"] diff --git a/utils/offload.py b/utils/offload.py new file mode 100644 index 0000000..b7bf203 --- /dev/null +++ b/utils/offload.py @@ -0,0 +1,77 @@ +# +# Force model to always use specified device +# City96 [Apache2] +# +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 = "other" + + def override(self, model, model_attr, device): + # set model/patcher attributes + 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) + + # move model to device + py_model = getattr(model, model_attr) + py_model.to = types.MethodType(torch.nn.Module.to, py_model) + py_model.to(device) + + # remove ability to move model + 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",) + TITLE = "Force/Set CLIP Device" + + 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",) + TITLE = "Force/Set VAE Device" + + def patch(self, vae, device): + return self.override(vae, "first_stage_model", torch.device(device)) + + +NODE_CLASS_MAPPINGS = { + "OverrideCLIPDevice": OverrideCLIPDevice, + "OverrideVAEDevice": OverrideVAEDevice, +} +NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()}