Added override device nodes for clip and vae
This commit is contained in:
@@ -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]",
|
||||
}
|
||||
@@ -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 ''}"
|
||||
|
||||
Reference in New Issue
Block a user