Merge PR #31 from Kosinkadink/develop Major features added + refactor

Major features added + refactor
This commit is contained in:
Jedrzej Kosinski
2023-11-30 10:46:16 -06:00
committed by GitHub
8 changed files with 984 additions and 273 deletions
+145 -8
View File
@@ -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
+566 -173
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
+33
View File
@@ -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,)
+74 -29
View File
@@ -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:
+77 -31
View File
@@ -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] 🛂🅐🅒🅝"
}
+12
View File
@@ -0,0 +1,12 @@
class AnimateDiffLoaderWithContext:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"image": ("IMAGE",),
},
}
RETURN_TYPES = ("MODEL",)
CATEGORY = ""
+76 -32
View File
@@ -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)))