From a78aeedd241a44ad01d5f7d0b598e0e836ea7023 Mon Sep 17 00:00:00 2001 From: peter942 Date: Thu, 14 Dec 2023 01:43:21 +0100 Subject: [PATCH] Bug fixing --- SteerableMotion.py | 65 +- imports/AdvancedControlNet.py | 751 ----------------- imports/AdvancedControlNet/control.py | 773 ++++++++++++++++++ imports/AdvancedControlNet/control_lllite.py | 1 + .../AdvancedControlNet/deprecated_nodes.py | 103 +++ .../latent_keyframe_nodes.py | 320 ++++++++ imports/AdvancedControlNet/logger.py | 36 + imports/AdvancedControlNet/nodes.py | 194 +++++ imports/AdvancedControlNet/reference_nodes.py | 12 + imports/AdvancedControlNet/weight_nodes.py | 201 +++++ imports/IPAdapterPlus.py | 35 +- 11 files changed, 1675 insertions(+), 816 deletions(-) delete mode 100644 imports/AdvancedControlNet.py create mode 100644 imports/AdvancedControlNet/control.py create mode 100644 imports/AdvancedControlNet/control_lllite.py create mode 100644 imports/AdvancedControlNet/deprecated_nodes.py create mode 100644 imports/AdvancedControlNet/latent_keyframe_nodes.py create mode 100644 imports/AdvancedControlNet/logger.py create mode 100644 imports/AdvancedControlNet/nodes.py create mode 100644 imports/AdvancedControlNet/reference_nodes.py create mode 100644 imports/AdvancedControlNet/weight_nodes.py diff --git a/SteerableMotion.py b/SteerableMotion.py index ada9468..dacaefd 100644 --- a/SteerableMotion.py +++ b/SteerableMotion.py @@ -1,21 +1,22 @@ +# Standard library imports from ast import literal_eval from io import BytesIO + +# Third-party library imports import torch import torchvision.transforms as TT from PIL import Image import matplotlib.pyplot as plt +# Local application/library specific imports import folder_paths - -from .imports.IPAdapterPlus import (IPAdapterApplyImport, prep_image,IPAdapterBatchEmbedsImport, IPAdapterEncoderImport,) -from .imports.AdvancedControlNet import ( +from .imports.IPAdapterPlus import (IPAdapterApplyImport, prep_image, IPAdapterEncoderImport,) +from .imports.AdvancedControlNet.latent_keyframe_nodes import ( calculate_weights, - LatentKeyframeInterpolationNodeImport, - ScaledSoftControlNetWeightsImport, - ControlNetLoaderAdvancedImport, - AdvancedControlNetApplyImport, - TimestepKeyframeNodeImport, + LatentKeyframeInterpolationNodeImport ) +from .imports.AdvancedControlNet.weight_nodes import ScaledSoftUniversalWeightsImport +from .imports.AdvancedControlNet.nodes import ControlNetLoaderAdvancedImport, AdvancedControlNetApplyImport,TimestepKeyframeNodeImport class BatchCreativeInterpolationNode: @classmethod @@ -261,8 +262,10 @@ class BatchCreativeInterpolationNode: cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights = [], [], [], [] last_key_frame_position = (keyframe_positions[-1]) + buffer - batches = [] - current_batch = [] + + embeds = [] + masks = [] + existing_embeds = [] for i, (start, end) in enumerate(influence_ranges): # set basic values @@ -298,13 +301,13 @@ class BatchCreativeInterpolationNode: # Import necessary modules latent_keyframe_interpolation_node = LatentKeyframeInterpolationNodeImport() - scaled_soft_control_net_weights = ScaledSoftControlNetWeightsImport() + scaled_soft_control_net_weights = ScaledSoftUniversalWeightsImport() timestep_keyframe_node = TimestepKeyframeNodeImport() control_net_loader = ControlNetLoaderAdvancedImport() apply_advanced_control_net = AdvancedControlNetApplyImport() ipadapter_application = IPAdapterApplyImport() ipadapter_encoder = IPAdapterEncoderImport() - ipadapter_batcher = IPAdapterBatchEmbedsImport() + # ipadapter_batcher = IPAdapterBatchEmbedsImport() # Load keyframe and append frame numbers and weights weights, frame_numbers, latent_keyframe = latent_keyframe_interpolation_node.load_keyframe( @@ -314,7 +317,7 @@ class BatchCreativeInterpolationNode: # Load weights and keyframe control_net_weights, _ = scaled_soft_control_net_weights.load_weights(soft_scaled_cn_weights_multiplier, False) - timestep_keyframe = timestep_keyframe_node.load_keyframe(start_percent=0.0, control_net_weights=control_net_weights, t2i_adapter_weights=None, latent_keyframe=latent_keyframe, prev_timestep_keyframe=None)[0] + timestep_keyframe = timestep_keyframe_node.load_keyframe(start_percent=0.0, control_net_weights=control_net_weights, latent_keyframe=latent_keyframe, prev_timestep_keyframe=None)[0] # Load and apply control net control_net = control_net_loader.load_controlnet(control_net_name, timestep_keyframe)[0] @@ -332,33 +335,29 @@ class BatchCreativeInterpolationNode: ipadapter_frame_numbers.append(ipa_frame_numbers) ipadapter_weights.append(ipa_weights) - # Create mask batch and apply ipadapter + + mask = create_mask_batch(last_key_frame_position, ipa_weights, frame_numbers) + # add mask to masks list + masks.append(mask) - masks = create_mask_batch(last_key_frame_position, weights, frame_numbers) + embed, = ipadapter_encoder.preprocess(clip_vision, prepped_image, True, 0.0, 1.0) + # add embeds to current batch + embeds.append(embed) - # Apply ipadapter - encoded, = ipadapter_encoder.preprocess(clip_vision, prepped_image, True, ipadapter_noise, 1.0, image_2=None, image_3=None, image_4=None, weight_2=1.0, weight_3=1.0, weight_4=1.0) + model, = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, image=None, weight_type="original", + noise=ipadapter_noise, embeds=embed, attn_mask=mask, start_at=0.0, end_at=1.0, unfold_batch=True) + - current_batch.append(encoded) + # print out the format for the embeds + + # merged_embeds = torch.cat(embeds, dim=1) - # If (i+1) is divisible by 8, start a new batch - if (i + 1) % 8 == 0: - batches.append(current_batch) - current_batch = [] + # stacked_masks = torch.stack(masks) + + # merged_masks = torch.cat(masks, dim=1) - # Add the last batch if it's not empty - if current_batch: - batches.append(current_batch) - # embeds = ipadapter_batcher.batch(self, embed1, embed2) - - for batch in batches: - # Combine all the encoded data in the batch into a single tensor - embeds = torch.cat(batch, dim=1) - # Apply ipadapter - model, = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, image=None, weight_type="original", noise=None, embeds=embeds, attn_mask=masks, start_at=0.0, end_at=1.0, unfold_batch=True) - comparison_diagram, = plot_weight_comparison(cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights, buffer) return comparison_diagram, positive, negative, model diff --git a/imports/AdvancedControlNet.py b/imports/AdvancedControlNet.py deleted file mode 100644 index 8e82867..0000000 --- a/imports/AdvancedControlNet.py +++ /dev/null @@ -1,751 +0,0 @@ -from typing import Union - -from collections.abc import Iterable -import folder_paths -import torch -import numpy as np -from torch import Tensor -from comfy.controlnet import ControlNet, T2IAdapter,broadcast_image_to -import comfy.utils -import comfy.controlnet as comfy_cn - -ControlNetWeightsTypeImport = list[float] -T2IAdapterWeightsTypeImport = list[float] - - -class LatentKeyframeImport: - def __init__(self, batch_index: int, strength: float) -> None: - self.batch_index = batch_index - self.strength = strength - -class LatentKeyframeGroupImport: - def __init__(self) -> None: - self.keyframes: list[LatentKeyframeImport] = [] - - def add(self, keyframe: LatentKeyframeImport) -> None: - added = False - # replace existing keyframe if same batch_index - for i in range(len(self.keyframes)): - if self.keyframes[i].batch_index == keyframe.batch_index: - self.keyframes[i] = keyframe - added = True - break - if not added: - self.keyframes.append(keyframe) - self.keyframes.sort(key=lambda k: k.batch_index) - - def get_index(self, index: int) -> Union[LatentKeyframeImport, None]: - try: - return self.keyframes[index] - except IndexError: - return None - - def __getitem__(self, index) -> LatentKeyframeImport: - return self.keyframes[index] - - def is_empty(self) -> bool: - return len(self.keyframes) == 0 - -class TimestepKeyframeImport: - def __init__(self, - start_percent: float = 0.0, - control_net_weights: ControlNetWeightsTypeImport = None, - t2i_adapter_weights: T2IAdapterWeightsTypeImport = None, - latent_keyframes: LatentKeyframeGroupImport = None, - default_latent_strength: float = 0.0) -> None: - self.start_percent = start_percent - self.control_net_weights = control_net_weights - self.t2i_adapter_weights = t2i_adapter_weights - self.latent_keyframes = latent_keyframes - self.default_latent_strength = default_latent_strength - - - @classmethod - def default(cls) -> 'TimestepKeyframeImport': - return cls(0.0) - -class TimestepKeyframeGroupImport: - def __init__(self) -> None: - self.keyframes: list[TimestepKeyframeImport] = [] - self.keyframes.append(TimestepKeyframeImport.default()) - - def add(self, keyframe: TimestepKeyframeImport) -> None: - added = False - # replace existing keyframe if same start_percent - for i in range(len(self.keyframes)): - if self.keyframes[i].start_percent == keyframe.start_percent: - self.keyframes[i] = keyframe - added = True - break - if not added: - self.keyframes.append(keyframe) - self.keyframes.sort(key=lambda k: k.start_percent) - - def get_index(self, index: int) -> Union[TimestepKeyframeImport, None]: - try: - return self.keyframes[index] - except IndexError: - return None - - def __getitem__(self, index) -> TimestepKeyframeImport: - return self.keyframes[index] - - def is_empty(self) -> bool: - return len(self.keyframes) == 0 - - @classmethod - def default(cls, keyframe: TimestepKeyframeImport) -> 'TimestepKeyframeGroupImport': - group = cls() - group.keyframes[0] = keyframe - return group - - - - -class AdvancedControlNetApplyImport: - @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": { - "mask_optional": ("MASK", ), - } - } - - RETURN_TYPES = ("CONDITIONING","CONDITIONING") - RETURN_NAMES = ("positive", "negative") - FUNCTION = "apply_controlnet" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" - - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=None): - if strength == 0: - return (positive, negative) - - control_hint = image.movedim(-1,1) - cnets = {} - - out = [] - for conditioning in [positive, negative]: - c = [] - - 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, (start_percent, end_percent)) - # set cond hint mask - if mask_optional is not None: - if is_advanced_controlnet(c_net): - # if not in the form of a batch, make it so - if len(mask_optional.shape) < 3: - mask_optional = mask_optional.unsqueeze(0) - c_net.set_cond_hint_mask(mask_optional) - 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) - out.append(c) - return (out[0], out[1]) - -class LatentKeyframeGroupNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "index_strengths": ("STRING", {"multiline": True, "default": ""}), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - "latent_optional": ("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[LatentKeyframeImport]: - 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(LatentKeyframeImport(i, strength)) - # parse individual indeces - else: - chosen_indeces.add(LatentKeyframeImport(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: LatentKeyframeGroupImport=None, - latent_image_opt=None): - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - curr_latent_keyframe = LatentKeyframeGroupImport() - - 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 LatentKeyframeInterpolationNodeImport: - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), - "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), - "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"], ), - "revert_direction_at_midpoint": ("BOOLEAN", {"default": False}), - }, - "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_excl: int, - strength_to: float, - interpolation: str, - revert_direction_at_midpoint: bool=False, - last_key_frame_position: int=0, - i=0, - number_of_items=0, - buffer=0, - prev_latent_keyframe: LatentKeyframeGroupImport=None): - - - - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - - curr_latent_keyframe = LatentKeyframeGroupImport() - - weights, frame_numbers = calculate_weights(batch_index_from, batch_index_to_excl, strength_from, strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position,i,number_of_items, buffer) - - for i, frame_number in enumerate(frame_numbers): - keyframe = LatentKeyframeImport(frame_number, float(weights[i])) - curr_latent_keyframe.add(keyframe) - - for latent_keyframe in prev_latent_keyframe.keyframes: - curr_latent_keyframe.add(latent_keyframe) - - - return (weights, frame_numbers, curr_latent_keyframe,) - -class ControlNetAdvancedImport(ControlNet): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroupImport, global_average_pooling=False, device=None): - super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device) - # initialize timestep_keyframes - self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroupImport() - self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0] - # initialize weights - 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.mask_cond_hint_original = None - self.mask_cond_hint = 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.mask_cond_hint_original = 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 - # perform special version of get_control that supports sliding context and masks - return self.sliding_get_control(x_noisy, t, cond, batched_number) - - def sliding_get_control(self, x_noisy: Tensor, 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 - - # make cond_hint appropriate dimensions - # TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present - if self.sub_idxs is not None or 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 - # if self.cond_hint_original length matches real latent count, need to subdivide it - if self.cond_hint_original.size(0) == self.full_latent_length: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device) - else: - 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 x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) - - # make mask appropriate dimensions, if present - if self.mask_cond_hint_original is not None: - if self.sub_idxs is not None or self.mask_cond_hint is None or x_noisy.shape[2] * 8 != self.mask_cond_hint.shape[1] or x_noisy.shape[3] * 8 != self.mask_cond_hint.shape[2]: - if self.mask_cond_hint is not None: - del self.mask_cond_hint - self.mask_cond_hint = None - # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM - # resize mask and match batch count - self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8) - actual_latent_length = x_noisy.shape[0] // batched_number - self.mask_cond_hint = comfy.utils.repeat_to_batch_size(self.mask_cond_hint, actual_latent_length if self.sub_idxs is None else self.full_latent_length) - if self.sub_idxs is not None: - self.mask_cond_hint = self.mask_cond_hint[self.sub_idxs] - # make cond_hint_mask length match x_noise - if x_noisy.shape[0] != self.mask_cond_hint.shape[0]: - self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number) - self.mask_cond_hint = self.mask_cond_hint.to(self.control_model.dtype).to(self.device) - - context = cond['c_crossattn'] - # uses 'y' in new ComfyUI update - y = cond.get('y', None) - if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI - y = cond.get('c_adm', None) - if y is not None: - y = y.to(self.control_model.dtype) - timestep = self.model_sampling_current.timestep(t) - x_noisy = self.model_sampling_current.calculate_input(t, x_noisy) - - control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=timestep.float(), 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: Tensor, current_timestep_keyframe: TimestepKeyframeImport, batched_number: int): - # apply strengths, and get batch indeces to default out - # AKA latents that should not be influenced by ControlNet - if current_timestep_keyframe.latent_keyframes is not None: - latent_count = x.size(0)//batched_number - indeces_to_default = 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: - 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_default: - indeces_to_default.remove(real_index) - # otherwise, see if batch_index is even included in this set of latents - else: - real_index = mapped_indeces.get(real_index, None) - if real_index is None: - continue - indeces_to_default.remove(real_index) - - # apply strength for each batched cond/uncond - for b in range(batched_number): - x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength - - # default them out by multiplying by default_latent_strength - for batch_index in indeces_to_default: - # apply default for each batched cond/uncond - for b in range(batched_number): - x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * current_timestep_keyframe.default_latent_strength - # apply masks - if self.mask_cond_hint is not None: - # first, resize mask to required dims - masks = prepare_mask_batch(self.mask_cond_hint, x.shape) - x[:] = x[:] * masks - - def copy(self): - c = ControlNetAdvancedImport(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) - self.copy_to(c) - return c - - def cleanup(self): - super().cleanup() - self.sub_idxs = None - self.full_latent_length = 0 - self.context_length = 0 - -class ControlNetLoaderAdvancedImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "control_net_name": (folder_paths.get_filename_list("controlnet"), ), - }, - "optional": { - "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), - } - } - - RETURN_TYPES = ("CONTROL_NET", ) - FUNCTION = "load_controlnet" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" - - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroupImport=None): - controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - controlnet = load_controlnet(controlnet_path, timestep_keyframe) - return (controlnet,) - -class TimestepKeyframeNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), - }, - "optional": { - "control_net_weights": ("CONTROL_NET_WEIGHTS", ), - "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), - "latent_keyframe": ("LATENT_KEYFRAME", ), - "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), - } - } - - RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) - FUNCTION = "load_keyframe" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, - start_percent: float, - control_net_weights: ControlNetWeightsTypeImport=None, - t2i_adapter_weights: T2IAdapterWeightsTypeImport=None, - latent_keyframe: LatentKeyframeGroupImport=None, - prev_timestep_keyframe: TimestepKeyframeGroupImport=None): - if not prev_timestep_keyframe: - prev_timestep_keyframe = TimestepKeyframeGroupImport() - keyframe = TimestepKeyframeImport(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe) - prev_timestep_keyframe.add(keyframe) - return (prev_timestep_keyframe,) - -class ScaledSoftControlNetWeightsImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-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, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_net_weights=weights))) - -class LatentKeyframeBatchedGroupNodeImport: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001}), - }, - "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), - } - } - - RETURN_TYPES = ("LATENT_KEYFRAME", ) - FUNCTION = "load_keyframe" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - - def load_keyframe(self, strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroupImport=None): - if not prev_latent_keyframe: - prev_latent_keyframe = LatentKeyframeGroupImport() - curr_latent_keyframe = LatentKeyframeGroupImport() - - # if received a normal float input, do nothing - if type(strengths) in (float, int): - print("No batched strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.") - # if iterable, attempt to create LatentKeyframes with chosen strengths - elif isinstance(strengths, Iterable): - for idx, strength in enumerate(strengths): - keyframe = LatentKeyframeImport(idx, strength) - curr_latent_keyframe.add(keyframe) - else: - raise ValueError(f"Expected strengths to be an iterable input, but was {type(strengths).__repr__}.") - - # replace values with prev_latent_keyframes - for latent_keyframe in prev_latent_keyframe.keyframes: - curr_latent_keyframe.add(latent_keyframe) - - return (curr_latent_keyframe,) - -class T2IAdapterAdvancedImport(T2IAdapter): - def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroupImport, channels_in, device=None): - super().__init__(t2i_model=t2i_model, channels_in=channels_in, device=device) - self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroupImport() - self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0] - first_weight = self.timestep_keyframes.keyframes[0].t2i_adapter_weights if self.timestep_keyframes.get_index(0) else None - 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)) - - 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 - 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: TimestepKeyframeImport, batched_number: int): - # For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters - # TODO: support masks - return - - def copy(self): - c = T2IAdapterAdvancedImport(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 is_advanced_controlnet(input_object): - return isinstance(input_object, ControlNetAdvancedImport) or isinstance(input_object, T2IAdapterAdvancedImport) - -def control_merge_inject(self, control_input, control_output, control_prev, output_dtype): - out = {'input':[], 'middle':[], 'output': []} - - if control_input is not None: - for i in range(len(control_input)): - key = 'input' - x = control_input[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) - - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - out[key].insert(0, x) - - if control_output is not None: - for i in range(len(control_output)): - if i == (len(control_output) - 1): - key = 'middle' - index = 0 - else: - key = 'output' - index = i - x = control_output[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) - - if self.global_average_pooling: - x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - - out[key].append(x) - if control_prev is not None: - for x in ['input', 'middle', 'output']: - o = out[x] - for i in range(len(control_prev[x])): - prev_val = control_prev[x][i] - if i >= len(o): - o.append(prev_val) - elif prev_val is not None: - if o[i] is None: - o[i] = prev_val - else: - o[i] += prev_val - return out - -def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): - mask = mask.clone() - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") - if match_dim1: - mask = torch.cat([mask] * shape[1], dim=1) - return mask - -def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroupImport=None, model=None): - control = comfy_cn.load_controlnet(ckpt_path, model=model) - # if exactly ControlNet returned, transform it into ControlNetAdvanced - if type(control) == ControlNet: - return ControlNetAdvancedImport(control.control_model, timestep_keyframe, global_average_pooling=control.global_average_pooling) - # if T2IAdapter returned, transform it into T2IAdapterAdvanced - elif isinstance(control, T2IAdapter): - return T2IAdapterAdvancedImport(control.t2i_model, timestep_keyframe, control.channels_in) - # otherwise, leave it be - probably a ControlLora for SDXL (no support for advanced stuff yet from here) - # TODO add ControlLoraAdvanced - return control - -def calculate_weights(batch_index_from, batch_index_to, strength_from, strength_to, interpolation,revert_direction_at_midpoint, last_key_frame_position,i, number_of_items,buffer): - - # Initialize variables based on the position of the keyframe - range_start = batch_index_from - range_end = batch_index_to - # if it's the first value, set influence range from 1.0 to 0.0 - if buffer > 0: - if i == 0: - range_start = 0 - elif i == 1: - range_start = buffer - else: - if i == 1: - range_start = 0 - - if i == number_of_items - 1: - range_end = last_key_frame_position - - steps = range_end - range_start - diff = strength_to - strength_from - - # Calculate index for interpolation - index = np.linspace(0, 1, steps // 2 + 1) if revert_direction_at_midpoint else np.linspace(0, 1, steps) - - # Calculate weights based on interpolation type - if interpolation == "linear": - weights = np.linspace(strength_from, strength_to, len(index)) - elif interpolation == "ease-in": - weights = diff * np.power(index, 2) + strength_from - elif interpolation == "ease-out": - weights = diff * (1 - np.power(1 - index, 2)) + strength_from - elif interpolation == "ease-in-out": - weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from - - # If it's a middle keyframe, mirror the weights - if revert_direction_at_midpoint: - weights = np.concatenate([weights, weights[::-1]]) - - # Generate frame numbers - frame_numbers = np.arange(range_start, range_start + len(weights)) - - # "Dropper" component: For keyframes with negative start, drop the weights - if range_start < 0 and i > 0: - drop_count = abs(range_start) - weights = weights[drop_count:] - frame_numbers = frame_numbers[drop_count:] - - # Dropper component: for keyframes a range_End is greater than last_key_frame_position, drop the weights - if range_end > last_key_frame_position and i < number_of_items - 1: - drop_count = range_end - last_key_frame_position - weights = weights[:-drop_count] - frame_numbers = frame_numbers[:-drop_count] - - return weights, frame_numbers \ No newline at end of file diff --git a/imports/AdvancedControlNet/control.py b/imports/AdvancedControlNet/control.py new file mode 100644 index 0000000..73bc1d6 --- /dev/null +++ b/imports/AdvancedControlNet/control.py @@ -0,0 +1,773 @@ +from typing import Union +from torch import Tensor +import torch + +import comfy.utils +import comfy.controlnet as comfy_cn +from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to + + +def get_properly_arranged_t2i_weights(initial_weights: list[float]): + new_weights = [] + new_weights.extend([initial_weights[0]]*3) + new_weights.extend([initial_weights[1]]*3) + new_weights.extend([initial_weights[2]]*3) + new_weights.extend([initial_weights[3]]*3) + return new_weights + + +class ControlWeightTypeImport: + DEFAULT = "default" + UNIVERSAL = "universal" + T2IADAPTER = "t2iadapter" + CONTROLNET = "controlnet" + CONTROLLORA = "controllora" + CONTROLLLLITE = "controllllite" + + +class ControlWeightsImport: + def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None): + self.weight_type = weight_type + self.base_multiplier = base_multiplier + self.flip_weights = flip_weights + self.weights = weights + if self.weights is not None and self.flip_weights: + self.weights.reverse() + self.weight_mask = weight_mask + + def get(self, idx: int) -> Union[float, Tensor]: + # if weights is not none, return index + if self.weights is not None: + return self.weights[idx] + return 1.0 + + @classmethod + def default(cls): + return cls(ControlWeightTypeImport.DEFAULT) + + @classmethod + def universal(cls, base_multiplier: float, flip_weights: bool=False): + return cls(ControlWeightTypeImport.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights) + + @classmethod + def universal_mask(cls, weight_mask: Tensor): + return cls(ControlWeightTypeImport.UNIVERSAL, weight_mask=weight_mask) + + @classmethod + def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*12 + return cls(ControlWeightTypeImport.T2IADAPTER, weights=weights,flip_weights=flip_weights) + + @classmethod + def controlnet(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*13 + return cls(ControlWeightTypeImport.CONTROLNET, weights=weights, flip_weights=flip_weights) + + @classmethod + def controllora(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*10 + return cls(ControlWeightTypeImport.CONTROLLORA, weights=weights, flip_weights=flip_weights) + + @classmethod + def controllllite(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + # TODO: make this have a real value + weights = [1.0]*200 + return cls(ControlWeightTypeImport.CONTROLLLLITE, weights=weights, flip_weights=flip_weights) + + +class StrengthInterpolationImport: + LINEAR = "linear" + EASE_IN = "ease-in" + EASE_OUT = "ease-out" + EASE_IN_OUT = "ease-in-out" + NONE = "none" + + +class LatentKeyframeImport: + def __init__(self, batch_index: int, strength: float) -> None: + self.batch_index = batch_index + self.strength = strength + + +# always maintain sorted state (by batch_index of LatentKeyframe) +class LatentKeyframeGroupImport: + def __init__(self) -> None: + self.keyframes: list[LatentKeyframeImport] = [] + + def add(self, keyframe: LatentKeyframeImport) -> None: + added = False + # replace existing keyframe if same batch_index + for i in range(len(self.keyframes)): + if self.keyframes[i].batch_index == keyframe.batch_index: + self.keyframes[i] = keyframe + added = True + break + if not added: + self.keyframes.append(keyframe) + self.keyframes.sort(key=lambda k: k.batch_index) + + def get_index(self, index: int) -> Union[LatentKeyframeImport, None]: + try: + return self.keyframes[index] + except IndexError: + return None + + def __getitem__(self, index) -> LatentKeyframeImport: + return self.keyframes[index] + + def is_empty(self) -> bool: + return len(self.keyframes) == 0 + + def clone(self) -> 'LatentKeyframeGroupImport': + cloned = LatentKeyframeGroupImport() + for tk in self.keyframes: + cloned.add(tk) + return cloned + + +class TimestepKeyframeImport: + def __init__(self, + start_percent: float = 0.0, + strength: float = 1.0, + interpolation: str = StrengthInterpolationImport.NONE, + control_weights: ControlWeightsImport = None, + latent_keyframes: LatentKeyframeGroupImport = None, + null_latent_kf_strength: float = 0.0, + inherit_missing: bool = True, + guarantee_usage: bool = True, + mask_hint_orig: Tensor = None) -> None: + self.start_percent = start_percent + self.start_t = 999999999.9 + self.strength = strength + self.interpolation = interpolation + self.control_weights = control_weights + self.latent_keyframes = latent_keyframes + self.null_latent_kf_strength = null_latent_kf_strength + self.inherit_missing = inherit_missing + self.guarantee_usage = guarantee_usage + self.mask_hint_orig = mask_hint_orig + + def has_control_weights(self): + return self.control_weights is not None + + def has_latent_keyframes(self): + return self.latent_keyframes is not None + + def has_mask_hint(self): + return self.mask_hint_orig is not None + + + @classmethod + def default(cls) -> 'TimestepKeyframeImport': + return cls(0.0) + + +# always maintain sorted state (by start_percent of TimestepKeyFrame) +class TimestepKeyframeGroupImport: + def __init__(self) -> None: + self.keyframes: list[TimestepKeyframeImport] = [] + self.keyframes.append(TimestepKeyframeImport.default()) + + def add(self, keyframe: TimestepKeyframeImport) -> None: + added = False + # replace existing keyframe if same start_percent + for i in range(len(self.keyframes)): + if self.keyframes[i].start_percent == keyframe.start_percent: + self.keyframes[i] = keyframe + added = True + break + if not added: + self.keyframes.append(keyframe) + self.keyframes.sort(key=lambda k: k.start_percent) + + def get_index(self, index: int) -> Union[TimestepKeyframeImport, None]: + try: + return self.keyframes[index] + except IndexError: + return None + + def has_index(self, index: int) -> int: + return index >=0 and index < len(self.keyframes) + + def __getitem__(self, index) -> TimestepKeyframeImport: + return self.keyframes[index] + + def __len__(self) -> int: + return len(self.keyframes) + + def is_empty(self) -> bool: + return len(self.keyframes) == 0 + + def clone(self) -> 'TimestepKeyframeGroupImport': + cloned = TimestepKeyframeGroupImport() + for tk in self.keyframes: + cloned.add(tk) + return cloned + + @classmethod + def default(cls, keyframe: TimestepKeyframeImport) -> 'TimestepKeyframeGroupImport': + group = cls() + group.keyframes[0] = keyframe + return group + + +# used to inject ControlNetAdvancedImport and T2IAdapterAdvancedImport control_merge function + + +class AdvancedControlBaseImport: + def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroupImport, weights_default: ControlWeightsImport): + self.base = base + self.compatible_weights = [ControlWeightTypeImport.UNIVERSAL] + self.add_compatible_weight(weights_default.weight_type) + # mask for which parts of controlnet output to keep + self.mask_cond_hint_original = None + self.mask_cond_hint = None + self.tk_mask_cond_hint_original = None + self.tk_mask_cond_hint = None + self.weight_mask_cond_hint = None + # actual index values + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 + # timesteps + self.t: Tensor = None + self.batched_number: int = None + # weights + override + self.weights: ControlWeightsImport = None + self.weights_default: ControlWeightsImport = weights_default + self.weights_override: ControlWeightsImport = None + # latent keyframe + override + self.latent_keyframes: LatentKeyframeGroupImport = None + self.latent_keyframe_override: LatentKeyframeGroupImport = None + # initialize timestep_keyframes + self.set_timestep_keyframes(timestep_keyframes) + # override some functions + self.get_control = self.get_control_inject + self.control_merge = self.control_merge_inject#.__get__(self, type(self)) + self.pre_run = self.pre_run_inject + self.cleanup = self.cleanup_inject + + def add_compatible_weight(self, control_weight_type: str): + self.compatible_weights.append(control_weight_type) + + def verify_all_weights(self, throw_error=True): + # first, check if override exists - if so, only need to check the override + if self.weights_override is not None: + if self.weights_override.weight_type not in self.compatible_weights: + msg = f"Weight override is type {self.weights_override.weight_type}, but loaded {type(self).__name__}" + \ + f"only supports {self.compatible_weights} weights." + raise WeightTypeExceptionImport(msg) + # otherwise, check all timestep keyframe weights + else: + for tk in self.timestep_keyframes.keyframes: + if tk.has_control_weights() and tk.control_weights.weight_type not in self.compatible_weights: + msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type" + \ + f"{tk.control_weights.weight_type}, but loaded {type(self).__name__} only supports {self.compatible_weights} weights." + raise WeightTypeExceptionImport(msg) + + def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroupImport): + self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroupImport() + # prepare first timestep_keyframe related stuff + self.current_timestep_keyframe = None + self.current_timestep_index = -1 + self.next_timestep_keyframe = None + self.weights = None + self.latent_keyframes = None + + def prepare_current_timestep(self, t: Tensor, batched_number: int): + self.t = t + self.batched_number = batched_number + # get current step percent + curr_t: float = t[0] + prev_index = self.current_timestep_index + # if has next index, loop through and see if need to switch + if self.timestep_keyframes.has_index(self.current_timestep_index+1): + for i in range(self.current_timestep_index+1, len(self.timestep_keyframes)): + eval_tk = self.timestep_keyframes[i] + # check if start percent is less or equal to curr_t + if eval_tk.start_t >= curr_t: + self.current_timestep_index = i + self.current_timestep_keyframe = eval_tk + # keep track of control weights, latent keyframes, and masks, + # accounting for inherit_missing + if self.current_timestep_keyframe.has_control_weights(): + self.weights = self.current_timestep_keyframe.control_weights + elif not self.current_timestep_keyframe.inherit_missing: + self.weights = self.weights_default + if self.current_timestep_keyframe.has_latent_keyframes(): + self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes + elif not self.current_timestep_keyframe.inherit_missing: + self.latent_keyframes = None + if self.current_timestep_keyframe.has_mask_hint(): + self.tk_mask_cond_hint_original = self.current_timestep_keyframe.mask_hint_orig + elif not self.current_timestep_keyframe.inherit_missing: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None + # if guarantee_usage, stop searching for other TKs + if self.current_timestep_keyframe.guarantee_usage: + break + # if eval_tk is outside of percent range, stop looking further + else: + break + + # if index changed, apply overrides + if prev_index != self.current_timestep_index: + if self.weights_override is not None: + self.weights = self.weights_override + if self.latent_keyframe_override is not None: + self.latent_keyframes = self.latent_keyframe_override + + # make sure weights and latent_keyframes are in a workable state + # Note: each AdvancedControlBaseImport should create their own get_universal_weights class + self.prepare_weights() + + def prepare_weights(self): + if self.weights is None or self.weights.weight_type == ControlWeightTypeImport.DEFAULT: + self.weights = self.weights_default + elif self.weights.weight_type == ControlWeightTypeImport.UNIVERSAL: + # if universal and weight_mask present, no need to convert + if self.weights.weight_mask is not None: + return + self.weights = self.get_universal_weights() + + def get_universal_weights(self) -> ControlWeightsImport: + return self.weights + + def set_cond_hint_mask(self, mask_hint): + self.mask_cond_hint_original = mask_hint + return self + + def pre_run_inject(self, model, percent_to_timestep_function): + self.base.pre_run(model, percent_to_timestep_function) + self.pre_run_advanced(model, percent_to_timestep_function) + + def pre_run_advanced(self, model, percent_to_timestep_function): + # for each timestep keyframe, calculate the start_t + for tk in self.timestep_keyframes.keyframes: + tk.start_t = percent_to_timestep_function(tk.start_percent) + # clear variables + self.cleanup_advanced() + + def get_control_inject(self, x_noisy, t, cond, batched_number): + # prepare timestep and everything related + self.prepare_current_timestep(t=t, batched_number=batched_number) + # if should not perform any actions for the controlnet, exit without doing any work + if self.strength == 0.0 or self.current_timestep_keyframe.strength == 0.0: + control_prev = None + if self.previous_controlnet is not None: + control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) + if control_prev is not None: + return control_prev + else: + return None + # otherwise, perform normal function + return self.get_control_advanced(x_noisy, t, cond, batched_number) + + def get_control_advanced(self, x_noisy, t, cond, batched_number): + pass + + def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]: + if self.weights.weight_mask is not None: + # prepare weight mask + self.prepare_weight_mask_cond_hint(x, self.batched_number) + # adjust mask for current layer and return + return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, layers=layers)) + return self.weights.get(idx=idx) + + def get_calc_pow(self, idx: int, layers: int) -> int: + return (layers-1)-idx + + def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): + # apply strengths, and get batch indeces to null out + # AKA latents that should not be influenced by ControlNet + if self.latent_keyframes is not None: + latent_count = x.size(0)//batched_number + indeces_to_null = 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 self.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_null: + indeces_to_null.remove(real_index) + # otherwise, see if batch_index is even included in this set of latents + else: + real_index = mapped_indeces.get(real_index, None) + if real_index is None: + continue + indeces_to_null.remove(real_index) + + # if real_index is outside the bounds of latents, don't apply + if real_index >= latent_count or real_index < 0: + continue + + # apply strength for each batched cond/uncond + for b in range(batched_number): + x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength + + # null them out by multiplying by null_latent_kf_strength + for batch_index in indeces_to_null: + # apply null for each batched cond/uncond + for b in range(batched_number): + x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength + # apply masks, resizing mask to required dims + if self.mask_cond_hint is not None: + masks = prepare_mask_batch(self.mask_cond_hint, x.shape) + x[:] = x[:] * masks + if self.tk_mask_cond_hint is not None: + masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape) + x[:] = x[:] * masks + # apply timestep keyframe strengths + if self.current_timestep_keyframe.strength != 1.0: + x[:] *= self.current_timestep_keyframe.strength + + def control_merge_inject(self: 'AdvancedControlBaseImport', control_input, control_output, control_prev, output_dtype): + out = {'input':[], 'middle':[], 'output': []} + + if control_input is not None: + for i in range(len(control_input)): + key = 'input' + x = control_input[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.batched_number) + + x *= self.strength * self.calc_weight(i, x, len(control_input)) + if x.dtype != output_dtype: + x = x.to(output_dtype) + out[key].insert(0, x) + + if control_output is not None: + for i in range(len(control_output)): + if i == (len(control_output) - 1): + key = 'middle' + index = 0 + else: + key = 'output' + index = i + x = control_output[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.batched_number) + + if self.global_average_pooling: + x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) + + x *= self.strength * self.calc_weight(i, x, len(control_output)) + if x.dtype != output_dtype: + x = x.to(output_dtype) + + out[key].append(x) + if control_prev is not None: + for x in ['input', 'middle', 'output']: + o = out[x] + for i in range(len(control_prev[x])): + prev_val = control_prev[x][i] + if i >= len(o): + o.append(prev_val) + elif prev_val is not None: + if o[i] is None: + o[i] = prev_val + else: + o[i] += prev_val + return out + + def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype) + self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) + + def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + + def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None): + return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True) + + def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False): + # make mask appropriate dimensions, if present + if orig_mask is not None: + out_mask = getattr(self, attr_name) + if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: + self._reset_attr(attr_name) + del out_mask + # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM + # resize mask and match batch count + multiplier = 1 if direct_attn else 8 + out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier) + actual_latent_length = x_noisy.shape[0] // batched_number + out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length) + if self.sub_idxs is not None: + out_mask = out_mask[self.sub_idxs] + # make cond_hint_mask length match x_noise + if x_noisy.shape[0] != out_mask.shape[0]: + out_mask = broadcast_image_to(out_mask, x_noisy.shape[0], batched_number) + # default dtype to be same as x_noisy + if dtype is None: + dtype = x_noisy.dtype + setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device)) + del out_mask + + def _reset_attr(self, attr_name, new_value=None): + if hasattr(self, attr_name): + delattr(self, attr_name) + setattr(self, attr_name, new_value) + + def cleanup_inject(self): + self.base.cleanup() + self.cleanup_advanced() + + def cleanup_advanced(self): + self.sub_idxs = None + self.full_latent_length = 0 + self.context_length = 0 + self.t = None + self.batched_number = None + self.weights = None + self.latent_keyframes = None + # timestep stuff + self.current_timestep_keyframe = None + self.next_timestep_keyframe = None + self.current_timestep_index = -1 + # clear mask hints + if self.mask_cond_hint is not None: + del self.mask_cond_hint + self.mask_cond_hint = None + if self.tk_mask_cond_hint_original is not None: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None + if self.tk_mask_cond_hint is not None: + del self.tk_mask_cond_hint + self.tk_mask_cond_hint = None + if self.weight_mask_cond_hint is not None: + del self.weight_mask_cond_hint + self.weight_mask_cond_hint = None + + def copy_to_advanced(self, copied: 'AdvancedControlBaseImport'): + copied.mask_cond_hint_original = self.mask_cond_hint_original + copied.weights_override = self.weights_override + copied.latent_keyframe_override = self.latent_keyframe_override + + +class ControlNetAdvancedImport(ControlNet, AdvancedControlBaseImport): + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroupImport, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): + super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + AdvancedControlBaseImport.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeightsImport.controlnet()) + + def get_universal_weights(self) -> ControlWeightsImport: + raw_weights = [(self.weights.base_multiplier ** float(12 - i)) for i in range(13)] + return ControlWeightsImport.controlnet(raw_weights, self.weights.flip_weights) + + def get_control_advanced(self, x_noisy, t, cond, batched_number): + # perform special version of get_control that supports sliding context and masks + return self.sliding_get_control(x_noisy, t, cond, batched_number) + + def sliding_get_control(self, x_noisy: Tensor, 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 + + dtype = self.control_model.dtype + if self.manual_cast_dtype is not None: + dtype = self.manual_cast_dtype + + output_dtype = x_noisy.dtype + # make cond_hint appropriate dimensions + # TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present + if self.sub_idxs is not None or 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 + # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling + if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + else: + 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(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) + + # prepare mask_cond_hint + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) + + context = cond['c_crossattn'] + # uses 'y' in new ComfyUI update + y = cond.get('y', None) + if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI + y = cond.get('c_adm', None) + if y is not None: + y = y.to(dtype) + timestep = self.model_sampling_current.timestep(t) + x_noisy = self.model_sampling_current.calculate_input(t, x_noisy) + + control = self.control_model(x=x_noisy.to(dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(dtype), y=y) + return self.control_merge(None, control, control_prev, output_dtype) + + def copy(self): + c = ControlNetAdvancedImport(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling, load_device=self.load_device, manual_cast_dtype=self.manual_cast_dtype) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + @staticmethod + def from_vanilla(v: ControlNet, timestep_keyframe: TimestepKeyframeGroupImport=None) -> 'ControlNetAdvancedImport': + return ControlNetAdvancedImport(control_model=v.control_model, timestep_keyframes=timestep_keyframe, + global_average_pooling=v.global_average_pooling, device=v.device, load_device=v.load_device, manual_cast_dtype=v.manual_cast_dtype) + + +class T2IAdapterAdvancedImport(T2IAdapter, AdvancedControlBaseImport): + def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroupImport, channels_in, device=None): + super().__init__(t2i_model=t2i_model, channels_in=channels_in, device=device) + AdvancedControlBaseImport.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeightsImport.t2iadapter()) + + def get_universal_weights(self) -> ControlWeightsImport: + raw_weights = [(self.weights.base_multiplier ** float(7 - i)) for i in range(8)] + raw_weights = [raw_weights[-8], raw_weights[-3], raw_weights[-2], raw_weights[-1]] + raw_weights = get_properly_arranged_t2i_weights(raw_weights) + return ControlWeightsImport.t2iadapter(raw_weights, self.weights.flip_weights) + + def get_calc_pow(self, idx: int, layers: int) -> int: + # match how T2IAdapterAdvancedImport deals with universal weights + indeces = [7 - i for i in range(8)] + indeces = [indeces[-8], indeces[-3], indeces[-2], indeces[-1]] + indeces = get_properly_arranged_t2i_weights(indeces) + return indeces[idx] + + def get_control_advanced(self, x_noisy, t, cond, batched_number): + # prepare timestep and everything related + self.prepare_current_timestep(t=t, batched_number=batched_number) + try: + # if sub indexes present, replace original hint with subsection + if self.sub_idxs is not None: + # cond hints + 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] + # mask hints + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + 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 copy(self): + c = T2IAdapterAdvancedImport(self.t2i_model, self.timestep_keyframes, self.channels_in) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + def cleanup(self): + super().cleanup() + self.cleanup_advanced() + + @staticmethod + def from_vanilla(v: T2IAdapter, timestep_keyframe: TimestepKeyframeGroupImport=None) -> 'T2IAdapterAdvancedImport': + return T2IAdapterAdvancedImport(t2i_model=v.t2i_model, timestep_keyframes=timestep_keyframe, channels_in=v.channels_in, device=v.device) + + +class ControlLoraAdvancedImport(ControlLora, AdvancedControlBaseImport): + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroupImport, global_average_pooling=False, device=None): + super().__init__(control_weights=control_weights, global_average_pooling=global_average_pooling, device=device) + AdvancedControlBaseImport.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeightsImport.controllora()) + # use some functions from ControlNetAdvancedImport + self.get_control_advanced = ControlNetAdvancedImport.get_control_advanced.__get__(self, type(self)) + self.sliding_get_control = ControlNetAdvancedImport.sliding_get_control.__get__(self, type(self)) + + def get_universal_weights(self) -> ControlWeightsImport: + raw_weights = [(self.weights.base_multiplier ** float(9 - i)) for i in range(10)] + return ControlWeightsImport.controllora(raw_weights, self.weights.flip_weights) + + def copy(self): + c = ControlLoraAdvancedImport(self.control_weights, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + def cleanup(self): + super().cleanup() + self.cleanup_advanced() + + @staticmethod + def from_vanilla(v: ControlLora, timestep_keyframe: TimestepKeyframeGroupImport=None) -> 'ControlLoraAdvancedImport': + return ControlLoraAdvancedImport(control_weights=v.control_weights, timestep_keyframes=timestep_keyframe, + global_average_pooling=v.global_average_pooling, device=v.device) + + +class ControlLLLiteAdvancedImport(ControlNet, AdvancedControlBaseImport): + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroupImport, device=None): + AdvancedControlBaseImport.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeightsImport.controllllite()) + + +def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroupImport=None, model=None): + control = comfy_cn.load_controlnet(ckpt_path, model=model) + # TODO: support controlnet-lllite + # if is None, see if is a non-vanilla ControlNet + # if control is None: + # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) + # # check if lllite + # if "lllite_unet" in controlnet_data: + # pass + return convert_to_advanced(control, timestep_keyframe=timestep_keyframe) + + +def convert_to_advanced(control, timestep_keyframe: TimestepKeyframeGroupImport=None): + # if already advanced, leave it be + if is_advanced_controlnet(control): + return control + # if exactly ControlNet returned, transform it into ControlNetAdvancedImport + if type(control) == ControlNet: + return ControlNetAdvancedImport.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) + # if exactly ControlLora returned, transform it into ControlLoraAdvancedImport + elif type(control) == ControlLora: + return ControlLoraAdvancedImport.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) + # if T2IAdapter returned, transform it into T2IAdapterAdvancedImport + elif isinstance(control, T2IAdapter): + return T2IAdapterAdvancedImport.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) + # otherwise, leave it be - might be something I am not supporting yet + return control + + +def is_advanced_controlnet(input_object): + return hasattr(input_object, "sub_idxs") + + +# adapted from comfy/sample.py +def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): + mask = mask.clone() + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") + if match_dim1: + mask = torch.cat([mask] * shape[1], dim=1) + return mask + + +# applies min-max normalization, from: +# https://stackoverflow.com/questions/68791508/min-max-normalization-of-a-tensor-in-pytorch +def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): + x_min, x_max = x.min(), x.max() + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + +def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + + +class WeightTypeExceptionImport(TypeError): + "Raised when weight not compatible with AdvancedControlBaseImport object" + pass diff --git a/imports/AdvancedControlNet/control_lllite.py b/imports/AdvancedControlNet/control_lllite.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/imports/AdvancedControlNet/control_lllite.py @@ -0,0 +1 @@ + diff --git a/imports/AdvancedControlNet/deprecated_nodes.py b/imports/AdvancedControlNet/deprecated_nodes.py new file mode 100644 index 0000000..a64ac9b --- /dev/null +++ b/imports/AdvancedControlNet/deprecated_nodes.py @@ -0,0 +1,103 @@ +import os + +import torch + +import numpy as np +from PIL import Image, ImageOps +from .control import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe +from .logger import logger + + +class LoadImagesFromDirectory: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "directory": ("STRING", {"default": ""}), + }, + "optional": { + "image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}), + "start_index": ("INT", {"default": 0, "min": 0, "step": 1}), + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "INT") + FUNCTION = "load_images" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/deprecated" + + 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) + if len(dir_files) == 0: + raise FileNotFoundError(f"No files in directory '{directory}'.") + + 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 = [] + + limit_images = False + if image_load_cap > 0: + limit_images = True + image_count = 0 + + for image_path in dir_files: + if os.path.isdir(image_path): + continue + if limit_images and image_count >= image_load_cap: + break + i = Image.open(image_path) + i = ImageOps.exif_transpose(i) + image = i.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + else: + mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") + images.append(image) + masks.append(mask) + image_count += 1 + + if len(images) == 0: + raise FileNotFoundError(f"No images could be loaded from directory '{directory}'.") + + return (torch.cat(images, dim=0), torch.stack(masks, dim=0), image_count) + + +class TimestepKeyframeNodeDeprecated: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + }, + "optional": { + "control_net_weights": ("CONTROL_NET_WEIGHTS", ), + "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), + "latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) + FUNCTION = "load_keyframe" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + start_percent: float, + control_net_weights: ControlWeights=None, + latent_keyframe: LatentKeyframeGroup=None, + prev_timestep_keyframe: TimestepKeyframeGroup=None): + if not prev_timestep_keyframe: + prev_timestep_keyframe = TimestepKeyframeGroup() + keyframe = TimestepKeyframe(start_percent, control_net_weights, latent_keyframe) + prev_timestep_keyframe.add(keyframe) + return (prev_timestep_keyframe,) diff --git a/imports/AdvancedControlNet/latent_keyframe_nodes.py b/imports/AdvancedControlNet/latent_keyframe_nodes.py new file mode 100644 index 0000000..0d6ee5e --- /dev/null +++ b/imports/AdvancedControlNet/latent_keyframe_nodes.py @@ -0,0 +1,320 @@ +from typing import Union +import numpy as np +from collections.abc import Iterable + +from .control import LatentKeyframeImport, LatentKeyframeGroupImport +from .control import StrengthInterpolationImport as SI +from .logger import logger + + +class LatentKeyframeNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + }, + "optional": { + "prev_latent_kf": ("LATENT_KEYFRAME", ), + } + } + + RETURN_NAMES = ("LATENT_KF", ) + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframe" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + batch_index: int, + strength: float, + prev_latent_kf: LatentKeyframeGroupImport=None, + prev_latent_keyframe: LatentKeyframeGroupImport=None, # old name + ): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() + keyframe = LatentKeyframeImport(batch_index, strength) + prev_latent_keyframe.add(keyframe) + return (prev_latent_keyframe,) + + +class LatentKeyframeGroupNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "index_strengths": ("STRING", {"multiline": True, "default": ""}), + }, + "optional": { + "prev_latent_kf": ("LATENT_KEYFRAME", ), + "latent_optional": ("LATENT", ), + "print_keyframes": ("BOOLEAN", {"default": False}) + } + } + + RETURN_NAMES = ("LATENT_KF", ) + 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[LatentKeyframeImport]: + if not latent_indeces: + return set() + int_latent_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) + # if latents were passed in, base indeces on known latent count + if len(int_latent_indeces) > 0: + for i in int_latent_indeces[start_index:end_index]: + chosen_indeces.add(LatentKeyframeImport(i, strength)) + # otherwise, assume indeces are valid + else: + for i in range(start_index, end_index): + chosen_indeces.add(LatentKeyframeImport(i, strength)) + # parse individual indeces + else: + chosen_indeces.add(LatentKeyframeImport(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_kf: LatentKeyframeGroupImport=None, + prev_latent_keyframe: LatentKeyframeGroupImport=None, # old name + latent_image_opt=None, + print_keyframes=False): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() + curr_latent_keyframe = LatentKeyframeGroupImport() + + 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) + + if print_keyframes: + for keyframe in curr_latent_keyframe.keyframes: + logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}") + + # replace values with prev_latent_keyframes + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + return (curr_latent_keyframe,) + + +class LatentKeyframeInterpolationNodeImport: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), + "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), + "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), + "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), + "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"], ), + "revert_direction_at_midpoint": ("BOOLEAN", {"default": False}), + }, + "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_excl: int, + strength_to: float, + interpolation: str, + revert_direction_at_midpoint: bool=False, + last_key_frame_position: int=0, + i=0, + number_of_items=0, + buffer=0, + prev_latent_keyframe: LatentKeyframeGroupImport=None): + + + + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() + + curr_latent_keyframe = LatentKeyframeGroupImport() + + weights, frame_numbers = calculate_weights(batch_index_from, batch_index_to_excl, strength_from, strength_to, interpolation, revert_direction_at_midpoint, last_key_frame_position,i,number_of_items, buffer) + + for i, frame_number in enumerate(frame_numbers): + keyframe = LatentKeyframeImport(frame_number, float(weights[i])) + curr_latent_keyframe.add(keyframe) + + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + + return (weights, frame_numbers, curr_latent_keyframe,) + +def calculate_weights(batch_index_from, batch_index_to, strength_from, strength_to, interpolation,revert_direction_at_midpoint, last_key_frame_position,i, number_of_items,buffer): + + # Initialize variables based on the position of the keyframe + range_start = batch_index_from + range_end = batch_index_to + # if it's the first value, set influence range from 1.0 to 0.0 + if buffer > 0: + if i == 0: + range_start = 0 + elif i == 1: + range_start = buffer + else: + if i == 1: + range_start = 0 + + if i == number_of_items - 1: + range_end = last_key_frame_position + + steps = range_end - range_start + diff = strength_to - strength_from + + # Calculate index for interpolation + index = np.linspace(0, 1, steps // 2 + 1) if revert_direction_at_midpoint else np.linspace(0, 1, steps) + + # Calculate weights based on interpolation type + if interpolation == "linear": + weights = np.linspace(strength_from, strength_to, len(index)) + elif interpolation == "ease-in": + weights = diff * np.power(index, 2) + strength_from + elif interpolation == "ease-out": + weights = diff * (1 - np.power(1 - index, 2)) + strength_from + elif interpolation == "ease-in-out": + weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from + + # If it's a middle keyframe, mirror the weights + if revert_direction_at_midpoint: + weights = np.concatenate([weights, weights[::-1]]) + + # Generate frame numbers + frame_numbers = np.arange(range_start, range_start + len(weights)) + + # "Dropper" component: For keyframes with negative start, drop the weights + if range_start < 0 and i > 0: + drop_count = abs(range_start) + weights = weights[drop_count:] + frame_numbers = frame_numbers[drop_count:] + + # Dropper component: for keyframes a range_End is greater than last_key_frame_position, drop the weights + if range_end > last_key_frame_position and i < number_of_items - 1: + drop_count = range_end - last_key_frame_position + weights = weights[:-drop_count] + frame_numbers = frame_numbers[:-drop_count] + + return weights, frame_numbers + +class LatentKeyframeBatchedGroupNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}), + }, + "optional": { + "prev_latent_kf": ("LATENT_KEYFRAME", ), + "print_keyframes": ("BOOLEAN", {"default": False}) + } + } + + RETURN_NAMES = ("LATENT_KF", ) + RETURN_TYPES = ("LATENT_KEYFRAME", ) + FUNCTION = "load_keyframe" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, float_strengths: Union[float, list[float]], + prev_latent_kf: LatentKeyframeGroupImport=None, + prev_latent_keyframe: LatentKeyframeGroupImport=None, # old name + print_keyframes=False): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf + if not prev_latent_keyframe: + prev_latent_keyframe = LatentKeyframeGroupImport() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() + curr_latent_keyframe = LatentKeyframeGroupImport() + + # if received a normal float input, do nothing + if type(float_strengths) in (float, int): + logger.info("No batched float_strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.") + # if iterable, attempt to create LatentKeyframes with chosen strengths + elif isinstance(float_strengths, Iterable): + for idx, strength in enumerate(float_strengths): + keyframe = LatentKeyframeImport(idx, strength) + curr_latent_keyframe.add(keyframe) + else: + raise ValueError(f"Expected strengths to be an iterable input, but was {type(float_strengths).__repr__}.") + + if print_keyframes: + for keyframe in curr_latent_keyframe.keyframes: + logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}") + + # replace values with prev_latent_keyframes + for latent_keyframe in prev_latent_keyframe.keyframes: + curr_latent_keyframe.add(latent_keyframe) + + return (curr_latent_keyframe,) diff --git a/imports/AdvancedControlNet/logger.py b/imports/AdvancedControlNet/logger.py new file mode 100644 index 0000000..b23b82f --- /dev/null +++ b/imports/AdvancedControlNet/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/imports/AdvancedControlNet/nodes.py b/imports/AdvancedControlNet/nodes.py new file mode 100644 index 0000000..011ad63 --- /dev/null +++ b/imports/AdvancedControlNet/nodes.py @@ -0,0 +1,194 @@ +import numpy as np +from torch import Tensor + +import folder_paths + +from .control import load_controlnet, convert_to_advanced, ControlWeightsImport, ControlWeightTypeImport,\ + LatentKeyframeGroupImport, TimestepKeyframeImport, TimestepKeyframeGroupImport, is_advanced_controlnet +from .control import StrengthInterpolationImport as SI +from .weight_nodes import DefaultWeightsImport, ScaledSoftMaskedUniversalWeightsImport, ScaledSoftUniversalWeightsImport, SoftControlNetWeightsImport, CustomControlNetWeightsImport, \ + SoftT2IAdapterWeightsImport, CustomT2IAdapterWeightsImport +from .latent_keyframe_nodes import LatentKeyframeGroupNodeImport, LatentKeyframeInterpolationNodeImport, LatentKeyframeBatchedGroupNodeImport, LatentKeyframeNodeImport +from .logger import logger + + +class TimestepKeyframeNodeImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + }, + "optional": { + "prev_timestep_kf": ("TIMESTEP_KEYFRAME", ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "cn_weights": ("CONTROL_NET_WEIGHTS", ), + "latent_keyframe": ("LATENT_KEYFRAME", ), + "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "inherit_missing": ("BOOLEAN", {"default": True}, ), + "guarantee_usage": ("BOOLEAN", {"default": True}, ), + "mask_optional": ("MASK", ), + #"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], {"default": SI.NONE}, ), + } + } + + RETURN_NAMES = ("TIMESTEP_KF", ) + RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) + FUNCTION = "load_keyframe" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + start_percent: float, + strength: float=1.0, + cn_weights: ControlWeightsImport=None, control_net_weights: ControlWeightsImport=None, # old name + latent_keyframe: LatentKeyframeGroupImport=None, + prev_timestep_kf: TimestepKeyframeGroupImport=None, prev_timestep_keyframe: TimestepKeyframeGroupImport=None, # old name + null_latent_kf_strength: float=0.0, + inherit_missing=True, + guarantee_usage=True, + mask_optional=None, + interpolation: str=SI.NONE,): + control_net_weights = control_net_weights if control_net_weights else cn_weights + prev_timestep_keyframe = prev_timestep_keyframe if prev_timestep_keyframe else prev_timestep_kf + if not prev_timestep_keyframe: + prev_timestep_keyframe = TimestepKeyframeGroupImport() + else: + prev_timestep_keyframe = prev_timestep_keyframe.clone() + keyframe = TimestepKeyframeImport(start_percent=start_percent, strength=strength, interpolation=interpolation, null_latent_kf_strength=null_latent_kf_strength, + control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage, + mask_hint_orig=mask_optional) + prev_timestep_keyframe.add(keyframe) + return (prev_timestep_keyframe,) + + +class ControlNetLoaderAdvancedImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + }, + "optional": { + "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" + + def load_controlnet(self, control_net_name, + timestep_keyframe: TimestepKeyframeGroupImport=None + ): + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + controlnet = load_controlnet(controlnet_path, timestep_keyframe) + return (controlnet,) + + +class DiffControlNetLoaderAdvancedImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "control_net_name": (folder_paths.get_filename_list("controlnet"), ) + }, + "optional": { + "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" + + def load_controlnet(self, control_net_name, model, + timestep_keyframe: TimestepKeyframeGroupImport=None + ): + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) + if is_advanced_controlnet(controlnet): + controlnet.verify_all_weights() + return (controlnet,) + + +class AdvancedControlNetApplyImport: + @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": { + "mask_optional": ("MASK", ), + "timestep_kf": ("TIMESTEP_KEYFRAME", ), + "latent_kf_override": ("LATENT_KEYFRAME", ), + "weights_override": ("CONTROL_NET_WEIGHTS", ), + } + } + + RETURN_TYPES = ("CONDITIONING","CONDITIONING") + RETURN_NAMES = ("positive", "negative") + FUNCTION = "apply_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" + + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, + mask_optional: Tensor=None, + timestep_kf: TimestepKeyframeGroupImport=None, latent_kf_override: LatentKeyframeGroupImport=None, + weights_override: ControlWeightsImport=None): + if strength == 0: + return (positive, negative) + + control_hint = image.movedim(-1,1) + cnets = {} + + out = [] + for conditioning in [positive, negative]: + c = [] + 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: + # copy, convert to advanced if needed, and set cond + c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent)) + if is_advanced_controlnet(c_net): + # apply optional parameters and overrides, if provided + if timestep_kf is not None: + c_net.set_timestep_keyframes(timestep_kf) + if latent_kf_override is not None: + c_net.latent_keyframe_override = latent_kf_override + if weights_override is not None: + c_net.weights_override = weights_override + # verify weights are compatible + c_net.verify_all_weights() + # set cond hint mask + if mask_optional is not None: + mask_optional = mask_optional.clone() + # if not in the form of a batch, make it so + if len(mask_optional.shape) < 3: + mask_optional = mask_optional.unsqueeze(0) + c_net.set_cond_hint_mask(mask_optional) + 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) + out.append(c) + return (out[0], out[1]) + + diff --git a/imports/AdvancedControlNet/reference_nodes.py b/imports/AdvancedControlNet/reference_nodes.py new file mode 100644 index 0000000..6879f97 --- /dev/null +++ b/imports/AdvancedControlNet/reference_nodes.py @@ -0,0 +1,12 @@ +class AnimateDiffLoaderWithContext: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("MODEL",) + CATEGORY = "" \ No newline at end of file diff --git a/imports/AdvancedControlNet/weight_nodes.py b/imports/AdvancedControlNet/weight_nodes.py new file mode 100644 index 0000000..c52e907 --- /dev/null +++ b/imports/AdvancedControlNet/weight_nodes.py @@ -0,0 +1,201 @@ +from torch import Tensor +import torch +from .control import TimestepKeyframeImport, TimestepKeyframeGroupImport, ControlWeightsImport, get_properly_arranged_t2i_weights, linear_conversion +from .logger import logger + + +WEIGHTS_RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT") + + +class DefaultWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self): + weights = ControlWeightsImport.default() + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class ScaledSoftMaskedUniversalWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mask": ("MASK", ), + "min_base_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + "max_base_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + #"lock_min": ("BOOLEAN", {"default": False}, ), + #"lock_max": ("BOOLEAN", {"default": False}, ), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False): + # normalize mask + mask = mask.clone() + x_min = 0.0 if lock_min else mask.min() + x_max = 1.0 if lock_max else mask.max() + if x_min == x_max: + mask = torch.ones_like(mask) * max_base_multiplier + else: + mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) + weights = ControlWeightsImport.universal_mask(weight_mask=mask) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class ScaledSoftUniversalWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self, base_multiplier, flip_weights): + weights = ControlWeightsImport.universal(base_multiplier=base_multiplier, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class SoftControlNetWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "weight_00": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_01": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_02": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_03": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_04": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_05": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_06": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_07": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_08": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_09": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" + + 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] + weights = ControlWeightsImport.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class CustomControlNetWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_04": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_05": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_06": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_07": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_08": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_09": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + } + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" + + 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] + weights = ControlWeightsImport.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class SoftT2IAdapterWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" + + def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): + weights = [weight_00, weight_01, weight_02, weight_03] + weights = get_properly_arranged_t2i_weights(weights) + weights = ControlWeightsImport.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) + + +class CustomT2IAdapterWeightsImport: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "flip_weights": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" + + def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): + weights = [weight_00, weight_01, weight_02, weight_03] + weights = get_properly_arranged_t2i_weights(weights) + weights = ControlWeightsImport.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroupImport.default(TimestepKeyframeImport(control_weights=weights))) diff --git a/imports/IPAdapterPlus.py b/imports/IPAdapterPlus.py index 2251521..573bdc3 100644 --- a/imports/IPAdapterPlus.py +++ b/imports/IPAdapterPlus.py @@ -167,10 +167,10 @@ def zeroed_hidden_states(clip_vision, batch_size): precision_scope = lambda a, b: contextlib.nullcontext(a) with precision_scope(comfy.model_management.get_autocast_device(clip_vision.load_device), torch.float32): - outputs = clip_vision.model(pixel_values, output_hidden_states=True) + outputs = clip_vision.model(pixel_values, intermediate_output=-2) # we only need the penultimate hidden states - outputs = outputs['hidden_states'][-2].cpu() if 'hidden_states' in outputs else None + outputs = outputs[1].to(comfy.model_management.intermediate_device()) return outputs @@ -433,41 +433,13 @@ class CrossAttentionPatchImport: return out.to(dtype=org_dtype) -class IPAdapterModelLoaderImport: - @classmethod - def INPUT_TYPES(s): - return {"required": { "ipadapter_file": (folder_paths.get_filename_list("ipadapter"), )}} - RETURN_TYPES = ("IPADAPTER",) - FUNCTION = "load_ipadapter_model" - - CATEGORY = "ipadapter" - - def load_ipadapter_model(self, ipadapter_file): - ckpt_path = folder_paths.get_full_path("ipadapter", ipadapter_file) - - model = comfy.utils.load_torch_file(ckpt_path, safe_load=True) - - if ckpt_path.lower().endswith(".safetensors"): - st_model = {"image_proj": {}, "ip_adapter": {}} - for key in model.keys(): - if key.startswith("image_proj."): - st_model["image_proj"][key.replace("image_proj.", "")] = model[key] - elif key.startswith("ip_adapter."): - st_model["ip_adapter"][key.replace("ip_adapter.", "")] = model[key] - model = st_model - - if not "ip_adapter" in model.keys() or not model["ip_adapter"]: - raise Exception("invalid IPAdapter model {}".format(ckpt_path)) - - return (model,) class IPAdapterApplyImport: @classmethod def INPUT_TYPES(s): return { "required": { - "ipadapter": ("IPADAPTER", ), "clip_vision": ("CLIP_VISION",), "image": ("IMAGE",), @@ -489,8 +461,6 @@ class IPAdapterApplyImport: CATEGORY = "ipadapter" def apply_ipadapter(self, ipadapter, model, weight, clip_vision=None, image=None, weight_type="original", noise=None, embeds=None, attn_mask=None, start_at=0.0, end_at=1.0, unfold_batch=False): - for attr in list(self.__dict__): - delattr(self, attr) self.dtype = model.model.diffusion_model.dtype self.device = comfy.model_management.get_torch_device() self.weight = weight @@ -763,6 +733,7 @@ class IPAdapterEncoderImport: + class IPAdapterBatchEmbedsImport: @classmethod def INPUT_TYPES(s):