Bug fixing

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