From 15e438100fb6f7e946da0112c08d44a78e8e54e2 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 18 Sep 2023 11:23:01 -0500 Subject: [PATCH] Updated to support sliding window usage by other extensions, added scaffolding code to support masks in controlnet, added logger for future use --- control.py | 108 +++++++++++++++++++++--- logger.py | 36 ++++++++ nodes.py | 242 ++++++++++++++++++++++++++--------------------------- 3 files changed, 255 insertions(+), 131 deletions(-) create mode 100644 logger.py diff --git a/control.py b/control.py index 2f20782..0732eb3 100644 --- a/control.py +++ b/control.py @@ -9,14 +9,12 @@ import inspect from ldm.modules.diffusionmodules.util import timestep_embedding -sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy")) - 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 as utils +import comfy.utils import comfy.model_management as model_management import comfy.model_detection as model_detection @@ -174,15 +172,61 @@ class ControlNetAdvanced(ControlNet): self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13 # mask for which parts of controlnet output to keep self.cond_hint_mask = None + # actual index values + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 # override control_merge self.control_merge = control_merge_inject.__get__(self, type(self)) + def set_cond_hint_mask(self, mask_hint): + self.cond_hint_mask = mask_hint + return self + def get_control(self, x_noisy, t, cond, batched_number): # need to reference t and batched_number later self.t = t self.batched_number = batched_number # TODO: choose TimestepKeyframe based on t - return super().get_control(x_noisy, t, cond, batched_number) + if self.sub_idxs is not None: + # perform special version of get_control + return self.sliding_get_control(x_noisy, t, cond, batched_number) + else: + return super().get_control(x_noisy, t, cond, batched_number) + + def sliding_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 None + + output_dtype = x_noisy.dtype + + # TODO: change this to not require cond_hint upscaling every step + if self.sub_idxs is not None or self.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 = comfy.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 self.cond_hint length matches real latent count, need to subdivide it + if self.cond_hint.size(0) == self.full_latent_length: + self.cond_hint = self.cond_hint[self.sub_idxs] + + 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) + + context = cond['c_crossattn'] + y = cond.get('c_adm', None) + if y is not None: + y = y.to(self.control_model.dtype) + control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=t, context=context.to(self.control_model.dtype), y=y) + return self.control_merge(None, control, control_prev, output_dtype) def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): if current_timestep_keyframe.latent_keyframes is not None: @@ -191,12 +235,28 @@ class ControlNetAdvanced(ControlNet): latent_count = x.size(0)//batched_number indeces_to_zero = set(range(latent_count)) + mapped_indeces = None + # if expecting subdivision, will need to translate between subset and actual idx values + if self.sub_idxs: + mapped_indeces = {} + for i, actual in enumerate(self.sub_idxs): + mapped_indeces[actual] = i for keyframe in current_timestep_keyframe.latent_keyframes: - if keyframe.batch_index in indeces_to_zero: - indeces_to_zero.remove(keyframe.batch_index) + real_index = keyframe.batch_index + # 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) + # otherwise, see if batch_index is even included in this set of latents + else: + real_index = mapped_indeces.get(keyframe.batch_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)+keyframe.batch_index] *= keyframe.strength + x[(latent_count*b)+real_index] *= keyframe.strength # zero them out by multiplying by zero for batch_index in indeces_to_zero: @@ -213,6 +273,12 @@ class ControlNetAdvanced(ControlNet): out = super().get_models() out.append(self.control_model_wrapped) return out + + def cleanup(self): + super().cleanup() + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 class T2IAdapterAdvanced(T2IAdapter): @@ -224,6 +290,10 @@ class T2IAdapterAdvanced(T2IAdapter): self.weights = first_weight if first_weight else [1.0]*12 # mask for which parts of controlnet output to keep self.cond_hint_mask = None + # actual index values + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 # override control_merge self.control_merge = control_merge_inject.__get__(self, type(self)) @@ -232,7 +302,19 @@ class T2IAdapterAdvanced(T2IAdapter): self.t = t self.batched_number = batched_number # TODO: choose TimestepKeyframe based on t - return super().get_control(x_noisy, t, cond, batched_number) + try: + # if sub indexes present, replace original hint with subsection + if self.sub_idxs is not None: + full_cond_hint_original = self.cond_hint_original + del self.cond_hint + self.cond_hint = None + self.cond_hint_original = full_cond_hint_original[self.sub_idxs] + return super().get_control(x_noisy, t, cond, batched_number) + finally: + if self.sub_idxs is not None: + # replace original cond hint + self.cond_hint_original = full_cond_hint_original + del full_cond_hint_original def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): # For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters @@ -242,10 +324,16 @@ class T2IAdapterAdvanced(T2IAdapter): c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in) self.copy_to(c) return c + + def cleanup(self): + super().cleanup() + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): - controlnet_data = utils.load_torch_file(ckpt_path, safe_load=True) + 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 @@ -253,7 +341,7 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo 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 = utils.unet_to_diffusers(controlnet_config) + 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" diff --git a/logger.py b/logger.py new file mode 100644 index 0000000..b23b82f --- /dev/null +++ b/logger.py @@ -0,0 +1,36 @@ +import sys +import copy +import logging + + +class ColoredFormatter(logging.Formatter): + COLORS = { + "DEBUG": "\033[0;36m", # CYAN + "INFO": "\033[0;32m", # GREEN + "WARNING": "\033[0;33m", # YELLOW + "ERROR": "\033[0;31m", # RED + "CRITICAL": "\033[0;37;41m", # WHITE ON RED + "RESET": "\033[0m", # RESET COLOR + } + + def format(self, record): + colored_record = copy.copy(record) + levelname = colored_record.levelname + seq = self.COLORS.get(levelname, self.COLORS["RESET"]) + colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + return super().format(colored_record) + + +# Create a new logger +logger = logging.getLogger("Advanced-ControlNet") +logger.propagate = False + +# Add handler if we don't have one. +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(ColoredFormatter("[%(name)s] - %(levelname)s - %(message)s")) + logger.addHandler(handler) + +# Configure logger +loglevel = logging.INFO +logger.setLevel(loglevel) diff --git a/nodes.py b/nodes.py index 3e16e5d..2d5ba18 100644 --- a/nodes.py +++ b/nodes.py @@ -8,11 +8,9 @@ from PIL import Image, ImageOps import folder_paths -sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy")) - -from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ +from .control import ControlNetAdvanced, T2IAdapterAdvanced, load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ LatentKeyframe, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup - +from .logger import logger def get_properly_arranged_t2i_weights(initial_weights: list[float]): new_weights = [] @@ -231,6 +229,106 @@ class LatentKeyframeNode: return (prev_latent_keyframe,) +class LatentKeyframeGroupNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "index_strengths": ("STRING", {"multiline": True, "default": ""}), + }, + "optional": { + "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "latent_image_opt": ("LATENT", ), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframes" + + CATEGORY = "adv-controlnet/keyframes" + + def validate_index(self, index: int, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: + # if part of range, do nothing + if is_range: + return index + # otherwise, validate index + # validate not out of range - only when latent_count is passed in + if latent_count > 0 and index > latent_count-1: + raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.") + # if negative, validate not out of range + if index < 0: + if not allow_negative: + raise IndexError(f"Negative indeces not allowed, but was {index}.") + conv_index = latent_count+index + if conv_index < 0: + raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.") + index = conv_index + return index + + def convert_to_index_int(self, raw_index: str, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int: + try: + return self.validate_index(int(raw_index), latent_count=latent_count, is_range=is_range, allow_negative=allow_negative) + except ValueError as e: + raise ValueError(f"index '{raw_index}' must be an integer.", e) + + def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframe]: + if not latent_indeces: + return set() + all_indeces = [i for i in range(0, latent_count)] + allow_negative = latent_count > 0 + chosen_indeces = set() + # parse string - allow positive ints, negative ints, and ranges separated by ':' + groups = latent_indeces.split(",") + groups = [g.strip() for g in groups] + for g in groups: + # parse strengths - default to 1.0 if no strength given + strength = 1.0 + if '=' in g: + g, strength_str = g.split("=", 1) + g = g.strip() + try: + strength = float(strength_str.strip()) + except ValueError as e: + raise ValueError(f"strength '{strength_str}' must be a float.", e) + if strength < 0: + raise ValueError(f"Strength '{strength}' cannot be negative.") + # parse range of indeces (e.g. 2:16) + if ':' in g: + index_range = g.split(":", 1) + index_range = [r.strip() for r in index_range] + start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative) + end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative) + for i in all_indeces[start_index:end_index]: + chosen_indeces.add(LatentKeyframe(i, strength)) + # parse individual indeces + else: + chosen_indeces.add(LatentKeyframe(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength)) + return chosen_indeces + + def load_keyframes(self, + index_strengths: str, + prev_latent_keyframe: LatentKeyframeGroup=None, + latent_image_opt=None): + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroup() + curr_latent_keyframe = LatentKeyframeGroup() + + latent_count = -1 + if latent_image_opt: + latent_count = latent_image_opt['samples'].size()[0] + latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count) + + for latent_keyframe in latent_keyframes: + curr_latent_keyframe.add(latent_keyframe) + + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + return (curr_latent_keyframe,) + + + + class ControlNetLoaderAdvanced: @classmethod def INPUT_TYPES(s): @@ -288,10 +386,12 @@ class ControlNetApplyAdvanced_AdvControlNet: "negative": ("CONDITIONING", ), "control_net": ("CONTROL_NET", ), "image": ("IMAGE", ), - "mask_opt": ("MASK", ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) + }, + "optional": { + "mask_opt": ("MASK", ), } } @@ -301,10 +401,12 @@ class ControlNetApplyAdvanced_AdvControlNet: CATEGORY = "adv-controlnet/loaders/conditioning" - def apply_controlnet(self, positive, negative, control_net, image, mask_opt, strength, start_percent, end_percent): + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_opt=None): if strength == 0: return (positive, negative) + if mask_opt is not None: + mask_hint = mask_opt.movedim(-1,1) control_hint = image.movedim(-1,1) cnets = {} @@ -319,6 +421,12 @@ class ControlNetApplyAdvanced_AdvControlNet: c_net = cnets[prev_cnet] else: c_net = control_net.copy().set_cond_hint(control_hint, strength, (1.0 - start_percent, 1.0 - end_percent)) + # TODO: finish mask implemention, does nothing right now + if mask_opt is not None: + if isinstance(c_net, ControlNetAdvanced) or isinstance(c_net, T2IAdapterAdvanced): + c_net.set_cond_hint_mask(mask_hint) + else: + logger c_net.set_previous_controlnet(prev_cnet) cnets[prev_cnet] = c_net @@ -330,115 +438,6 @@ class ControlNetApplyAdvanced_AdvControlNet: return (out[0], out[1]) -class ControlNetApplyPartialBatch: # NOT USED: was used for a different test, has useful index parsing code though - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "control_net": ("CONTROL_NET", ), - "image": ("IMAGE", ), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), - "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) - }, - "optional": { - "latent_image": ("LATENT", ), - "latent_indeces": ("STRING", {"default": ""}), - } - } - - RETURN_TYPES = ("CONDITIONING","CONDITIONING") - RETURN_NAMES = ("positive", "negative") - FUNCTION = "apply_controlnet" - - CATEGORY = "adv-controlnet/conditioning" - - def validate_index(self, index: int, latent_count: int, is_range: bool = False) -> int: - # if part of range, do nothing - if is_range: - return index - # otherwise, validate index - # validate not out of range - if index > latent_count-1: - raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.") - # if negative, validate not out of range - if index < 0: - conv_index = latent_count+index - if conv_index < 0: - raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.") - index = conv_index - return index - - def convert_to_index_int(self, raw_index: str, is_range: bool = False) -> int: - try: - return self.validate_index(int(raw_index), is_range=is_range) - except ValueError as e: - raise ValueError(f"index '{raw_index}' must be an integer.", e) - - def convert_to_indeces(self, latent_indeces: str, latent_count: int) -> set[int]: - if not latent_indeces: - return set() - all_indeces = [i for i in range(0, latent_count)] - chosen_indeces = set() - # parse string - allow positive ints, negative ints, and ranges separated by ':' - groups = latent_indeces.split(",") - groups = [g.strip() for g in groups] - for g in groups: - # parse range of indeces (e.g. 2:16) - if ':' in g: - index_range = g.split(":", 1) - index_range = [r.strip() for r in index_range] - start_index = self.convert_to_index_int(index_range[0], is_range=True) - end_index = self.convert_to_index_int(index_range[1], is_range=True) - for i in all_indeces[start_index, end_index]: - chosen_indeces.add(i) - # parse individual indeces - else: - chosen_indeces.add(self.convert_to_index_int(g)) - return chosen_indeces - - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, latent_image=None, latent_indeces: str=None): - if strength == 0: - return (positive, negative) - - latent_count = 1 - if latent_image: - latent_count = latent_image['samples'].size()[0] - indeces_to_apply = self.convert_to_indeces(latent_indeces, latent_count) - - control_hint = image.movedim(-1,1) - cnets = {} - - evaluating_positive = True - out = [] - for conditioning in [positive, negative]: - c = [] - if evaluating_positive and latent_count > 1: - # should copy positive conditioning to match latent_count - if len(conditioning) < latent_count: - pass - for t in conditioning: - d = t[1].copy() - - prev_cnet = d.get('control', None) - if prev_cnet in cnets: - c_net = cnets[prev_cnet] - else: - c_net = control_net.copy().set_cond_hint(control_hint, strength, (1.0 - start_percent, 1.0 - end_percent)) - c_net.set_previous_controlnet(prev_cnet) - cnets[prev_cnet] = c_net - - d['control'] = c_net - d['control_apply_to_uncond'] = False - n = [t[0], d] - c.append(n) - evaluating_positive = False - out.append(c) - return (out[0], out[1]) - - class LoadImagesFromDirectory: @classmethod def INPUT_TYPES(s): @@ -447,7 +446,8 @@ class LoadImagesFromDirectory: "directory": ("STRING", {"default": ""}), }, "optional": { - "image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}) + "image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}), + "start_index": ("INT", {"default": 0, "min": 0, "step": 1}), } } @@ -456,7 +456,7 @@ class LoadImagesFromDirectory: CATEGORY = "adv-controlnet/image" - def load_images(self, directory: str, image_load_cap: int = 0): + def load_images(self, directory: str, image_load_cap: int = 0, start_index: int = 0): if not os.path.isdir(directory): raise FileNotFoundError(f"Directory '{directory} cannot be found.'") dir_files = os.listdir(directory) @@ -465,6 +465,8 @@ class LoadImagesFromDirectory: dir_files = sorted(dir_files) dir_files = [os.path.join(directory, x) for x in dir_files] + # start at start_index + dir_files = dir_files[start_index:] images = [] masks = [] @@ -506,8 +508,7 @@ NODE_CLASS_MAPPINGS = { # Keyframes "TimestepKeyframe": TimestepKeyframeNode, "LatentKeyframe": LatentKeyframeNode, - # Conditioning - # "ControlNetApplyPartialBatch": ControlNetApplyPartialBatch, + "LatentKeyframeGroup": LatentKeyframeGroupNode, # Loaders "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, @@ -525,8 +526,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Keyframes "TimestepKeyframe": "Timestep Keyframe", "LatentKeyframe": "Latent Keyframe", - # Conditioning - # "ControlNetApplyPartialBatch": "Apply ControlNet (Partial Batch)", + "LatentKeyframeGroup": "Latent Keyframe Group", # Loaders "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced)", "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced)",