Added advanced ControlNet loaders, added customizable ControlNet weight nodes
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
+303
@@ -0,0 +1,303 @@
|
||||
import sys
|
||||
import os
|
||||
|
||||
|
||||
import torch
|
||||
import contextlib
|
||||
import copy
|
||||
import inspect
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
|
||||
|
||||
from comfy.cldm import cldm
|
||||
from comfy.t2i_adapter import adapter
|
||||
|
||||
from comfy.sd import ControlBase, broadcast_image_to
|
||||
import comfy.utils as utils
|
||||
import comfy.model_management as model_management
|
||||
import comfy.model_detection as model_detection
|
||||
|
||||
|
||||
ControlNetWeightsType = list[float]
|
||||
T2IAdapterWeightsType = list[float]
|
||||
|
||||
|
||||
# Copied from comfy.sd, weights modified
|
||||
class ControlNetAdvanced(ControlBase):
|
||||
def __init__(self, control_model, weights: ControlNetWeightsType, global_average_pooling=False, device=None):
|
||||
super().__init__(device)
|
||||
self.control_model = control_model
|
||||
self.weights = weights
|
||||
self.global_average_pooling = global_average_pooling
|
||||
|
||||
def get_control(self, x_noisy, t, cond, batched_number):
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
||||
|
||||
if self.timestep_range is not None:
|
||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||
if control_prev is not None:
|
||||
return control_prev
|
||||
else:
|
||||
return {}
|
||||
|
||||
output_dtype = x_noisy.dtype
|
||||
if self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
|
||||
if self.cond_hint is not None:
|
||||
del self.cond_hint
|
||||
self.cond_hint = None
|
||||
self.cond_hint = utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
|
||||
if x_noisy.shape[0] != self.cond_hint.shape[0]:
|
||||
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
|
||||
|
||||
if self.control_model.dtype == torch.float16:
|
||||
precision_scope = torch.autocast
|
||||
else:
|
||||
precision_scope = contextlib.nullcontext
|
||||
|
||||
with precision_scope(model_management.get_autocast_device(self.device)):
|
||||
self.control_model = model_management.load_if_low_vram(self.control_model)
|
||||
context = torch.cat(cond['c_crossattn'], 1)
|
||||
y = cond.get('c_adm', None)
|
||||
control = self.control_model(x=x_noisy, hint=self.cond_hint, timesteps=t, context=context, y=y)
|
||||
self.control_model = model_management.unload_if_low_vram(self.control_model)
|
||||
out = {'middle':[], 'output': []}
|
||||
autocast_enabled = torch.is_autocast_enabled()
|
||||
|
||||
for i in range(len(control)):
|
||||
if i == (len(control) - 1):
|
||||
key = 'middle'
|
||||
index = 0
|
||||
else:
|
||||
key = 'output'
|
||||
index = i
|
||||
x = control[i]
|
||||
if self.global_average_pooling:
|
||||
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
|
||||
|
||||
#multiplier = 1#0.825**float(12-i)
|
||||
#print(f"$$$ multiplier: {multiplier}")
|
||||
x *= self.strength*self.weights[i]
|
||||
if x.dtype != output_dtype and not autocast_enabled:
|
||||
x = x.to(output_dtype)
|
||||
|
||||
if control_prev is not None and key in control_prev:
|
||||
prev = control_prev[key][index]
|
||||
if prev is not None:
|
||||
x += prev
|
||||
out[key].append(x)
|
||||
if control_prev is not None and 'input' in control_prev:
|
||||
out['input'] = control_prev['input']
|
||||
return out
|
||||
|
||||
def copy(self):
|
||||
c = ControlNetAdvanced(self.control_model, self.weights, global_average_pooling=self.global_average_pooling)
|
||||
self.copy_to(c)
|
||||
return c
|
||||
|
||||
def get_models(self):
|
||||
out = super().get_models()
|
||||
out.append(self.control_model)
|
||||
return out
|
||||
|
||||
|
||||
def load_controlnet(ckpt_path, control_net_weights: ControlNetWeightsType=None, t2i_adapter_weights: T2IAdapterWeightsType=None, model=None):
|
||||
controlnet_data = utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
|
||||
controlnet_config = None
|
||||
if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: #diffusers format
|
||||
use_fp16 = model_management.should_use_fp16()
|
||||
controlnet_config = model_detection.model_config_from_diffusers_unet(controlnet_data, use_fp16).unet_config
|
||||
diffusers_keys = utils.unet_to_diffusers(controlnet_config)
|
||||
diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"
|
||||
diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"
|
||||
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
k_in = "controlnet_down_blocks.{}{}".format(count, s)
|
||||
k_out = "zero_convs.{}.0{}".format(count, s)
|
||||
if k_in not in controlnet_data:
|
||||
loop = False
|
||||
break
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
if count == 0:
|
||||
k_in = "controlnet_cond_embedding.conv_in{}".format(s)
|
||||
else:
|
||||
k_in = "controlnet_cond_embedding.blocks.{}{}".format(count - 1, s)
|
||||
k_out = "input_hint_block.{}{}".format(count * 2, s)
|
||||
if k_in not in controlnet_data:
|
||||
k_in = "controlnet_cond_embedding.conv_out{}".format(s)
|
||||
loop = False
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
|
||||
new_sd = {}
|
||||
for k in diffusers_keys:
|
||||
if k in controlnet_data:
|
||||
new_sd[diffusers_keys[k]] = controlnet_data.pop(k)
|
||||
|
||||
controlnet_data = new_sd
|
||||
|
||||
pth_key = 'control_model.zero_convs.0.0.weight'
|
||||
pth = False
|
||||
key = 'zero_convs.0.0.weight'
|
||||
if pth_key in controlnet_data:
|
||||
pth = True
|
||||
key = pth_key
|
||||
prefix = "control_model."
|
||||
elif key in controlnet_data:
|
||||
prefix = ""
|
||||
else:
|
||||
net = load_t2i_adapter(controlnet_data, t2i_adapter_weights)
|
||||
if net is None:
|
||||
print("error checkpoint does not contain controlnet or t2i adapter data", ckpt_path)
|
||||
return net
|
||||
|
||||
if controlnet_config is None:
|
||||
use_fp16 = model_management.should_use_fp16()
|
||||
controlnet_config = model_detection.model_config_from_unet(controlnet_data, prefix, use_fp16).unet_config
|
||||
controlnet_config.pop("out_channels")
|
||||
controlnet_config["hint_channels"] = 3
|
||||
control_model = cldm.ControlNet(**controlnet_config)
|
||||
|
||||
if pth:
|
||||
if 'difference' in controlnet_data:
|
||||
if model is not None:
|
||||
m = model.patch_model()
|
||||
model_sd = m.state_dict()
|
||||
for x in controlnet_data:
|
||||
c_m = "control_model."
|
||||
if x.startswith(c_m):
|
||||
sd_key = "diffusion_model.{}".format(x[len(c_m):])
|
||||
if sd_key in model_sd:
|
||||
cd = controlnet_data[x]
|
||||
cd += model_sd[sd_key].type(cd.dtype).to(cd.device)
|
||||
model.unpatch_model()
|
||||
else:
|
||||
print("WARNING: Loaded a diff controlnet without a model. It will very likely not work.")
|
||||
|
||||
class WeightsLoader(torch.nn.Module):
|
||||
pass
|
||||
w = WeightsLoader()
|
||||
w.control_model = control_model
|
||||
missing, unexpected = w.load_state_dict(controlnet_data, strict=False)
|
||||
else:
|
||||
missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False)
|
||||
print(missing, unexpected)
|
||||
|
||||
if use_fp16:
|
||||
control_model = control_model.half()
|
||||
|
||||
global_average_pooling = False
|
||||
if ckpt_path.endswith("_shuffle.pth") or ckpt_path.endswith("_shuffle.safetensors") or ckpt_path.endswith("_shuffle_fp16.safetensors"): #TODO: smarter way of enabling global_average_pooling
|
||||
global_average_pooling = True
|
||||
|
||||
control = ControlNetAdvanced(control_model, control_net_weights, global_average_pooling=global_average_pooling)
|
||||
return control
|
||||
|
||||
|
||||
# Copied from comfy.sd, weights modified
|
||||
class T2IAdapterAdvanced(ControlBase):
|
||||
def __init__(self, t2i_model, weights: T2IAdapterWeightsType, channels_in, device=None):
|
||||
super().__init__(device)
|
||||
self.t2i_model = t2i_model
|
||||
self.weights = weights if weights else [1.0]*4
|
||||
self.channels_in = channels_in
|
||||
self.control_input = None
|
||||
|
||||
def get_control(self, x_noisy, t, cond, batched_number):
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
||||
|
||||
if self.timestep_range is not None:
|
||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||
if control_prev is not None:
|
||||
return control_prev
|
||||
else:
|
||||
return {}
|
||||
|
||||
if self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
|
||||
if self.cond_hint is not None:
|
||||
del self.cond_hint
|
||||
self.control_input = None
|
||||
self.cond_hint = None
|
||||
self.cond_hint = utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").float().to(self.device)
|
||||
if self.channels_in == 1 and self.cond_hint.shape[1] > 1:
|
||||
self.cond_hint = torch.mean(self.cond_hint, 1, keepdim=True)
|
||||
if x_noisy.shape[0] != self.cond_hint.shape[0]:
|
||||
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
|
||||
if self.control_input is None:
|
||||
self.t2i_model.to(self.device)
|
||||
self.control_input = self.t2i_model(self.cond_hint)
|
||||
self.t2i_model.cpu()
|
||||
|
||||
output_dtype = x_noisy.dtype
|
||||
out = {'input':[]}
|
||||
|
||||
autocast_enabled = torch.is_autocast_enabled()
|
||||
print(f"$$$$ t2i control_input len: {len(self.control_input)}")
|
||||
for i in range(len(self.control_input)):
|
||||
key = 'input'
|
||||
x = self.control_input[i] * self.strength
|
||||
if x.dtype != output_dtype and not autocast_enabled:
|
||||
x = x.to(output_dtype)
|
||||
|
||||
if control_prev is not None and key in control_prev:
|
||||
index = len(control_prev[key]) - i * 3 - 3
|
||||
prev = control_prev[key][index]
|
||||
if prev is not None:
|
||||
x += prev
|
||||
out[key].insert(0, None)
|
||||
out[key].insert(0, None)
|
||||
out[key].insert(0, x)
|
||||
|
||||
if control_prev is not None and 'input' in control_prev:
|
||||
for i in range(len(out['input'])):
|
||||
if out['input'][i] is None:
|
||||
out['input'][i] = control_prev['input'][i]
|
||||
if control_prev is not None and 'middle' in control_prev:
|
||||
out['middle'] = control_prev['middle']
|
||||
if control_prev is not None and 'output' in control_prev:
|
||||
out['output'] = control_prev['output']
|
||||
return out
|
||||
|
||||
def copy(self):
|
||||
c = T2IAdapterAdvanced(self.t2i_model, self.weights, self.channels_in)
|
||||
self.copy_to(c)
|
||||
return c
|
||||
|
||||
|
||||
def load_t2i_adapter(t2i_data, weights: T2IAdapterWeightsType=None):
|
||||
keys = t2i_data.keys()
|
||||
if 'adapter' in keys:
|
||||
t2i_data = t2i_data['adapter']
|
||||
keys = t2i_data.keys()
|
||||
if "body.0.in_conv.weight" in keys:
|
||||
cin = t2i_data['body.0.in_conv.weight'].shape[1]
|
||||
model_ad = adapter.Adapter_light(cin=cin, channels=[320, 640, 1280, 1280], nums_rb=4)
|
||||
elif 'conv_in.weight' in keys:
|
||||
cin = t2i_data['conv_in.weight'].shape[1]
|
||||
channel = t2i_data['conv_in.weight'].shape[0]
|
||||
ksize = t2i_data['body.0.block2.weight'].shape[2]
|
||||
use_conv = False
|
||||
down_opts = list(filter(lambda a: a.endswith("down_opt.op.weight"), keys))
|
||||
if len(down_opts) > 0:
|
||||
use_conv = True
|
||||
model_ad = adapter.Adapter(cin=cin, channels=[channel, channel*2, channel*4, channel*4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv)
|
||||
else:
|
||||
return None
|
||||
model_ad.load_state_dict(t2i_data)
|
||||
return T2IAdapterAdvanced(model_ad, weights, cin // 64)
|
||||
@@ -0,0 +1,228 @@
|
||||
import sys
|
||||
import os
|
||||
|
||||
import folder_paths
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
|
||||
from comfy.sd import ControlBase
|
||||
|
||||
from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType
|
||||
|
||||
|
||||
class ScaledSoftControlNetWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ([False, True], ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "controlnet weights"
|
||||
|
||||
def load_weights(self, base_multiplier, flip_weights):
|
||||
weights = [(base_multiplier ** float(12 - i)) for i in range(13)]
|
||||
if flip_weights:
|
||||
weights.reverse()
|
||||
return (weights, )
|
||||
|
||||
|
||||
class SoftControlNetWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ([False, True], ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "controlnet weights"
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12]
|
||||
if flip_weights:
|
||||
weights.reverse()
|
||||
return (weights,)
|
||||
|
||||
|
||||
class CustomControlNetWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ([False, True], ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "controlnet weights"
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12]
|
||||
if flip_weights:
|
||||
weights.reverse()
|
||||
return (weights,)
|
||||
|
||||
|
||||
class SoftT2IAdapterWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ([False, True], ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", )
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "controlnet weights"
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
if flip_weights:
|
||||
weights.reverse()
|
||||
return (weights,)
|
||||
|
||||
|
||||
class CustomT2IAdapterWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ([False, True], ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", )
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "controlnet weights"
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
if flip_weights:
|
||||
weights.reverse()
|
||||
return (weights,)
|
||||
|
||||
|
||||
class ControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
},
|
||||
"optional": {
|
||||
"control_net_weights": ("CONTROL_NET_WEIGHTS", ),
|
||||
"t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "loaders"
|
||||
|
||||
def load_controlnet(self, control_net_name, control_net_weights: ControlNetWeightsType=None, t2i_adapter_weights: T2IAdapterWeightsType=None):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, control_net_weights, t2i_adapter_weights)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class DiffControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), )
|
||||
},
|
||||
"optional": {
|
||||
"control_net_weights": ("CONTROL_NET_WEIGHTS", ),
|
||||
"t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "loaders"
|
||||
|
||||
def load_controlnet(self, control_net_name, control_net_weights, t2i_adapter_weights, model):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, control_net_weights, t2i_adapter_weights, model)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
# NODE MAPPING
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# Loaders
|
||||
"ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
|
||||
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
|
||||
# Weights
|
||||
"ScaledSoftControlNetWeights": ScaledSoftControlNetWeights,
|
||||
"SoftControlNetWeights": SoftControlNetWeights,
|
||||
"CustomControlNetWeights": CustomControlNetWeights,
|
||||
"SoftT2IAdapterWeights": SoftT2IAdapterWeights,
|
||||
"CustomT2IAdapterWeights": CustomT2IAdapterWeights,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Loaders
|
||||
"ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced)",
|
||||
"DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced)",
|
||||
# Weights
|
||||
"ScaledSoftControlNetWeights": "Scaled Soft ControlNet Weights",
|
||||
"SoftControlNetWeights": "Soft ControlNet Weights",
|
||||
"CustomControlNetWeights": "Custom ControlNet Weights",
|
||||
"SoftT2IAdapterWeights": "Soft T2IAdapter Weights",
|
||||
"CustomT2IAdapterWeights": "Custom T2IAdapter Weights",
|
||||
}
|
||||
Reference in New Issue
Block a user