diff --git a/control.py b/control.py index 0732eb3..0bf49fc 100644 --- a/control.py +++ b/control.py @@ -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 diff --git a/nodes.py b/nodes.py index 2d5ba18..51ddead 100644 --- a/nodes.py +++ b/nodes.py @@ -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)",