Fix SparseCtrl Issue
This commit is contained in:
+2
-14
@@ -9,7 +9,6 @@ from PIL import Image
|
||||
import matplotlib.pyplot as plt
|
||||
# Local application/library specific imports
|
||||
from .imports.ComfyUI_IPAdapter_plus.IPAdapterPlus import IPAdapterBatchImport, IPAdapterTiledBatchImport, IPAdapterTiledImport, PrepImageForClipVisionImport, IPAdapterAdvancedImport, IPAdapterNoiseImport
|
||||
from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport
|
||||
from .imports.ComfyUI_Frame_Interpolation.vfi_models.film import FILM_VFIImport
|
||||
import matplotlib
|
||||
import gc
|
||||
@@ -47,7 +46,7 @@ class BatchCreativeInterpolationNode:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","CONDITIONING","CONDITIONING","MODEL","SPARSE_METHOD","INT", "INT", "STRING")
|
||||
RETURN_TYPES = ("IMAGE","CONDITIONING","CONDITIONING","MODEL","STRING","INT", "INT", "STRING")
|
||||
RETURN_NAMES = ("GRAPH","POSITIVE","NEGATIVE","MODEL","KEYFRAME_POSITIONS","BATCH_SIZE", "BUFFER","FRAMES_TO_DROP")
|
||||
FUNCTION = "combined_function"
|
||||
|
||||
@@ -299,11 +298,6 @@ class BatchCreativeInterpolationNode:
|
||||
shifted_keyframes_position = [position + buffer - 2 for position in keyframe_positions]
|
||||
shifted_keyframe_positions_string = ','.join(str(pos) for pos in shifted_keyframes_position)
|
||||
|
||||
# GET SPARSE INDEXES
|
||||
sparseindexmethod = SparseIndexMethodNodeImport()
|
||||
sparse_indexes, = sparseindexmethod.get_method(shifted_keyframe_positions_string)
|
||||
|
||||
# ADD BUFFER TO KEYFRAME POSITIONS
|
||||
if buffer > 0:
|
||||
# add front buffer
|
||||
keyframe_positions = [position + buffer - 1 for position in keyframe_positions]
|
||||
@@ -562,13 +556,7 @@ class BatchCreativeInterpolationNode:
|
||||
model, *_ = tiled_ipa_application.apply_tiled(model=model, ipadapter=ipadapter, image=torch.cat(bin.bigImageBatch, dim=0), weight=[x * detail_ipa_advanced_settings["ipa_weight"] for x in bin.weight_schedule], weight_type=detail_ipa_advanced_settings["ipa_weight_type"], start_at=detail_ipa_advanced_settings["ipa_starts_at"], end_at=detail_ipa_advanced_settings["ipa_ends_at"], clip_vision=clip_vision,sharpening=0.1,image_negative=negative_noise,embeds_scaling=detail_ipa_advanced_settings["ipa_embeds_scaling"], encode_batch_size=1, image_schedule=bin.image_schedule)
|
||||
|
||||
comparison_diagram, = plot_weight_comparison(all_cn_frame_numbers, all_cn_weights, all_ipa_frame_numbers, all_ipa_weights, buffer)
|
||||
return comparison_diagram, positive, negative, model, sparse_indexes, last_key_frame_position, buffer, shifted_keyframes_position
|
||||
|
||||
|
||||
# import the class FILM_VFI from ComfyUI-Frame-Interpolation/vfi_models/film/__init__.py
|
||||
|
||||
# from .imports.AdvancedControlNet.nodes_sparsectrl import SparseIndexMethodNodeImport
|
||||
|
||||
return comparison_diagram, positive, negative, model, shifted_keyframe_positions_string, last_key_frame_position, buffer, shifted_keyframes_position
|
||||
|
||||
class RemoveAndInterpolateFramesNode:
|
||||
@classmethod
|
||||
|
||||
@@ -1,773 +0,0 @@
|
||||
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
|
||||
@@ -1 +0,0 @@
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
#taken from: https://github.com/lllyasviel/ControlNet
|
||||
#and modified
|
||||
#and then taken from comfy/cldm/cldm.py and modified again
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import math
|
||||
import numpy as np
|
||||
from typing import Iterable, Union
|
||||
import torch
|
||||
import torch as th
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
from einops import rearrange, repeat
|
||||
|
||||
from comfy.ldm.modules.diffusionmodules.util import (
|
||||
zero_module,
|
||||
timestep_embedding,
|
||||
)
|
||||
|
||||
from comfy.cldm.cldm import ControlNet as ControlNetCLDM
|
||||
from comfy.ldm.modules.attention import SpatialTransformer
|
||||
from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample
|
||||
from comfy.ldm.util import exists
|
||||
from comfy.ldm.modules.attention import default, optimized_attention
|
||||
from comfy.ldm.modules.attention import FeedForward, SpatialTransformer
|
||||
from comfy.controlnet import broadcast_image_to
|
||||
from comfy.utils import repeat_to_batch_size
|
||||
import comfy.ops
|
||||
|
||||
# from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class SparseMethodImport(ABC):
|
||||
SPREAD = "spread"
|
||||
INDEX = "index"
|
||||
def __init__(self, method: str):
|
||||
self.method = method
|
||||
|
||||
@abstractmethod
|
||||
def get_indexes(self, hint_length: int, full_length: int) -> list[int]:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
class SparseIndexMethodImport(SparseMethodImport):
|
||||
def __init__(self, idxs: list[int]):
|
||||
super().__init__(self.INDEX)
|
||||
self.idxs = idxs
|
||||
|
||||
def get_indexes(self, hint_length: int, full_length: int) -> list[int]:
|
||||
orig_hint_length = hint_length
|
||||
if hint_length > full_length:
|
||||
hint_length = full_length
|
||||
# if idxs is less than hint_length, throw error
|
||||
if len(self.idxs) < hint_length:
|
||||
err_msg = f"There are not enough indexes ({len(self.idxs)}) provided to fit the usable {hint_length} input images."
|
||||
if orig_hint_length != hint_length:
|
||||
err_msg = f"{err_msg} (original input images: {orig_hint_length})"
|
||||
raise ValueError(err_msg)
|
||||
# cap idxs to hint_length
|
||||
idxs = self.idxs[:hint_length]
|
||||
new_idxs = []
|
||||
real_idxs = set()
|
||||
for idx in idxs:
|
||||
if idx < 0:
|
||||
real_idx = full_length+idx
|
||||
if real_idx in real_idxs:
|
||||
raise ValueError(f"Index '{idx}' maps to '{real_idx}' and is duplicate - indexes in Sparse Index Method must be unique.")
|
||||
else:
|
||||
real_idx = idx
|
||||
if real_idx in real_idxs:
|
||||
raise ValueError(f"Index '{idx}' is duplicate (or a negative index is equivalent) - indexes in Sparse Index Method must be unique.")
|
||||
real_idxs.add(real_idx)
|
||||
new_idxs.append(real_idx)
|
||||
return new_idxs
|
||||
|
||||
@@ -1,103 +0,0 @@
|
||||
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,)
|
||||
@@ -1,244 +0,0 @@
|
||||
from typing import Union
|
||||
|
||||
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,
|
||||
weights: int,
|
||||
frame_numbers: float):
|
||||
|
||||
|
||||
curr_latent_keyframe = LatentKeyframeGroupImport()
|
||||
|
||||
for i, frame_number in enumerate(frame_numbers):
|
||||
keyframe = LatentKeyframeImport(frame_number, float(weights[i]))
|
||||
curr_latent_keyframe.add(keyframe)
|
||||
|
||||
return (curr_latent_keyframe,)
|
||||
|
||||
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,)
|
||||
@@ -1,36 +0,0 @@
|
||||
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)
|
||||
@@ -1,194 +0,0 @@
|
||||
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])
|
||||
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
from torch import Tensor
|
||||
|
||||
import folder_paths
|
||||
from nodes import VAEEncode
|
||||
import comfy.utils
|
||||
|
||||
# from .utils import TimestepKeyframeGroup
|
||||
from .control_sparsectrl import SparseIndexMethodImport
|
||||
# from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced
|
||||
|
||||
|
||||
|
||||
class SparseIndexMethodNodeImport:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"indexes": ("STRING", {"default": "0"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SPARSE_METHOD",)
|
||||
FUNCTION = "get_method"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def get_method(self, indexes: str):
|
||||
idxs = []
|
||||
unique_idxs = set()
|
||||
# get indeces from string
|
||||
str_idxs = [x.strip() for x in indexes.strip().split(",")]
|
||||
for str_idx in str_idxs:
|
||||
try:
|
||||
idx = int(str_idx)
|
||||
if idx in unique_idxs:
|
||||
raise ValueError(f"'{idx}' is duplicated; indexes must be unique.")
|
||||
idxs.append(idx)
|
||||
unique_idxs.add(idx)
|
||||
except ValueError:
|
||||
raise ValueError(f"'{str_idx}' is not a valid integer index.")
|
||||
if len(idxs) == 0:
|
||||
raise ValueError(f"No indexes were listed in Sparse Index Method.")
|
||||
return (SparseIndexMethodImport(idxs),)
|
||||
|
||||
@@ -1,12 +0,0 @@
|
||||
class AnimateDiffLoaderWithContext:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = ""
|
||||
@@ -1,201 +0,0 @@
|
||||
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)))
|
||||
Reference in New Issue
Block a user