Merge pull request #4 from ArtVentureX/refactor
Refactor load_controlnet to use original loading function & some minor fixes
This commit is contained in:
+27
-161
@@ -1,22 +1,8 @@
|
||||
import sys
|
||||
import os
|
||||
|
||||
|
||||
import torch
|
||||
import contextlib
|
||||
import copy
|
||||
import inspect
|
||||
|
||||
from ldm.modules.diffusionmodules.util import timestep_embedding
|
||||
|
||||
from comfy.cldm import cldm
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.controlnet import ControlBase, ControlNet, T2IAdapter, broadcast_image_to, ControlLora
|
||||
import comfy.t2i_adapter as t2i_adapter
|
||||
import comfy.utils
|
||||
import comfy.model_management as model_management
|
||||
import comfy.model_detection as model_detection
|
||||
import comfy.controlnet as comfy_cn
|
||||
from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to
|
||||
|
||||
ControlNetWeightsType = list[float]
|
||||
T2IAdapterWeightsType = list[float]
|
||||
@@ -243,37 +229,36 @@ class ControlNetAdvanced(ControlNet):
|
||||
mapped_indeces[actual] = i
|
||||
for keyframe in current_timestep_keyframe.latent_keyframes:
|
||||
real_index = keyframe.batch_index
|
||||
# if negative, count from end
|
||||
if real_index < 0:
|
||||
real_index += latent_count if self.sub_idxs is None else self.full_latent_length
|
||||
|
||||
# if not mapping indeces, what you see is what you get
|
||||
if mapped_indeces is None:
|
||||
if real_index in indeces_to_zero:
|
||||
indeces_to_zero.remove(keyframe.batch_index)
|
||||
indeces_to_zero.remove(real_index)
|
||||
# otherwise, see if batch_index is even included in this set of latents
|
||||
else:
|
||||
real_index = mapped_indeces.get(keyframe.batch_index, None)
|
||||
real_index = mapped_indeces.get(real_index, None)
|
||||
if real_index is None:
|
||||
continue
|
||||
indeces_to_zero.remove(real_index)
|
||||
|
||||
|
||||
# apply strength for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+real_index] *= keyframe.strength
|
||||
x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength
|
||||
|
||||
# zero them out by multiplying by zero
|
||||
for batch_index in indeces_to_zero:
|
||||
# apply zero for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+batch_index] *= 0.0
|
||||
x[(latent_count*b)+batch_index] = 0.0
|
||||
|
||||
def copy(self):
|
||||
c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, 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_wrapped)
|
||||
return out
|
||||
|
||||
def cleanup(self):
|
||||
super().cleanup()
|
||||
self.sub_idxs = None
|
||||
@@ -333,141 +318,22 @@ class T2IAdapterAdvanced(T2IAdapter):
|
||||
|
||||
|
||||
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
|
||||
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
if "lora_controlnet" in controlnet_data:
|
||||
return ControlLora(controlnet_data) # TODO: apply weights to ControlLora
|
||||
def load_t2i_adapter(t2i_data):
|
||||
adapter = comfy_cn.load_t2i_adapter(t2i_data)
|
||||
return T2IAdapterAdvanced(adapter.t2i_model, timestep_keyframe, adapter.channels_in)
|
||||
|
||||
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.unet_config_from_diffusers_unet(controlnet_data, use_fp16)
|
||||
diffusers_keys = comfy.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"
|
||||
# override load_t2i_adapter
|
||||
original_load_t2i_adapter = comfy_cn.load_t2i_adapter
|
||||
comfy_cn.load_t2i_adapter = load_t2i_adapter
|
||||
|
||||
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
|
||||
try:
|
||||
control = comfy_cn.load_controlnet(ckpt_path, model=model)
|
||||
if isinstance(control, T2IAdapterAdvanced):
|
||||
return control
|
||||
|
||||
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)
|
||||
|
||||
leftover_keys = controlnet_data.keys()
|
||||
if len(leftover_keys) > 0:
|
||||
print("leftover keys:", leftover_keys)
|
||||
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, timestep_keyframe)
|
||||
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"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
|
||||
control_model = cldm.ControlNet(**controlnet_config)
|
||||
|
||||
if pth:
|
||||
if 'difference' in controlnet_data:
|
||||
if model is not None:
|
||||
model_management.load_models_gpu([model])
|
||||
model_sd = model.model_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)
|
||||
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, timestep_keyframe, global_average_pooling=global_average_pooling)
|
||||
return control
|
||||
|
||||
|
||||
def load_t2i_adapter(t2i_data, timestep_keyframes: TimestepKeyframeGroup=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 = t2i_adapter.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
|
||||
xl = False
|
||||
if cin == 256 or cin == 768:
|
||||
xl = True
|
||||
model_ad = t2i_adapter.adapter.Adapter(cin=cin, channels=[channel, channel*2, channel*4, channel*4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv, xl=xl)
|
||||
else:
|
||||
return None
|
||||
missing, unexpected = model_ad.load_state_dict(t2i_data)
|
||||
if len(missing) > 0:
|
||||
print("t2i missing", missing)
|
||||
|
||||
if len(unexpected) > 0:
|
||||
print("t2i unexpected", unexpected)
|
||||
|
||||
return T2IAdapterAdvanced(model_ad, timestep_keyframes, model_ad.input_channels)
|
||||
return ControlNetAdvanced(control.control_model, timestep_keyframe, global_average_pooling=control.global_average_pooling)
|
||||
except:
|
||||
raise
|
||||
finally:
|
||||
# restore original load_t2i_adapter
|
||||
comfy_cn.load_t2i_adapter = original_load_t2i_adapter
|
||||
|
||||
@@ -327,6 +327,71 @@ class LatentKeyframeGroupNode:
|
||||
return (curr_latent_keyframe,)
|
||||
|
||||
|
||||
class LatentKeyframeTimingNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"batch_index_from": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}),
|
||||
"batch_index_to": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}),
|
||||
"strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.00001}, ),
|
||||
"strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.00001}, ),
|
||||
"timming": (["linear", "ease-in", "ease-out", "ease-in-out"], ),
|
||||
"flip_weights": ([False, True], ),
|
||||
},
|
||||
"optional": {
|
||||
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT_KEYFRAME", )
|
||||
FUNCTION = "load_keyframe"
|
||||
CATEGORY = "adv-controlnet/keyframes"
|
||||
|
||||
def load_keyframe(self,
|
||||
batch_index_from: int,
|
||||
strength_from: float,
|
||||
batch_index_to: int,
|
||||
strength_to: float,
|
||||
timming: str,
|
||||
flip_weights: bool,
|
||||
prev_latent_keyframe: LatentKeyframeGroup=None):
|
||||
|
||||
if (batch_index_from > batch_index_to):
|
||||
raise ValueError("batch_index_from must be less than or equal to batch_index_to.")
|
||||
|
||||
if (batch_index_from < 0 and batch_index_to >= 0):
|
||||
raise ValueError("batch_index_from and batch_index_to must be either both positive or both negative.")
|
||||
|
||||
if (strength_to < strength_from):
|
||||
raise ValueError("strength_to must be greater than or equal to strength_from.")
|
||||
|
||||
if not prev_latent_keyframe:
|
||||
prev_latent_keyframe = LatentKeyframeGroup()
|
||||
|
||||
steps = batch_index_to - batch_index_from + 1
|
||||
diff = strength_to - strength_from
|
||||
if timming == "linear":
|
||||
weights = np.linspace(strength_from, strength_to, steps)
|
||||
elif timming == "ease-in":
|
||||
index = np.linspace(0, 1, steps)
|
||||
weights = diff * np.power(index, 2) + strength_from
|
||||
elif timming == "ease-out":
|
||||
index = np.linspace(0, 1, steps)
|
||||
weights = diff * (1 - np.power(1 - index, 2)) + strength_from
|
||||
elif timming == "ease-in-out":
|
||||
index = np.linspace(0, 1, steps)
|
||||
weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from
|
||||
|
||||
if flip_weights:
|
||||
weights = np.flip(weights)
|
||||
|
||||
for i in range(steps):
|
||||
keyframe = LatentKeyframe(batch_index_from + i, float(weights[i]))
|
||||
print("keyframe", batch_index_from + i, ":", weights[i])
|
||||
prev_latent_keyframe.add(keyframe)
|
||||
|
||||
return (prev_latent_keyframe,)
|
||||
|
||||
|
||||
class ControlNetLoaderAdvanced:
|
||||
@@ -509,6 +574,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"TimestepKeyframe": TimestepKeyframeNode,
|
||||
"LatentKeyframe": LatentKeyframeNode,
|
||||
"LatentKeyframeGroup": LatentKeyframeGroupNode,
|
||||
"LatentKeyframeTiming": LatentKeyframeTimingNode,
|
||||
# Loaders
|
||||
"ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
|
||||
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
|
||||
@@ -527,6 +593,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"TimestepKeyframe": "Timestep Keyframe",
|
||||
"LatentKeyframe": "Latent Keyframe",
|
||||
"LatentKeyframeGroup": "Latent Keyframe Group",
|
||||
"LatentKeyframeTiming": "Latent Keyframe Timing",
|
||||
# Loaders
|
||||
"ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced)",
|
||||
"DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced)",
|
||||
|
||||
Reference in New Issue
Block a user