diff --git a/README.md b/README.md index ed3607b..f0ef159 100644 --- a/README.md +++ b/README.md @@ -1,14 +1,151 @@ # ComfyUI-Advanced-ControlNet -These custom nodes allow for scheduling ControlNet strength across latents in the same batch (WORKING) and across timesteps (IN PROGRESS). +Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, and ControlLoRAs. Kohya Controllllite support coming soon. -Custom weights can also be applied to ControlNets and T2IAdapters to mimic the "My prompt is more important" functionality in AUTOMATIC1111's ControlNet extension. +Custom weights allow replication of the "My prompt is more important" feature of Auto1111's sd-webui ControlNet extension. -TODO: -- Other handy nodes -- Finish and update this README for other workflows +ControlNet preprocessors are available through [comfyui_controlnet_aux](https://github.com/Fannovel16/comfyui_controlnet_aux) nodes -## Workflows +## Features +- Timestep and latent strength scheduling +- Attention masks +- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling. +- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows. -### AnimateDiff Workflows -***Latent Keyframes*** identify which latents in a batch the ControlNet should apply to, and at what strength. They connect to a ***Timestep Keyframe*** to identify at what point in the generation to kick in (for basic use, start_percent on the Timestep Keyframe should be 0.0). Latent Keyframe nodes can be chained to apply the ControlNet to multiple keyframes at various strengths. +## Table of Contents: +- [Scheduling Explanation](#scheduling-explanation) +- [Nodes](#nodes) +- [Usage](#usage) (will fill this out soon) + +# Scheduling Explanation + +The two core concepts for scheduling are ***Timestep Keyframes*** and ***Latent Keyframes***. + +***Timestep Keyframes*** hold the values that guide the settings for a controlnet, and begin to take effect based on their start_percent, which corresponds to the percentage of the sampling process. They can contain masks for the strengths of each latent, control_net_weights, and latent_keyframes (specific strengths for each latent), all optional. + +***Latent Keyframes*** determine the strength of the controlnet for specific latents - all they contain is the batch_index of the latent, and the strength the controlnet should apply for that latent. As a concept, latent keyframes achieve the same affect as a uniform mask with the chosen strength value. + +![advcn_image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/e6275264-6c3f-4246-a319-111ee48f4cd9) + +# Nodes + +The ControlNet nodes provided here are the ***Apply Advanced ControlNet*** and ***Load Advanced ControlNet Model*** (or diff) nodes. The vanilla ControlNet nodes are also compatible, and can be used almost interchangeably - the only difference is that **at least one of these nodes must be used** for Advanced versions of ControlNets to be used (important for sliding context sampling, like with AnimateDiff-Evolved). + +Key: +- 🟩 - required inputs +- 🟨 - optional inputs +- 🟦 - start as widgets, can be converted to inputs +- 🟥 - optional input/output, but not recommended to use unless needed +- 🟪 - output + +## Apply Advanced ControlNet +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/dc541d41-70df-4a71-b832-efa65af98f06) + +Same functionality as the vanilla Apply Advanced ControlNet (Advanced) node, except with Advanced ControlNet features added to it. Automatically converts any ControlNet from ControlNet loaders into Advanced versions. + +### Inputs +- 🟩***positive***: conditioning (positive). +- 🟩***negative***: conditioning (negative). +- 🟩***control_net***: loaded controlnet; will be converted to Advanced version automatically by this node, if it's a supported type. +- 🟩***image***: images to guide controlnets - if the loaded controlnet requires it, they must preprocessed images. If one image provided, will be used for all latents. If more images provided, will use each image separately for each latent. If not enough images to meet latent count, will repeat the images from the beginning to match vanilla ControlNet functionality. +- 🟨***mask_optional***: attention masks to apply to controlnets; basically, decides what part of the image the controlnet to apply to (and the relative strength, if the mask is not binary). Same as image input, if you provide more than one mask, each can apply to a different latent. +- 🟨***timestep_kf***: timestep keyframes to guide controlnet effect throughout sampling steps. +- 🟨***latent_kf_override***: override for latent keyframes, useful if no other features from timestep keyframes is needed. *NOTE: this latent keyframe will be applied to ALL timesteps, regardless if there are other latent keyframes attached to connected timestep keyframes.* +- 🟨***weights_override***: override for weights, useful if no other features from timestep keyframes is needed. *NOTE: this weight will be applied to ALL timesteps, regardless if there are other weights attached to connected timestep keyframes.* +- 🟦***strength***: strength of controlnet; 1.0 is full strength, 0.0 is no effect at all. +- 🟦***start_percent***: sampling step percentage at which controlnet should start to be applied - no matter what start_percent is set on timestep keyframes, they won't take effect until this start_percent is reached. +- 🟦***stop_percent***: sampling step percentage at which controlnet should stop being applied - no matter what start_percent is set on timestep keyframes, they won't take effect once this end_percent is reached. + +### Outputs +- 🟪***positive***: conditioning (positive) with applied controlnets +- 🟪***negative***: conditioning (negative) with applied controlnets + +## Load Advanced ControlNet Model +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/4a7f58a9-783d-4da4-bf82-bc9c167e4722) + +Loads a ControlNet model and converts it into an Advanced version that supports all the features in this repo. When used with **Apply Advanced ControlNet** node, there is no reason to use the timestep_keyframe input on this node - use timestep_kf on the Apply node instead. + +### Inputs +- 🟥***timestep_keyframe***: optional and likely unnecessary input to have ControlNet use selected timestep_keyframes - should not be used unless you need to. Useful if this node is not attached to **Apply Advanced ControlNet** node, but still want to use Timestep Keyframe, or to use TK_SHORTCUT outputs from ControlWeights in the same scenario. Will be overriden by the timestep_kf input on **Apply Advanced ControlNet** node, if one is provided there. +- 🟨***model***: model to plug into the diff version of the node. Some controlnets are designed for receive the model; if you don't know what this does, you probably don't want tot use the diff version of the node. + +### Outputs +- 🟪***CONTROL_NET***: loaded Advanced ControlNet + +## Timestep Keyframe +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/c6f2a86e-fc96-4f8b-b976-7c2062a6eba2) + +Scheduling node across timesteps (sampling steps) based on the set start_percent. Chaining Timestep Keyframes allows ControlNet scheduling across sampling steps (percentage-wise), through a timestep keyframe schedule. + +### Inputs +- 🟨***prev_timestep_kf***: used to chain Timestep Keyframes together to create a schedule. The order does not matter - the Timestep Keyframes sort themselves automatically by their start_percent. *Any Timestep Keyframe contained in the prev_timestep_keyframe that contains the same start_percent as the Timestep Keyframe will be overwritten.* +- 🟨***cn_weights***: weights to apply to controlnet while this Timestep Keyframe is in effect. Must be compatible with the loaded controlnet, or will throw an error explaining what weight types are compatible. If inherit_missing is True, if no control_net_weight is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a weight_override, the weight_override will be used during sampling instead of control_net_weight.* +- 🟨***latent_keyframe***: latent keyframes to apply to controlnet while this Timestep Keyframe is in effect. If inherit_missing is True, if no latent_keyframe is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a latent_kf_override, the latent_lf_override will be used during sampling instead of latent_keyframe.* +- 🟨***mask_optional***: attention masks to apply to controlnets; basically, decides what part of the image the controlnet to apply to (and the relative strength, if the mask is not binary). Same as mask_optional on the Apply Advanced ControlNet node, can apply either one maks to all latents, or individual masks for each latent. If inherit_missing is True, if no mask_optional is passed in, will attempt to reuse the last-used mask_optional in the timestep keyframe schedule. It is NOT overriden by mask_optional on the Apply Advanced ControlNet node; will be used together. +- 🟦***start_percent***: sampling step percentage at which this Timestep Keyframe qualifies to be used. Acts as the 'key' for the Timestep Keyframe in the timestep keyframe schedule. +- 🟦***strength***: strength of the controlnet; multiplies the controlnet by this value, basically, applied alongside the strength on the Apply ControlNet node. If set to 0.0 will not have any effect during the duration of this Timestep Keyframe's effect, and will increase sampling speed by not doing any work. +- 🟦***null_latent_kf_strength***: strength to assign to latents that are unaccounted for in the passed in latent_keyframes. Has no effect if no latent_keyframes are passed in, or no batch_indeces are unaccounted in the latent_keyframes for during sampling. +- 🟦***inherit_missing***: determines if should reuse values from previous Timestep Keyframes for optional values (control_net_weights, latent_keyframe, and mask_option) that are not included on this TimestepKeyframe. To inherit only specific inputs, use default inputs. +- 🟦***guarantee_usage***: when true, even if a Timestep Keyframe's start_percent ahead of this one in the schedule is closer to current sampling percentage, this Timestep Keyframe will still be used for one step before moving on to the next selected Timestep Keyframe in the following step. Whether the Timestep Keyframe is used or not, its inputs will still be accounted for inherit_missing purposes. + +### Outputs +- 🟪***TIMESTEP_KF***: the created Timestep Keyframe, that can either be linked to another or into a Timestep Keyframe input. + +## Latent Keyframe +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/7eb2cc4c-255c-4f32-b09b-699f713fada3) + +A singular Latent Keyframe, selects the strength for a specific batch_index. If batch_index is not present during sampling, will simply have no effect. Can be chained with any other Latent Keyframe-type node to create a latent keyframe schedule. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If a Latent Keyframe contained in prev_latent_keyframes have the same batch_index as this Latent Keyframe, they will take priority over this node's value.* +- 🟦***batch_index***: index of latent in batch to apply controlnet strength to. Acts as the 'key' for the Latent Keyframe in the latent keyframe schedule. +- 🟦***strength***: strength of controlnet to apply to the corresponding latent. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +## Latent Keyframe Group +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/5ce3b795-f5fc-4dc3-ae30-a4c7f87e278c) + +Allows to create Latent Keyframes via individual indeces or python-style ranges. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟨***latent_optional***: the latents expected to be passed in for sampling; only required if you wish to use negative indeces (will be automatically converted to real values). +- 🟦***index_strengths***: string list of indeces or python-style ranges of indeces to assign strengths to. If latent_optional is passed in, can contain negative indeces or ranges that contain negative numbers, python-style. The different indeces must be comma separated. Individual latents can be specified by ```batch_index=strength```, like ```0=0.9```. Ranges can be specified by ```start_index_inclusive:end_index_exclusive=strength```, like ```0:8=strength```. Negative indeces are possible when latents_optional has an input, with a string such as ```0,-4=0.25```. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +## Latent Keyframe Interpolation +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/7986c737-83b9-46bc-aab0-ae4c368df446) + +Allows to create Latent Keyframes with interpolated values in a range. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟦***batch_index_from***: starting batch_index of range, included. +- 🟦***batch_index_to***: end batch_index of range, excluded (python-style range). +- 🟦***strength_from***: starting strength of interpolation. +- 🟦***strength_to***: end strength of interpolation. +- 🟦***interpolation***: the method of interpolation. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +## Latent Keyframe Batched Group +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/6cec701f-6183-4aeb-af5c-cac76f5591b7) + +Allows to create Latent Keyframes via a list of floats, such as with Batch Value Schedule from [ComfyUI_FizzNodes](https://github.com/FizzleDorf/ComfyUI_FizzNodes) nodes. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟦***float_strengths***: a list of floats, that will correspond to the strength of each Latent Keyframe; the batch_index is the index of each float value in the list. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +# There are more nodes to document and show usage - will add this soon! TODO diff --git a/control/control.py b/control/control.py index c851fd2..18cbffd 100644 --- a/control/control.py +++ b/control/control.py @@ -4,11 +4,87 @@ import torch import comfy.utils import comfy.controlnet as comfy_cn -from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to +from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to -ControlNetWeightsType = list[float] -T2IAdapterWeightsType = list[float] +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 ControlWeightType: + DEFAULT = "default" + UNIVERSAL = "universal" + T2IADAPTER = "t2iadapter" + CONTROLNET = "controlnet" + CONTROLLORA = "controllora" + CONTROLLLLITE = "controllllite" + + +class ControlWeights: + 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(ControlWeightType.DEFAULT) + + @classmethod + def universal(cls, base_multiplier: float, flip_weights: bool=False): + return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights) + + @classmethod + def universal_mask(cls, weight_mask: Tensor): + return cls(ControlWeightType.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(ControlWeightType.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(ControlWeightType.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(ControlWeightType.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(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights) + + +class StrengthInterpolation: + LINEAR = "linear" + EASE_IN = "ease-in" + EASE_OUT = "ease-out" + EASE_IN_OUT = "ease-in-out" + NONE = "none" class LatentKeyframe: @@ -46,19 +122,43 @@ class LatentKeyframeGroup: def is_empty(self) -> bool: return len(self.keyframes) == 0 + def clone(self) -> 'LatentKeyframeGroup': + cloned = LatentKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + class TimestepKeyframe: def __init__(self, start_percent: float = 0.0, - control_net_weights: ControlNetWeightsType = None, - t2i_adapter_weights: T2IAdapterWeightsType = None, + strength: float = 1.0, + interpolation: str = StrengthInterpolation.NONE, + control_weights: ControlWeights = None, latent_keyframes: LatentKeyframeGroup = None, - default_latent_strength: float = 0.0) -> 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.control_net_weights = control_net_weights - self.t2i_adapter_weights = t2i_adapter_weights + self.start_t = 999999999.9 + self.strength = strength + self.interpolation = interpolation + self.control_weights = control_weights self.latent_keyframes = latent_keyframes - self.default_latent_strength = default_latent_strength + 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 @@ -90,12 +190,24 @@ class TimestepKeyframeGroup: except IndexError: return None + def has_index(self, index: int) -> int: + return index >=0 and index < len(self.keyframes) + def __getitem__(self, index) -> TimestepKeyframe: 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) -> 'TimestepKeyframeGroup': + cloned = TimestepKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + @classmethod def default(cls, keyframe: TimestepKeyframe) -> 'TimestepKeyframeGroup': group = cls() @@ -104,83 +216,358 @@ class TimestepKeyframeGroup: # used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function -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 -class ControlNetAdvanced(ControlNet): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, 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 TimestepKeyframeGroup() - 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 +class AdvancedControlBase: + def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights): + self.base = base + self.compatible_weights = [ControlWeightType.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 - # override control_merge - self.control_merge = control_merge_inject.__get__(self, type(self)) + # timesteps + self.t: Tensor = None + self.batched_number: int = None + # weights + override + self.weights: ControlWeights = None + self.weights_default: ControlWeights = weights_default + self.weights_override: ControlWeights = None + # latent keyframe + override + self.latent_keyframes: LatentKeyframeGroup = None + self.latent_keyframe_override: LatentKeyframeGroup = 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 WeightTypeException(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 WeightTypeException(msg) + + def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroup): + self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() + # 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 AdvancedControlBase 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 == ControlWeightType.DEFAULT: + self.weights = self.weights_default + elif self.weights.weight_type == ControlWeightType.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) -> ControlWeights: + return self.weights 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 + 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: 'AdvancedControlBase', 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: 'AdvancedControlBase'): + 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 ControlNetAdvanced(ControlNet, AdvancedControlBase): + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None): + super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controlnet()) + + def get_universal_weights(self) -> ControlWeights: + raw_weights = [(self.weights.base_multiplier ** float(12 - i)) for i in range(13)] + return ControlWeights.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) @@ -204,31 +591,16 @@ class ControlNetAdvanced(ControlNet): 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: + # 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(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) + # prepare mask_cond_hint + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=self.control_model.dtype) context = cond['c_crossattn'] # uses 'y' in new ComfyUI update @@ -243,90 +615,49 @@ class ControlNetAdvanced(ControlNet): 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: TimestepKeyframe, 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 = ControlNetAdvanced(self.control_model, 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.sub_idxs = None - self.full_latent_length = 0 - self.context_length = 0 + + @staticmethod + def from_vanilla(v: ControlNet, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlNetAdvanced': + return ControlNetAdvanced(control_model=v.control_model, timestep_keyframes=timestep_keyframe, + global_average_pooling=v.global_average_pooling, device=v.device) -class T2IAdapterAdvanced(T2IAdapter): +class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroup, 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 TimestepKeyframeGroup() - 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 + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.t2iadapter()) + + def get_universal_weights(self) -> ControlWeights: + 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 ControlWeights.t2iadapter(raw_weights, self.weights.flip_weights) + + def get_calc_pow(self, idx: int, layers: int) -> int: + # match how T2IAdapterAdvanced 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: @@ -334,38 +665,85 @@ class T2IAdapterAdvanced(T2IAdapter): self.cond_hint_original = full_cond_hint_original del full_cond_hint_original - def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): - # For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters - # TODO: support masks - return - def copy(self): c = T2IAdapterAdvanced(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.sub_idxs = None - self.full_latent_length = 0 - self.context_length = 0 + self.cleanup_advanced() + + @staticmethod + def from_vanilla(v: T2IAdapter, timestep_keyframe: TimestepKeyframeGroup=None) -> 'T2IAdapterAdvanced': + return T2IAdapterAdvanced(t2i_model=v.t2i_model, timestep_keyframes=timestep_keyframe, channels_in=v.channels_in, device=v.device) + + +class ControlLoraAdvanced(ControlLora, AdvancedControlBase): + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None): + super().__init__(control_weights=control_weights, global_average_pooling=global_average_pooling, device=device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllora()) + # use some functions from ControlNetAdvanced + self.get_control_advanced = ControlNetAdvanced.get_control_advanced.__get__(self, type(self)) + self.sliding_get_control = ControlNetAdvanced.sliding_get_control.__get__(self, type(self)) + + def get_universal_weights(self) -> ControlWeights: + raw_weights = [(self.weights.base_multiplier ** float(9 - i)) for i in range(10)] + return ControlWeights.controllora(raw_weights, self.weights.flip_weights) + + def copy(self): + c = ControlLoraAdvanced(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: TimestepKeyframeGroup=None) -> 'ControlLoraAdvanced': + return ControlLoraAdvanced(control_weights=v.control_weights, timestep_keyframes=timestep_keyframe, + global_average_pooling=v.global_average_pooling, device=v.device) + + +class ControlLLLiteAdvanced(ControlNet, AdvancedControlBase): + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, device=None): + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite()) def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=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: TimestepKeyframeGroup=None): + # if already advanced, leave it be + if is_advanced_controlnet(control): + return control # if exactly ControlNet returned, transform it into ControlNetAdvanced if type(control) == ControlNet: - return ControlNetAdvanced(control.control_model, timestep_keyframe, global_average_pooling=control.global_average_pooling) + return ControlNetAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) + # if exactly ControlLora returned, transform it into ControlLoraAdvanced + elif type(control) == ControlLora: + return ControlLoraAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) # if T2IAdapter returned, transform it into T2IAdapterAdvanced elif isinstance(control, T2IAdapter): - return T2IAdapterAdvanced(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 T2IAdapterAdvanced.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 isinstance(input_object, ControlNetAdvanced) or isinstance(input_object, T2IAdapterAdvanced) + return hasattr(input_object, "sub_idxs") # adapted from comfy/sample.py @@ -375,3 +753,18 @@ def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim 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 WeightTypeException(TypeError): + "Raised when weight not compatible with AdvancedControlBase object" + pass diff --git a/control/control_lllite.py b/control/control_lllite.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/control/control_lllite.py @@ -0,0 +1 @@ + diff --git a/control/deprecated_nodes.py b/control/deprecated_nodes.py index 25b7169..a64ac9b 100644 --- a/control/deprecated_nodes.py +++ b/control/deprecated_nodes.py @@ -4,6 +4,7 @@ import torch import numpy as np from PIL import Image, ImageOps +from .control import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe from .logger import logger @@ -68,3 +69,35 @@ class LoadImagesFromDirectory: 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/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index d12935c..2fde61e 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -3,6 +3,7 @@ import numpy as np from collections.abc import Iterable from .control import LatentKeyframe, LatentKeyframeGroup +from .control import StrengthInterpolation as SI from .logger import logger @@ -12,13 +13,14 @@ class LatentKeyframeNode: 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.00001}, ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframe" @@ -27,9 +29,14 @@ class LatentKeyframeNode: def load_keyframe(self, batch_index: int, strength: float, - prev_latent_keyframe: LatentKeyframeGroup=None): + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=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 = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() keyframe = LatentKeyframe(batch_index, strength) prev_latent_keyframe.add(keyframe) return (prev_latent_keyframe,) @@ -43,11 +50,13 @@ class LatentKeyframeGroupNode: "index_strengths": ("STRING", {"multiline": True, "default": ""}), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), "latent_optional": ("LATENT", ), + "print_keyframes": ("BOOLEAN", {"default": False}) } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframes" @@ -80,7 +89,7 @@ class LatentKeyframeGroupNode: def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframe]: if not latent_indeces: return set() - all_indeces = [i for i in range(0, latent_count)] + 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 ':' @@ -104,8 +113,14 @@ class LatentKeyframeGroupNode: index_range = [r.strip() for r in index_range] start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative) end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative) - for i in all_indeces[start_index:end_index]: - chosen_indeces.add(LatentKeyframe(i, strength)) + # 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(LatentKeyframe(i, strength)) + # otherwise, assume indeces are valid + else: + for i in range(start_index, end_index): + chosen_indeces.add(LatentKeyframe(i, strength)) # parse individual indeces else: chosen_indeces.add(LatentKeyframe(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength)) @@ -113,10 +128,15 @@ class LatentKeyframeGroupNode: def load_keyframes(self, index_strengths: str, - prev_latent_keyframe: LatentKeyframeGroup=None, - latent_image_opt=None): + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=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 = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() latent_count = -1 @@ -125,9 +145,13 @@ class LatentKeyframeGroupNode: latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count) for latent_keyframe in latent_keyframes: - logger.info(f"keyframe {latent_keyframe.batch_index}:{latent_keyframe.strength}") 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) @@ -141,15 +165,17 @@ class LatentKeyframeInterpolationNode: "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"], ), + "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT], ), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "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" @@ -160,7 +186,9 @@ class LatentKeyframeInterpolationNode: batch_index_to_excl: int, strength_to: float, interpolation: str, - prev_latent_keyframe: LatentKeyframeGroup=None): + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=None, # old name + print_keyframes=False): if (batch_index_from > batch_index_to_excl): raise ValueError("batch_index_from must be less than or equal to batch_index_to.") @@ -168,28 +196,34 @@ class LatentKeyframeInterpolationNode: if (batch_index_from < 0 and batch_index_to_excl >= 0): raise ValueError("batch_index_from and batch_index_to must be either both positive or both negative.") + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() steps = batch_index_to_excl - batch_index_from diff = strength_to - strength_from - if interpolation == "linear": + if interpolation == SI.LINEAR: weights = np.linspace(strength_from, strength_to, steps) - elif interpolation == "ease-in": + elif interpolation == SI.EASE_IN: index = np.linspace(0, 1, steps) weights = diff * np.power(index, 2) + strength_from - elif interpolation == "ease-out": + elif interpolation == SI.EASE_OUT: index = np.linspace(0, 1, steps) weights = diff * (1 - np.power(1 - index, 2)) + strength_from - elif interpolation == "ease-in-out": + elif interpolation == SI.EASE_IN_OUT: index = np.linspace(0, 1, steps) weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from for i in range(steps): keyframe = LatentKeyframe(batch_index_from + i, float(weights[i])) - logger.info(f"keyframe {batch_index_from + i}:{weights[i]}") curr_latent_keyframe.add(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: @@ -203,33 +237,44 @@ class LatentKeyframeBatchedGroupNode: def INPUT_TYPES(s): return { "required": { - "strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001}), + "float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "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, strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroup=None): + def load_keyframe(self, float_strengths: Union[float, list[float]], + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=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 = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() # if received a normal float input, do nothing - if type(strengths) in (float, int): - logger.info("No batched strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.") + 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(strengths, Iterable): - for idx, strength in enumerate(strengths): + elif isinstance(float_strengths, Iterable): + for idx, strength in enumerate(float_strengths): keyframe = LatentKeyframe(idx, strength) curr_latent_keyframe.add(keyframe) - logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}") else: - raise ValueError(f"Expected strengths to be an iterable input, but was {type(strengths).__repr__}.") + 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: diff --git a/control/nodes.py b/control/nodes.py index e3d4f64..3794ac7 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -1,10 +1,12 @@ import numpy as np +from torch import Tensor import folder_paths -from .control import ControlNetAdvanced, T2IAdapterAdvanced, load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ +from .control import load_controlnet, convert_to_advanced, ControlWeights, ControlWeightType,\ LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet -from .weight_nodes import ScaledSoftControlNetWeights, SoftControlNetWeights, CustomControlNetWeights, \ +from .control import StrengthInterpolation as SI +from .weight_nodes import DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode from .deprecated_nodes import LoadImagesFromDirectory @@ -19,13 +21,19 @@ class TimestepKeyframeNode: "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", ), + "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", ), - "prev_timestep_keyframe": ("TIMESTEP_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" @@ -33,13 +41,24 @@ class TimestepKeyframeNode: def load_keyframe(self, start_percent: float, - control_net_weights: ControlNetWeightsType=None, - t2i_adapter_weights: T2IAdapterWeightsType=None, + strength: float=1.0, + cn_weights: ControlWeights=None, control_net_weights: ControlWeights=None, # old name latent_keyframe: LatentKeyframeGroup=None, - prev_timestep_keyframe: TimestepKeyframeGroup=None): + prev_timestep_kf: TimestepKeyframeGroup=None, prev_timestep_keyframe: TimestepKeyframeGroup=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 = TimestepKeyframeGroup() - keyframe = TimestepKeyframe(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe) + else: + prev_timestep_keyframe = prev_timestep_keyframe.clone() + keyframe = TimestepKeyframe(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,) @@ -55,13 +74,15 @@ class ControlNetLoaderAdvanced: "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), } } - + RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name, + timestep_keyframe: TimestepKeyframeGroup=None + ): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe) return (controlnet,) @@ -83,11 +104,15 @@ class DiffControlNetLoaderAdvanced: RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup, model): + def load_controlnet(self, control_net_name, model, + timestep_keyframe: TimestepKeyframeGroup=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,) @@ -106,6 +131,9 @@ class AdvancedControlNetApply: }, "optional": { "mask_optional": ("MASK", ), + "timestep_kf": ("TIMESTEP_KEYFRAME", ), + "latent_kf_override": ("LATENT_KEYFRAME", ), + "weights_override": ("CONTROL_NET_WEIGHTS", ), } } @@ -113,9 +141,12 @@ class AdvancedControlNetApply: RETURN_NAMES = ("positive", "negative") FUNCTION = "apply_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=None): + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, + mask_optional: Tensor=None, + timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None, + weights_override: ControlWeights=None): if strength == 0: return (positive, negative) @@ -132,10 +163,21 @@ class AdvancedControlNetApply: 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): + # 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) @@ -159,17 +201,19 @@ NODE_CLASS_MAPPINGS = { "LatentKeyframeGroup": LatentKeyframeGroupNode, "LatentKeyframeBatchedGroup": LatentKeyframeBatchedGroupNode, "LatentKeyframeTiming": LatentKeyframeInterpolationNode, + # Conditioning + "ACN_AdvancedControlNetApply": AdvancedControlNetApply, # Loaders "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, - # Conditioning - "ACN_AdvancedControlNetApply": AdvancedControlNetApply, # Weights - "ScaledSoftControlNetWeights": ScaledSoftControlNetWeights, + "ScaledSoftControlNetWeights": ScaledSoftUniversalWeights, + "ScaledSoftMaskedUniversalWeights": ScaledSoftMaskedUniversalWeights, "SoftControlNetWeights": SoftControlNetWeights, "CustomControlNetWeights": CustomControlNetWeights, "SoftT2IAdapterWeights": SoftT2IAdapterWeights, "CustomT2IAdapterWeights": CustomT2IAdapterWeights, + "ACN_DefaultUniversalWeights": DefaultWeights, # Image "LoadImagesFromDirectory": LoadImagesFromDirectory } @@ -181,17 +225,19 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LatentKeyframeGroup": "Latent Keyframe Group 🛂🅐🅒🅝", "LatentKeyframeBatchedGroup": "Latent Keyframe Batched Group 🛂🅐🅒🅝", "LatentKeyframeTiming": "Latent Keyframe Interpolation 🛂🅐🅒🅝", - # Loaders - "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", - "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", # Conditioning "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", + # Loaders + "ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝", + "DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝", # Weights - "ScaledSoftControlNetWeights": "Scaled Soft ControlNet Weights 🛂🅐🅒🅝", - "SoftControlNetWeights": "Soft ControlNet Weights 🛂🅐🅒🅝", - "CustomControlNetWeights": "Custom ControlNet Weights 🛂🅐🅒🅝", - "SoftT2IAdapterWeights": "Soft T2IAdapter Weights 🛂🅐🅒🅝", - "CustomT2IAdapterWeights": "Custom T2IAdapter Weights 🛂🅐🅒🅝", + "ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftMaskedUniversalWeights": "Scaled Soft Masked Weights 🛂🅐🅒🅝", + "SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝", + "CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝", + "SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝", + "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", + "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # Image "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" } diff --git a/control/reference_nodes.py b/control/reference_nodes.py new file mode 100644 index 0000000..6879f97 --- /dev/null +++ b/control/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/control/weight_nodes.py b/control/weight_nodes.py index 2015c8a..f80d607 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -1,36 +1,80 @@ -from .control import TimestepKeyframe, TimestepKeyframeGroup +from torch import Tensor +import torch +from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion from .logger import logger -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 +WEIGHTS_RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT") -class ScaledSoftControlNetWeights: +class DefaultWeights: + @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 = ControlWeights.default() + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) + + +class ScaledSoftMaskedUniversalWeights: @classmethod def INPUT_TYPES(s): return { "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "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 = ControlWeights.universal_mask(weight_mask=mask) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) + + +class ScaledSoftUniversalWeights: + @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 = [(base_multiplier ** float(12 - i)) for i in range(13)] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.universal(base_multiplier=base_multiplier, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class SoftControlNetWeights: @@ -56,17 +100,17 @@ class SoftControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/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] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class CustomControlNetWeights: @@ -92,17 +136,17 @@ class CustomControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/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] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class SoftT2IAdapterWeights: @@ -118,17 +162,17 @@ class SoftT2IAdapterWeights: }, } - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/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] - if flip_weights: - weights.reverse() weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights))) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class CustomT2IAdapterWeights: @@ -144,14 +188,14 @@ class CustomT2IAdapterWeights: }, } - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/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] - if flip_weights: - weights.reverse() weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights))) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))