Merge PR #42 from Kosinkadink/develop - SparseCtrl support
Added SparseCtrl support - RGB + scribble
This commit is contained in:
+349
-578
@@ -1,561 +1,18 @@
|
||||
from typing import Union
|
||||
from typing import Callable, Union
|
||||
from torch import Tensor
|
||||
import torch
|
||||
import os
|
||||
|
||||
import comfy.utils
|
||||
import comfy.model_management
|
||||
import comfy.model_detection
|
||||
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 ControlWeightType:
|
||||
DEFAULT = "default"
|
||||
UNIVERSAL = "universal"
|
||||
T2IADAPTER = "t2iadapter"
|
||||
CONTROLNET = "controlnet"
|
||||
CONTROLLORA = "controllora"
|
||||
CONTROLLLLITE = "controllllite"
|
||||
|
||||
|
||||
class ControlWeights:
|
||||
def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None):
|
||||
self.weight_type = weight_type
|
||||
self.base_multiplier = base_multiplier
|
||||
self.flip_weights = flip_weights
|
||||
self.weights = weights
|
||||
if self.weights is not None and self.flip_weights:
|
||||
self.weights.reverse()
|
||||
self.weight_mask = weight_mask
|
||||
|
||||
def get(self, idx: int) -> Union[float, Tensor]:
|
||||
# if weights is not none, return index
|
||||
if self.weights is not None:
|
||||
return self.weights[idx]
|
||||
return 1.0
|
||||
|
||||
@classmethod
|
||||
def default(cls):
|
||||
return cls(ControlWeightType.DEFAULT)
|
||||
|
||||
@classmethod
|
||||
def universal(cls, base_multiplier: float, flip_weights: bool=False):
|
||||
return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def universal_mask(cls, weight_mask: Tensor):
|
||||
return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask)
|
||||
|
||||
@classmethod
|
||||
def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*12
|
||||
return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controlnet(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*13
|
||||
return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controllora(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*10
|
||||
return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controllllite(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
# TODO: make this have a real value
|
||||
weights = [1.0]*200
|
||||
return cls(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
|
||||
class StrengthInterpolation:
|
||||
LINEAR = "linear"
|
||||
EASE_IN = "ease-in"
|
||||
EASE_OUT = "ease-out"
|
||||
EASE_IN_OUT = "ease-in-out"
|
||||
NONE = "none"
|
||||
|
||||
|
||||
class LatentKeyframe:
|
||||
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 LatentKeyframeGroup:
|
||||
def __init__(self) -> None:
|
||||
self.keyframes: list[LatentKeyframe] = []
|
||||
|
||||
def add(self, keyframe: LatentKeyframe) -> 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[LatentKeyframe, None]:
|
||||
try:
|
||||
return self.keyframes[index]
|
||||
except IndexError:
|
||||
return None
|
||||
|
||||
def __getitem__(self, index) -> LatentKeyframe:
|
||||
return self.keyframes[index]
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self) -> 'LatentKeyframeGroup':
|
||||
cloned = LatentKeyframeGroup()
|
||||
for tk in self.keyframes:
|
||||
cloned.add(tk)
|
||||
return cloned
|
||||
|
||||
|
||||
class TimestepKeyframe:
|
||||
def __init__(self,
|
||||
start_percent: float = 0.0,
|
||||
strength: float = 1.0,
|
||||
interpolation: str = StrengthInterpolation.NONE,
|
||||
control_weights: ControlWeights = None,
|
||||
latent_keyframes: LatentKeyframeGroup = 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) -> 'TimestepKeyframe':
|
||||
return cls(0.0)
|
||||
|
||||
|
||||
# always maintain sorted state (by start_percent of TimestepKeyFrame)
|
||||
class TimestepKeyframeGroup:
|
||||
def __init__(self) -> None:
|
||||
self.keyframes: list[TimestepKeyframe] = []
|
||||
self.keyframes.append(TimestepKeyframe.default())
|
||||
|
||||
def add(self, keyframe: TimestepKeyframe) -> 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[TimestepKeyframe, 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) -> TimestepKeyframe:
|
||||
return self.keyframes[index]
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.keyframes)
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self) -> 'TimestepKeyframeGroup':
|
||||
cloned = TimestepKeyframeGroup()
|
||||
for tk in self.keyframes:
|
||||
cloned.add(tk)
|
||||
return cloned
|
||||
|
||||
@classmethod
|
||||
def default(cls, keyframe: TimestepKeyframe) -> 'TimestepKeyframeGroup':
|
||||
group = cls()
|
||||
group.keyframes[0] = keyframe
|
||||
return group
|
||||
|
||||
|
||||
# used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function
|
||||
|
||||
|
||||
class AdvancedControlBase:
|
||||
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights):
|
||||
self.base = base
|
||||
self.compatible_weights = [ControlWeightType.UNIVERSAL]
|
||||
self.add_compatible_weight(weights_default.weight_type)
|
||||
# mask for which parts of controlnet output to keep
|
||||
self.mask_cond_hint_original = None
|
||||
self.mask_cond_hint = None
|
||||
self.tk_mask_cond_hint_original = None
|
||||
self.tk_mask_cond_hint = None
|
||||
self.weight_mask_cond_hint = None
|
||||
# actual index values
|
||||
self.sub_idxs = None
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
# timesteps
|
||||
self.t: Tensor = None
|
||||
self.batched_number: int = None
|
||||
# weights + override
|
||||
self.weights: ControlWeights = None
|
||||
self.weights_default: ControlWeights = weights_default
|
||||
self.weights_override: ControlWeights = None
|
||||
# latent keyframe + override
|
||||
self.latent_keyframes: LatentKeyframeGroup = None
|
||||
self.latent_keyframe_override: LatentKeyframeGroup = None
|
||||
# initialize timestep_keyframes
|
||||
self.set_timestep_keyframes(timestep_keyframes)
|
||||
# override some functions
|
||||
self.get_control = self.get_control_inject
|
||||
self.control_merge = self.control_merge_inject#.__get__(self, type(self))
|
||||
self.pre_run = self.pre_run_inject
|
||||
self.cleanup = self.cleanup_inject
|
||||
|
||||
def add_compatible_weight(self, control_weight_type: str):
|
||||
self.compatible_weights.append(control_weight_type)
|
||||
|
||||
def verify_all_weights(self, throw_error=True):
|
||||
# first, check if override exists - if so, only need to check the override
|
||||
if self.weights_override is not None:
|
||||
if self.weights_override.weight_type not in self.compatible_weights:
|
||||
msg = f"Weight override is type {self.weights_override.weight_type}, but loaded {type(self).__name__}" + \
|
||||
f"only supports {self.compatible_weights} weights."
|
||||
raise WeightTypeException(msg)
|
||||
# otherwise, check all timestep keyframe weights
|
||||
else:
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
if tk.has_control_weights() and tk.control_weights.weight_type not in self.compatible_weights:
|
||||
msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type" + \
|
||||
f"{tk.control_weights.weight_type}, but loaded {type(self).__name__} only supports {self.compatible_weights} weights."
|
||||
raise WeightTypeException(msg)
|
||||
|
||||
def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroup):
|
||||
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
|
||||
# prepare first timestep_keyframe related stuff
|
||||
self.current_timestep_keyframe = None
|
||||
self.current_timestep_index = -1
|
||||
self.next_timestep_keyframe = None
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
|
||||
def prepare_current_timestep(self, t: Tensor, batched_number: int):
|
||||
self.t = t
|
||||
self.batched_number = batched_number
|
||||
# get current step percent
|
||||
curr_t: float = t[0]
|
||||
prev_index = self.current_timestep_index
|
||||
# if has next index, loop through and see if need to switch
|
||||
if self.timestep_keyframes.has_index(self.current_timestep_index+1):
|
||||
for i in range(self.current_timestep_index+1, len(self.timestep_keyframes)):
|
||||
eval_tk = self.timestep_keyframes[i]
|
||||
# check if start percent is less or equal to curr_t
|
||||
if eval_tk.start_t >= curr_t:
|
||||
self.current_timestep_index = i
|
||||
self.current_timestep_keyframe = eval_tk
|
||||
# keep track of control weights, latent keyframes, and masks,
|
||||
# accounting for inherit_missing
|
||||
if self.current_timestep_keyframe.has_control_weights():
|
||||
self.weights = self.current_timestep_keyframe.control_weights
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
self.weights = self.weights_default
|
||||
if self.current_timestep_keyframe.has_latent_keyframes():
|
||||
self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
self.latent_keyframes = None
|
||||
if self.current_timestep_keyframe.has_mask_hint():
|
||||
self.tk_mask_cond_hint_original = self.current_timestep_keyframe.mask_hint_orig
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
del self.tk_mask_cond_hint_original
|
||||
self.tk_mask_cond_hint_original = None
|
||||
# if guarantee_usage, stop searching for other TKs
|
||||
if self.current_timestep_keyframe.guarantee_usage:
|
||||
break
|
||||
# if eval_tk is outside of percent range, stop looking further
|
||||
else:
|
||||
break
|
||||
|
||||
# if index changed, apply overrides
|
||||
if prev_index != self.current_timestep_index:
|
||||
if self.weights_override is not None:
|
||||
self.weights = self.weights_override
|
||||
if self.latent_keyframe_override is not None:
|
||||
self.latent_keyframes = self.latent_keyframe_override
|
||||
|
||||
# make sure weights and latent_keyframes are in a workable state
|
||||
# Note: each AdvancedControlBase should create their own get_universal_weights class
|
||||
self.prepare_weights()
|
||||
|
||||
def prepare_weights(self):
|
||||
if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT:
|
||||
self.weights = self.weights_default
|
||||
elif self.weights.weight_type == ControlWeightType.UNIVERSAL:
|
||||
# if universal and weight_mask present, no need to convert
|
||||
if self.weights.weight_mask is not None:
|
||||
return
|
||||
self.weights = self.get_universal_weights()
|
||||
|
||||
def get_universal_weights(self) -> ControlWeights:
|
||||
return self.weights
|
||||
|
||||
def set_cond_hint_mask(self, mask_hint):
|
||||
self.mask_cond_hint_original = mask_hint
|
||||
return self
|
||||
|
||||
def pre_run_inject(self, model, percent_to_timestep_function):
|
||||
self.base.pre_run(model, percent_to_timestep_function)
|
||||
self.pre_run_advanced(model, percent_to_timestep_function)
|
||||
|
||||
def pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
# for each timestep keyframe, calculate the start_t
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
tk.start_t = percent_to_timestep_function(tk.start_percent)
|
||||
# clear variables
|
||||
self.cleanup_advanced()
|
||||
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number):
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
# if should not perform any actions for the controlnet, exit without doing any work
|
||||
if self.strength == 0.0 or self.current_timestep_keyframe.strength == 0.0:
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
||||
if control_prev is not None:
|
||||
return control_prev
|
||||
else:
|
||||
return None
|
||||
# otherwise, perform normal function
|
||||
return self.get_control_advanced(x_noisy, t, cond, batched_number)
|
||||
|
||||
def get_control_advanced(self, x_noisy, t, cond, batched_number):
|
||||
pass
|
||||
|
||||
def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]:
|
||||
if self.weights.weight_mask is not None:
|
||||
# prepare weight mask
|
||||
self.prepare_weight_mask_cond_hint(x, self.batched_number)
|
||||
# adjust mask for current layer and return
|
||||
return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, layers=layers))
|
||||
return self.weights.get(idx=idx)
|
||||
|
||||
def get_calc_pow(self, idx: int, layers: int) -> int:
|
||||
return (layers-1)-idx
|
||||
|
||||
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int):
|
||||
# apply strengths, and get batch indeces to null out
|
||||
# AKA latents that should not be influenced by ControlNet
|
||||
if self.latent_keyframes is not None:
|
||||
latent_count = x.size(0)//batched_number
|
||||
indeces_to_null = set(range(latent_count))
|
||||
mapped_indeces = None
|
||||
# if expecting subdivision, will need to translate between subset and actual idx values
|
||||
if self.sub_idxs:
|
||||
mapped_indeces = {}
|
||||
for i, actual in enumerate(self.sub_idxs):
|
||||
mapped_indeces[actual] = i
|
||||
for keyframe in self.latent_keyframes:
|
||||
real_index = keyframe.batch_index
|
||||
# if negative, count from end
|
||||
if real_index < 0:
|
||||
real_index += latent_count if self.sub_idxs is None else self.full_latent_length
|
||||
|
||||
# if not mapping indeces, what you see is what you get
|
||||
if mapped_indeces is None:
|
||||
if real_index in indeces_to_null:
|
||||
indeces_to_null.remove(real_index)
|
||||
# otherwise, see if batch_index is even included in this set of latents
|
||||
else:
|
||||
real_index = mapped_indeces.get(real_index, None)
|
||||
if real_index is None:
|
||||
continue
|
||||
indeces_to_null.remove(real_index)
|
||||
|
||||
# if real_index is outside the bounds of latents, don't apply
|
||||
if real_index >= latent_count or real_index < 0:
|
||||
continue
|
||||
|
||||
# apply strength for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength
|
||||
|
||||
# null them out by multiplying by null_latent_kf_strength
|
||||
for batch_index in indeces_to_null:
|
||||
# apply null for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength
|
||||
# apply masks, resizing mask to required dims
|
||||
if self.mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
if self.tk_mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
# apply timestep keyframe strengths
|
||||
if self.current_timestep_keyframe.strength != 1.0:
|
||||
x[:] *= self.current_timestep_keyframe.strength
|
||||
|
||||
def control_merge_inject(self: 'AdvancedControlBase', control_input, control_output, control_prev, output_dtype):
|
||||
out = {'input':[], 'middle':[], 'output': []}
|
||||
|
||||
if control_input is not None:
|
||||
for i in range(len(control_input)):
|
||||
key = 'input'
|
||||
x = control_input[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_input))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
out[key].insert(0, x)
|
||||
|
||||
if control_output is not None:
|
||||
for i in range(len(control_output)):
|
||||
if i == (len(control_output) - 1):
|
||||
key = 'middle'
|
||||
index = 0
|
||||
else:
|
||||
key = 'output'
|
||||
index = i
|
||||
x = control_output[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
if self.global_average_pooling:
|
||||
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
|
||||
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_output))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
|
||||
out[key].append(x)
|
||||
if control_prev is not None:
|
||||
for x in ['input', 'middle', 'output']:
|
||||
o = out[x]
|
||||
for i in range(len(control_prev[x])):
|
||||
prev_val = control_prev[x][i]
|
||||
if i >= len(o):
|
||||
o.append(prev_val)
|
||||
elif prev_val is not None:
|
||||
if o[i] is None:
|
||||
o[i] = prev_val
|
||||
else:
|
||||
o[i] += prev_val
|
||||
return out
|
||||
|
||||
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype)
|
||||
self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype)
|
||||
|
||||
def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype)
|
||||
|
||||
def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None):
|
||||
return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True)
|
||||
|
||||
def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
|
||||
# make mask appropriate dimensions, if present
|
||||
if orig_mask is not None:
|
||||
out_mask = getattr(self, attr_name)
|
||||
if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]:
|
||||
self._reset_attr(attr_name)
|
||||
del out_mask
|
||||
# TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM
|
||||
# resize mask and match batch count
|
||||
multiplier = 1 if direct_attn else 8
|
||||
out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier)
|
||||
actual_latent_length = x_noisy.shape[0] // batched_number
|
||||
out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
|
||||
if self.sub_idxs is not None:
|
||||
out_mask = out_mask[self.sub_idxs]
|
||||
# make cond_hint_mask length match x_noise
|
||||
if x_noisy.shape[0] != out_mask.shape[0]:
|
||||
out_mask = broadcast_image_to(out_mask, x_noisy.shape[0], batched_number)
|
||||
# default dtype to be same as x_noisy
|
||||
if dtype is None:
|
||||
dtype = x_noisy.dtype
|
||||
setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device))
|
||||
del out_mask
|
||||
|
||||
def _reset_attr(self, attr_name, new_value=None):
|
||||
if hasattr(self, attr_name):
|
||||
delattr(self, attr_name)
|
||||
setattr(self, attr_name, new_value)
|
||||
|
||||
def cleanup_inject(self):
|
||||
self.base.cleanup()
|
||||
self.cleanup_advanced()
|
||||
|
||||
def cleanup_advanced(self):
|
||||
self.sub_idxs = None
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
self.t = None
|
||||
self.batched_number = None
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
# timestep stuff
|
||||
self.current_timestep_keyframe = None
|
||||
self.next_timestep_keyframe = None
|
||||
self.current_timestep_index = -1
|
||||
# clear mask hints
|
||||
if self.mask_cond_hint is not None:
|
||||
del self.mask_cond_hint
|
||||
self.mask_cond_hint = None
|
||||
if self.tk_mask_cond_hint_original is not None:
|
||||
del self.tk_mask_cond_hint_original
|
||||
self.tk_mask_cond_hint_original = None
|
||||
if self.tk_mask_cond_hint is not None:
|
||||
del self.tk_mask_cond_hint
|
||||
self.tk_mask_cond_hint = None
|
||||
if self.weight_mask_cond_hint is not None:
|
||||
del self.weight_mask_cond_hint
|
||||
self.weight_mask_cond_hint = None
|
||||
|
||||
def copy_to_advanced(self, copied: 'AdvancedControlBase'):
|
||||
copied.mask_cond_hint_original = self.mask_cond_hint_original
|
||||
copied.weights_override = self.weights_override
|
||||
copied.latent_keyframe_override = self.latent_keyframe_override
|
||||
from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper
|
||||
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException,
|
||||
manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory)
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class ControlNetAdvanced(ControlNet, AdvancedControlBase):
|
||||
@@ -711,20 +168,210 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase):
|
||||
global_average_pooling=v.global_average_pooling, device=v.device)
|
||||
|
||||
|
||||
class ControlLLLiteAdvanced(ControlNet, AdvancedControlBase):
|
||||
def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, device=None):
|
||||
class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
|
||||
# This ControlNet is more of an attention patch than a traditional controlnet
|
||||
# So, the pre_run will be responsible for a lot of the functionality,
|
||||
# while the usual get_control is mostly used to set some values
|
||||
def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None):
|
||||
super().__init__(device)
|
||||
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite())
|
||||
self.already_patched = False
|
||||
|
||||
def set_cond_hint(self, *args, **kwargs):
|
||||
super().set_cond_hint(*args, **kwargs)
|
||||
# cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1)
|
||||
self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0
|
||||
|
||||
def pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function)
|
||||
logger.info(f"In ControlLLLiteAdvanced pre_run_advanced! {self.already_patched}")
|
||||
# perform patches if not already patches
|
||||
if not self.already_patched:
|
||||
self.already_patched = True
|
||||
|
||||
def get_control(self, x_noisy: Tensor, t, cond, batched_number):
|
||||
logger.info("In ControlLLLiteAdvanced get_control!")
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
# perform other controlnets
|
||||
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
|
||||
|
||||
def get_models(self):
|
||||
logger.info(f"In ControlLLLiteAdvanced get_models!")
|
||||
# get_models is called once at the start of every KSampler run - use to reset already_patched status
|
||||
self.already_patched = False
|
||||
out = super().get_models()
|
||||
return out
|
||||
|
||||
def copy(self):
|
||||
c = ControlLLLiteAdvanced(self.timestep_keyframes)
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
def cleanup(self):
|
||||
super().cleanup()
|
||||
self.cleanup_advanced()
|
||||
self.already_patched = False
|
||||
|
||||
|
||||
class SparseCtrlAdvanced(ControlNetAdvanced):
|
||||
def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None):
|
||||
super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
self.add_compatible_weight(ControlWeightType.SPARSECTRL)
|
||||
self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints
|
||||
self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default()
|
||||
self.latent_format = None
|
||||
self.preprocessed = False
|
||||
|
||||
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
|
||||
# normal ControlNet stuff
|
||||
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
|
||||
# set actual input length on motion model
|
||||
actual_length = x_noisy.size(0)//batched_number
|
||||
full_length = actual_length if self.sub_idxs is None else self.full_latent_length
|
||||
self.control_model.set_actual_length(actual_length=actual_length, full_length=full_length)
|
||||
# prepare cond_hint, if needed
|
||||
dim_mult = 1 if self.control_model.use_simplified_conditioning_embedding else 8
|
||||
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2]*dim_mult != self.cond_hint.shape[2] or x_noisy.shape[3]*dim_mult != self.cond_hint.shape[3]:
|
||||
# clear out cond_hint and conditioning_mask
|
||||
if self.cond_hint is not None:
|
||||
del self.cond_hint
|
||||
self.cond_hint = None
|
||||
# first, figure out which cond idxs are relevant, and where they fit in
|
||||
cond_idxs = self.sparse_settings.sparse_method.get_indexes(hint_length=self.cond_hint_original.size(0), full_length=full_length)
|
||||
|
||||
range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs
|
||||
hint_idxs = [] # idxs in cond_idxs
|
||||
local_idxs = [] # idx to pun in final cond_hint
|
||||
for i,cond_idx in enumerate(cond_idxs):
|
||||
if cond_idx in range_idxs:
|
||||
hint_idxs.append(i)
|
||||
local_idxs.append(range_idxs.index(cond_idx))
|
||||
# sub_cond_hint now contains the hints relevant to current x_noisy
|
||||
sub_cond_hint = self.cond_hint_original[hint_idxs].to(dtype).to(self.device)
|
||||
|
||||
# scale cond_hints to match noisy input
|
||||
if self.control_model.use_simplified_conditioning_embedding:
|
||||
# RGB SparseCtrl; the inputs are latents - use bilinear to avoid blocky artifacts
|
||||
sub_cond_hint = self.latent_format.process_in(sub_cond_hint) # multiplies by model scale factor
|
||||
sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "nearest-exact", "center").to(dtype).to(self.device)
|
||||
else:
|
||||
# other SparseCtrl; inputs are typical images
|
||||
sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device)
|
||||
# prepare cond_hint (b, c, h ,w)
|
||||
cond_shape = list(sub_cond_hint.shape)
|
||||
cond_shape[0] = len(range_idxs)
|
||||
self.cond_hint = torch.zeros(cond_shape).to(dtype).to(self.device)
|
||||
self.cond_hint[local_idxs] = sub_cond_hint[:]
|
||||
# prepare cond_mask (b, 1, h, w)
|
||||
cond_shape[1] = 1
|
||||
cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device)
|
||||
cond_mask[local_idxs] = 1.0
|
||||
# combine cond_hint and cond_mask into (b, c+1, h, w)
|
||||
if not self.sparse_settings.merged:
|
||||
self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1)
|
||||
del sub_cond_hint
|
||||
del cond_mask
|
||||
# make cond_hint match x_noisy batch
|
||||
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']
|
||||
y = cond.get('y', 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 pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
super().pre_run_advanced(model, percent_to_timestep_function)
|
||||
if type(self.cond_hint_original) == PreprocSparseRGBWrapper:
|
||||
if not self.control_model.use_simplified_conditioning_embedding:
|
||||
raise ValueError("Any model besides RGB SparseCtrl should NOT have its images go through the RGB SparseCtrl preprocessor.")
|
||||
self.cond_hint_original = self.cond_hint_original.condhint
|
||||
self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint
|
||||
if self.control_model.motion_holder is not None:
|
||||
self.control_model.motion_holder.motion_wrapper.reset()
|
||||
self.control_model.motion_holder.motion_wrapper.set_strength(self.sparse_settings.motion_strength)
|
||||
self.control_model.motion_holder.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale)
|
||||
|
||||
def cleanup_advanced(self):
|
||||
super().cleanup_advanced()
|
||||
if self.latent_format is not None:
|
||||
del self.latent_format
|
||||
self.latent_format = None
|
||||
|
||||
def copy(self):
|
||||
c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype)
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
|
||||
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
|
||||
control = comfy_cn.load_controlnet(ckpt_path, model=model)
|
||||
# TODO: support controlnet-lllite
|
||||
# if is None, see if is a non-vanilla ControlNet
|
||||
# if control is None:
|
||||
# controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
# # check if lllite
|
||||
# if "lllite_unet" in controlnet_data:
|
||||
# pass
|
||||
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
control = None
|
||||
# check if a non-vanilla ControlNet
|
||||
controlnet_type = ControlWeightType.DEFAULT
|
||||
has_controlnet_key = False
|
||||
has_motion_modules_key = False
|
||||
for key in controlnet_data:
|
||||
# LLLLite check
|
||||
if "lllite" in key:
|
||||
logger.info("ControlLLLite controlnet!")
|
||||
controlnet_type = ControlWeightType.CONTROLLLLITE
|
||||
break
|
||||
# SparseCtrl check
|
||||
elif "motion_modules" in key:
|
||||
has_motion_modules_key = True
|
||||
elif "controlnet" in key:
|
||||
has_controlnet_key = True
|
||||
if has_controlnet_key and has_motion_modules_key:
|
||||
controlnet_type = ControlWeightType.SPARSECTRL
|
||||
|
||||
if controlnet_type != ControlWeightType.DEFAULT:
|
||||
if controlnet_type == ControlWeightType.CONTROLLLLITE:
|
||||
raise NotImplementedError("ControlLLLite has not been fully implemented yet!")
|
||||
control = ControlLLLiteAdvanced(timestep_keyframes=timestep_keyframe)
|
||||
# load Controll
|
||||
elif controlnet_type == ControlWeightType.SPARSECTRL:
|
||||
control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model)
|
||||
# otherwise, load vanilla ControlNet
|
||||
else:
|
||||
try:
|
||||
# hacky way of getting load_torch_file in load_controlnet to use already-present controlnet_data and not redo loading
|
||||
orig_load_torch_file = comfy.utils.load_torch_file
|
||||
comfy.utils.load_torch_file = load_torch_file_with_dict_factory(controlnet_data, orig_load_torch_file)
|
||||
control = comfy_cn.load_controlnet(ckpt_path, model=model)
|
||||
finally:
|
||||
comfy.utils.load_torch_file = orig_load_torch_file
|
||||
return convert_to_advanced(control, timestep_keyframe=timestep_keyframe)
|
||||
|
||||
|
||||
@@ -749,25 +396,149 @@ 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
|
||||
def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=SparseSettings.default(), model=None) -> SparseCtrlAdvanced:
|
||||
if controlnet_data is None:
|
||||
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
# first, separate out motion part from normal controlnet part and attempt to load that portion
|
||||
motion_data = {}
|
||||
for key in list(controlnet_data.keys()):
|
||||
if "temporal" in key:
|
||||
motion_data[key] = controlnet_data.pop(key)
|
||||
if len(motion_data) == 0:
|
||||
raise ValueError(f"No motion-related keys in '{ckpt_path}'; not a valid SparseCtrl model!")
|
||||
motion_wrapper: SparseCtrlMotionWrapper = SparseCtrlMotionWrapper(motion_data).to(comfy.model_management.unet_dtype())
|
||||
missing, unexpected = motion_wrapper.load_state_dict(motion_data)
|
||||
if len(missing) > 0 or len(unexpected) > 0:
|
||||
logger.info(f"SparseCtrlMotionWrapper: {missing}, {unexpected}")
|
||||
|
||||
# now, load as if it was a normal controlnet - mostly copied from comfy load_controlnet function
|
||||
controlnet_config = None
|
||||
is_diffusers = False
|
||||
use_simplified_conditioning_embedding = False
|
||||
if "controlnet_cond_embedding.conv_in.weight" in controlnet_data:
|
||||
is_diffusers = True
|
||||
if "controlnet_cond_embedding.weight" in controlnet_data:
|
||||
is_diffusers = True
|
||||
use_simplified_conditioning_embedding = True
|
||||
if is_diffusers: #diffusers format
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
controlnet_config = comfy.model_detection.unet_config_from_diffusers_unet(controlnet_data, unet_dtype)
|
||||
diffusers_keys = comfy.utils.unet_to_diffusers(controlnet_config)
|
||||
diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"
|
||||
diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"
|
||||
|
||||
# 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
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
k_in = "controlnet_down_blocks.{}{}".format(count, s)
|
||||
k_out = "zero_convs.{}.0{}".format(count, s)
|
||||
if k_in not in controlnet_data:
|
||||
loop = False
|
||||
break
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
# normal conditioning embedding
|
||||
if not use_simplified_conditioning_embedding:
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
if count == 0:
|
||||
k_in = "controlnet_cond_embedding.conv_in{}".format(s)
|
||||
else:
|
||||
k_in = "controlnet_cond_embedding.blocks.{}{}".format(count - 1, s)
|
||||
k_out = "input_hint_block.{}{}".format(count * 2, s)
|
||||
if k_in not in controlnet_data:
|
||||
k_in = "controlnet_cond_embedding.conv_out{}".format(s)
|
||||
loop = False
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
# simplified conditioning embedding
|
||||
else:
|
||||
count = 0
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
k_in = "controlnet_cond_embedding{}".format(s)
|
||||
k_out = "input_hint_block.{}{}".format(count, s)
|
||||
diffusers_keys[k_in] = k_out
|
||||
|
||||
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
|
||||
new_sd = {}
|
||||
for k in diffusers_keys:
|
||||
if k in controlnet_data:
|
||||
new_sd[diffusers_keys[k]] = controlnet_data.pop(k)
|
||||
|
||||
leftover_keys = controlnet_data.keys()
|
||||
if len(leftover_keys) > 0:
|
||||
logger.info("leftover keys:", leftover_keys)
|
||||
controlnet_data = new_sd
|
||||
|
||||
class WeightTypeException(TypeError):
|
||||
"Raised when weight not compatible with AdvancedControlBase object"
|
||||
pass
|
||||
pth_key = 'control_model.zero_convs.0.0.weight'
|
||||
pth = False
|
||||
key = 'zero_convs.0.0.weight'
|
||||
if pth_key in controlnet_data:
|
||||
pth = True
|
||||
key = pth_key
|
||||
prefix = "control_model."
|
||||
elif key in controlnet_data:
|
||||
prefix = ""
|
||||
else:
|
||||
raise ValueError("The provided model is not a valid SparseCtrl model! [ErrorCode: HORSERADISH]")
|
||||
|
||||
if controlnet_config is None:
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, unet_dtype, True).unet_config
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device)
|
||||
if manual_cast_dtype is not None:
|
||||
controlnet_config["operations"] = manual_cast_clean_groupnorm
|
||||
else:
|
||||
controlnet_config["operations"] = disable_weight_init_clean_groupnorm
|
||||
controlnet_config.pop("out_channels")
|
||||
# get proper hint channels
|
||||
if use_simplified_conditioning_embedding:
|
||||
controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
|
||||
controlnet_config["use_simplified_conditioning_embedding"] = use_simplified_conditioning_embedding
|
||||
else:
|
||||
controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
|
||||
controlnet_config["use_simplified_conditioning_embedding"] = use_simplified_conditioning_embedding
|
||||
control_model = SparseControlNet(**controlnet_config)
|
||||
|
||||
if pth:
|
||||
if 'difference' in controlnet_data:
|
||||
if model is not None:
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
model_sd = model.model_state_dict()
|
||||
for x in controlnet_data:
|
||||
c_m = "control_model."
|
||||
if x.startswith(c_m):
|
||||
sd_key = "diffusion_model.{}".format(x[len(c_m):])
|
||||
if sd_key in model_sd:
|
||||
cd = controlnet_data[x]
|
||||
cd += model_sd[sd_key].type(cd.dtype).to(cd.device)
|
||||
else:
|
||||
logger.warning("WARNING: Loaded a diff SparseCtrl without a model. It will very likely not work.")
|
||||
|
||||
class WeightsLoader(torch.nn.Module):
|
||||
pass
|
||||
w = WeightsLoader()
|
||||
w.control_model = control_model
|
||||
missing, unexpected = w.load_state_dict(controlnet_data, strict=False)
|
||||
else:
|
||||
missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False)
|
||||
if len(missing) > 0 or len(unexpected) > 0:
|
||||
logger.info(f"SparseCtrl ControlNet: {missing}, {unexpected}")
|
||||
|
||||
global_average_pooling = False
|
||||
filename = os.path.splitext(ckpt_path)[0]
|
||||
if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling
|
||||
global_average_pooling = True
|
||||
|
||||
# both motion portion and controlnet portions are loaded; bring them together if using motion model
|
||||
if sparse_settings.use_motion:
|
||||
motion_wrapper.inject(control_model)
|
||||
|
||||
control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, sparse_settings=sparse_settings, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
return control
|
||||
|
||||
@@ -1 +1,240 @@
|
||||
# adapted from https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI
|
||||
# basically, all the LLLite core code is from there, which I then combined with
|
||||
# Advanced-ControlNet features and QoL
|
||||
import math
|
||||
import torch
|
||||
import os
|
||||
|
||||
import comfy
|
||||
|
||||
|
||||
def extra_options_to_module_prefix(extra_options):
|
||||
# extra_options = {'transformer_index': 2, 'block_index': 8, 'original_shape': [2, 4, 128, 128], 'block': ('input', 7), 'n_heads': 20, 'dim_head': 64}
|
||||
|
||||
# block is: [('input', 4), ('input', 5), ('input', 7), ('input', 8), ('middle', 0),
|
||||
# ('output', 0), ('output', 1), ('output', 2), ('output', 3), ('output', 4), ('output', 5)]
|
||||
# transformer_index is: [0, 1, 2, 3, 4, 5, 6, 7, 8], for each block
|
||||
# block_index is: 0-1 or 0-9, depends on the block
|
||||
# input 7 and 8, middle has 10 blocks
|
||||
|
||||
# make module name from extra_options
|
||||
block = extra_options["block"]
|
||||
block_index = extra_options["block_index"]
|
||||
if block[0] == "input":
|
||||
module_pfx = f"lllite_unet_input_blocks_{block[1]}_1_transformer_blocks_{block_index}"
|
||||
elif block[0] == "middle":
|
||||
module_pfx = f"lllite_unet_middle_block_1_transformer_blocks_{block_index}"
|
||||
elif block[0] == "output":
|
||||
module_pfx = f"lllite_unet_output_blocks_{block[1]}_1_transformer_blocks_{block_index}"
|
||||
else:
|
||||
raise Exception("invalid block name")
|
||||
return module_pfx
|
||||
|
||||
|
||||
def load_control_net_lllite_patch(path, cond_image, multiplier, num_steps, start_percent, end_percent):
|
||||
# calculate start and end step
|
||||
start_step = math.floor(num_steps * start_percent * 0.01) if start_percent > 0 else 0
|
||||
end_step = math.floor(num_steps * end_percent * 0.01) if end_percent > 0 else num_steps
|
||||
|
||||
# load weights
|
||||
ctrl_sd = comfy.utils.load_torch_file(path, safe_load=True)
|
||||
|
||||
# split each weights for each module
|
||||
module_weights = {}
|
||||
for key, value in ctrl_sd.items():
|
||||
fragments = key.split(".")
|
||||
module_name = fragments[0]
|
||||
weight_name = ".".join(fragments[1:])
|
||||
|
||||
if module_name not in module_weights:
|
||||
module_weights[module_name] = {}
|
||||
module_weights[module_name][weight_name] = value
|
||||
|
||||
# load each module
|
||||
modules = {}
|
||||
for module_name, weights in module_weights.items():
|
||||
# kohya planned to do something about how these should be chosen, so I'm not touching this
|
||||
# since I am not familiar with the logic for this
|
||||
if "conditioning1.4.weight" in weights:
|
||||
depth = 3
|
||||
elif weights["conditioning1.2.weight"].shape[-1] == 4:
|
||||
depth = 2
|
||||
else:
|
||||
depth = 1
|
||||
|
||||
module = LLLiteModule(
|
||||
name=module_name,
|
||||
is_conv2d=weights["down.0.weight"].ndim == 4,
|
||||
in_dim=weights["down.0.weight"].shape[1],
|
||||
depth=depth,
|
||||
cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2,
|
||||
mlp_dim=weights["down.0.weight"].shape[0],
|
||||
multiplier=multiplier,
|
||||
num_steps=num_steps,
|
||||
start_step=start_step,
|
||||
end_step=end_step,
|
||||
)
|
||||
info = module.load_state_dict(weights)
|
||||
modules[module_name] = module
|
||||
if len(modules) == 1:
|
||||
module.is_first = True
|
||||
|
||||
print(f"loaded {path} successfully, {len(modules)} modules")
|
||||
|
||||
for module in modules.values():
|
||||
module.set_cond_image(cond_image)
|
||||
|
||||
class control_net_lllite_patch:
|
||||
def __init__(self, modules):
|
||||
self.modules = modules
|
||||
|
||||
def __call__(self, q, k, v, extra_options):
|
||||
module_pfx = extra_options_to_module_prefix(extra_options)
|
||||
|
||||
is_attn1 = q.shape[-1] == k.shape[-1] # self attention
|
||||
if is_attn1:
|
||||
module_pfx = module_pfx + "_attn1"
|
||||
else:
|
||||
module_pfx = module_pfx + "_attn2"
|
||||
|
||||
module_pfx_to_q = module_pfx + "_to_q"
|
||||
module_pfx_to_k = module_pfx + "_to_k"
|
||||
module_pfx_to_v = module_pfx + "_to_v"
|
||||
|
||||
if module_pfx_to_q in self.modules:
|
||||
q = q + self.modules[module_pfx_to_q](q)
|
||||
if module_pfx_to_k in self.modules:
|
||||
k = k + self.modules[module_pfx_to_k](k)
|
||||
if module_pfx_to_v in self.modules:
|
||||
v = v + self.modules[module_pfx_to_v](v)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def to(self, device):
|
||||
for d in self.modules.keys():
|
||||
self.modules[d] = self.modules[d].to(device)
|
||||
return self
|
||||
|
||||
return control_net_lllite_patch(modules)
|
||||
|
||||
|
||||
class LLLiteModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
is_conv2d: bool,
|
||||
in_dim: int,
|
||||
depth: int,
|
||||
cond_emb_dim: int,
|
||||
mlp_dim: int,
|
||||
multiplier: int,
|
||||
num_steps: int,
|
||||
start_step: int,
|
||||
end_step: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.name = name
|
||||
self.is_conv2d = is_conv2d
|
||||
self.multiplier = multiplier
|
||||
self.num_steps = num_steps
|
||||
self.start_step = start_step
|
||||
self.end_step = end_step
|
||||
self.is_first = False
|
||||
|
||||
modules = []
|
||||
modules.append(torch.nn.Conv2d(3, cond_emb_dim // 2, kernel_size=4, stride=4, padding=0)) # to latent (from VAE) size*2
|
||||
if depth == 1:
|
||||
modules.append(torch.nn.ReLU(inplace=True))
|
||||
modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=2, stride=2, padding=0))
|
||||
elif depth == 2:
|
||||
modules.append(torch.nn.ReLU(inplace=True))
|
||||
modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=4, stride=4, padding=0))
|
||||
elif depth == 3:
|
||||
# kernel size 8 is too large, so set it to 4
|
||||
modules.append(torch.nn.ReLU(inplace=True))
|
||||
modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim // 2, kernel_size=4, stride=4, padding=0))
|
||||
modules.append(torch.nn.ReLU(inplace=True))
|
||||
modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=2, stride=2, padding=0))
|
||||
|
||||
self.conditioning1 = torch.nn.Sequential(*modules)
|
||||
|
||||
if self.is_conv2d:
|
||||
self.down = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(in_dim, mlp_dim, kernel_size=1, stride=1, padding=0),
|
||||
torch.nn.ReLU(inplace=True),
|
||||
)
|
||||
self.mid = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(mlp_dim + cond_emb_dim, mlp_dim, kernel_size=1, stride=1, padding=0),
|
||||
torch.nn.ReLU(inplace=True),
|
||||
)
|
||||
self.up = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(mlp_dim, in_dim, kernel_size=1, stride=1, padding=0),
|
||||
)
|
||||
else:
|
||||
self.down = torch.nn.Sequential(
|
||||
torch.nn.Linear(in_dim, mlp_dim),
|
||||
torch.nn.ReLU(inplace=True),
|
||||
)
|
||||
self.mid = torch.nn.Sequential(
|
||||
torch.nn.Linear(mlp_dim + cond_emb_dim, mlp_dim),
|
||||
torch.nn.ReLU(inplace=True),
|
||||
)
|
||||
self.up = torch.nn.Sequential(
|
||||
torch.nn.Linear(mlp_dim, in_dim),
|
||||
)
|
||||
|
||||
self.depth = depth
|
||||
self.cond_image = None
|
||||
self.cond_emb = None
|
||||
self.current_step = 0
|
||||
|
||||
# @torch.inference_mode()
|
||||
def set_cond_image(self, cond_image):
|
||||
# print("set_cond_image", self.name)
|
||||
self.cond_image = cond_image
|
||||
self.cond_emb = None
|
||||
self.current_step = 0
|
||||
|
||||
def forward(self, x):
|
||||
if self.num_steps > 0:
|
||||
if self.current_step < self.start_step:
|
||||
self.current_step += 1
|
||||
return torch.zeros_like(x)
|
||||
elif self.current_step >= self.end_step:
|
||||
if self.is_first and self.current_step == self.end_step:
|
||||
print(f"end LLLite: step {self.current_step}")
|
||||
self.current_step += 1
|
||||
if self.current_step >= self.num_steps:
|
||||
self.current_step = 0 # reset
|
||||
return torch.zeros_like(x)
|
||||
else:
|
||||
if self.is_first and self.current_step == self.start_step:
|
||||
print(f"start LLLite: step {self.current_step}")
|
||||
self.current_step += 1
|
||||
if self.current_step >= self.num_steps:
|
||||
self.current_step = 0 # reset
|
||||
|
||||
if self.cond_emb is None:
|
||||
# print(f"cond_emb is None, {self.name}")
|
||||
cx = self.conditioning1(self.cond_image.to(x.device, dtype=x.dtype))
|
||||
if not self.is_conv2d:
|
||||
# reshape / b,c,h,w -> b,h*w,c
|
||||
n, c, h, w = cx.shape
|
||||
cx = cx.view(n, c, h * w).permute(0, 2, 1)
|
||||
self.cond_emb = cx
|
||||
|
||||
cx: torch.Tensor = self.cond_emb
|
||||
# print(f"forward {self.name}, {cx.shape}, {x.shape}")
|
||||
|
||||
# x in uncond/cond doubles batch size
|
||||
if x.shape[0] != cx.shape[0]:
|
||||
if self.is_conv2d:
|
||||
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1, 1)
|
||||
else:
|
||||
# print("x.shape[0] != cx.shape[0]", x.shape[0], cx.shape[0])
|
||||
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1)
|
||||
|
||||
cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2)
|
||||
cx = self.mid(cx)
|
||||
cx = self.up(cx)
|
||||
return cx * self.multiplier
|
||||
@@ -0,0 +1,892 @@
|
||||
#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 SparseControlNet(ControlNetCLDM):
|
||||
def __init__(self, *args,**kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
hint_channels = kwargs.get("hint_channels")
|
||||
operations: disable_weight_init_clean_groupnorm = kwargs.get("operations", disable_weight_init_clean_groupnorm)
|
||||
device = kwargs.get("device", None)
|
||||
self.use_simplified_conditioning_embedding = kwargs.get("use_simplified_conditioning_embedding", False)
|
||||
if self.use_simplified_conditioning_embedding:
|
||||
self.input_hint_block = TimestepEmbedSequential(
|
||||
zero_module(operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device)),
|
||||
#zero_module(operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device)),
|
||||
)
|
||||
self.motion_holder: MotionWrapperHolder = None
|
||||
|
||||
def set_actual_length(self, actual_length: int, full_length: int):
|
||||
if self.motion_holder is not None:
|
||||
self.motion_holder.motion_wrapper.set_video_length(video_length=actual_length, full_length=full_length)
|
||||
|
||||
def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs):
|
||||
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
# SparseCtrl sets noisy input to zeros
|
||||
x = torch.zeros_like(x)
|
||||
guided_hint = self.input_hint_block(hint, emb, context)
|
||||
|
||||
outs = []
|
||||
|
||||
hs = []
|
||||
if self.num_classes is not None:
|
||||
assert y.shape[0] == x.shape[0]
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
h = x
|
||||
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
|
||||
if guided_hint is not None:
|
||||
h = module(h, emb, context)
|
||||
h += guided_hint
|
||||
guided_hint = None
|
||||
else:
|
||||
h = module(h, emb, context)
|
||||
outs.append(zero_conv(h, emb, context))
|
||||
|
||||
h = self.middle_block(h, emb, context)
|
||||
outs.append(self.middle_block_out(h, emb, context))
|
||||
|
||||
return outs
|
||||
|
||||
|
||||
class PreprocSparseRGBWrapper:
|
||||
error_msg = "Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise). It cannot be used for anything else that accepts IMAGE input."
|
||||
def __init__(self, condhint: Tensor):
|
||||
self.condhint = condhint
|
||||
|
||||
def movedim(self, *args, **kwargs):
|
||||
return self
|
||||
|
||||
def __getattr__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
if name != "condhint":
|
||||
raise AttributeError(self.error_msg)
|
||||
super().__setattr__(name, value)
|
||||
|
||||
def __iter__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __next__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __len__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __getitem__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __setitem__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
|
||||
class SparseSettings:
|
||||
def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False):
|
||||
self.sparse_method = sparse_method
|
||||
self.use_motion = use_motion
|
||||
self.motion_strength = motion_strength
|
||||
self.motion_scale = motion_scale
|
||||
self.merged = merged
|
||||
|
||||
@classmethod
|
||||
def default(cls):
|
||||
return SparseSettings(sparse_method=SparseSpreadMethod(), use_motion=True)
|
||||
|
||||
|
||||
class SparseMethod(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 SparseSpreadMethod(SparseMethod):
|
||||
UNIFORM = "uniform"
|
||||
STARTING = "starting"
|
||||
ENDING = "ending"
|
||||
CENTER = "center"
|
||||
|
||||
LIST = [UNIFORM, STARTING, ENDING, CENTER]
|
||||
|
||||
def __init__(self, spread=UNIFORM):
|
||||
super().__init__(self.SPREAD)
|
||||
self.spread = spread
|
||||
|
||||
def get_indexes(self, hint_length: int, full_length: int) -> list[int]:
|
||||
# if hint_length >= full_length, limit hints to full_length
|
||||
if hint_length >= full_length:
|
||||
return list(range(full_length))
|
||||
# handle special case of 1 hint image
|
||||
if hint_length == 1:
|
||||
if self.spread in [self.UNIFORM, self.STARTING]:
|
||||
return [0]
|
||||
elif self.spread == self.ENDING:
|
||||
return [full_length-1]
|
||||
elif self.spread == self.CENTER:
|
||||
# return second (of three) values as the center
|
||||
return [np.linspace(0, full_length-1, 3, endpoint=True, dtype=int)[1]]
|
||||
else:
|
||||
raise ValueError(f"Unrecognized spread: {self.spread}")
|
||||
# otherwise, handle other cases
|
||||
if self.spread == self.UNIFORM:
|
||||
return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int))
|
||||
elif self.spread == self.STARTING:
|
||||
# make split 1 larger, remove last element
|
||||
return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1]
|
||||
elif self.spread == self.ENDING:
|
||||
# make split 1 larger, remove first element
|
||||
return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[1:]
|
||||
elif self.spread == self.CENTER:
|
||||
# if hint length is not 3 greater than full length, do STARTING behavior
|
||||
if full_length-hint_length < 3:
|
||||
return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1]
|
||||
# otherwise, get linspace of 2 greater than needed, then cut off first and last
|
||||
return list(np.linspace(0, full_length-1, hint_length+2, endpoint=True, dtype=int))[1:-1]
|
||||
return ValueError(f"Unrecognized spread: {self.spread}")
|
||||
|
||||
|
||||
class SparseIndexMethod(SparseMethod):
|
||||
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
|
||||
|
||||
|
||||
#########################################
|
||||
# motion-related portion of controlnet
|
||||
class BlockType:
|
||||
UP = "up"
|
||||
DOWN = "down"
|
||||
MID = "mid"
|
||||
|
||||
def get_down_block_max(mm_state_dict: dict[str, Tensor]) -> int:
|
||||
return get_block_max(mm_state_dict, "down_blocks")
|
||||
|
||||
def get_up_block_max(mm_state_dict: dict[str, Tensor]) -> int:
|
||||
return get_block_max(mm_state_dict, "up_blocks")
|
||||
|
||||
def get_block_max(mm_state_dict: dict[str, Tensor], block_name: str) -> int:
|
||||
# keep track of biggest down_block count in module
|
||||
biggest_block = -1
|
||||
for key in mm_state_dict.keys():
|
||||
if block_name in key:
|
||||
try:
|
||||
block_int = key.split(".")[1]
|
||||
block_num = int(block_int)
|
||||
if block_num > biggest_block:
|
||||
biggest_block = block_num
|
||||
except ValueError:
|
||||
pass
|
||||
return biggest_block
|
||||
|
||||
def has_mid_block(mm_state_dict: dict[str, Tensor]):
|
||||
# check if keys contain mid_block
|
||||
for key in mm_state_dict.keys():
|
||||
if key.startswith("mid_block."):
|
||||
return True
|
||||
return False
|
||||
|
||||
def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str=None) -> int:
|
||||
# use pos_encoder.pe entries to determine max length - [1, {max_length}, {320|640|1280}]
|
||||
for key in mm_state_dict.keys():
|
||||
if key.endswith("pos_encoder.pe"):
|
||||
return mm_state_dict[key].size(1) # get middle dim
|
||||
raise ValueError(f"No pos_encoder.pe found in SparseCtrl state_dict - {mm_name} is not a valid SparseCtrl model!")
|
||||
|
||||
|
||||
class MotionWrapperHolder:
|
||||
def __init__(self, motion_wrapper: 'SparseCtrlMotionWrapper'):
|
||||
self.motion_wrapper = motion_wrapper
|
||||
|
||||
|
||||
class SparseCtrlMotionWrapper(nn.Module):
|
||||
def __init__(self, mm_state_dict: dict[str, Tensor]):
|
||||
super().__init__()
|
||||
self.down_blocks: Iterable[MotionModule] = None
|
||||
self.up_blocks: Iterable[MotionModule] = None
|
||||
self.mid_block: MotionModule = None
|
||||
self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, "")
|
||||
layer_channels = (320, 640, 1280, 1280)
|
||||
if get_down_block_max(mm_state_dict) > -1:
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
for c in layer_channels:
|
||||
self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN))
|
||||
if get_up_block_max(mm_state_dict) > -1:
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
for c in reversed(layer_channels):
|
||||
self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP))
|
||||
if has_mid_block(mm_state_dict):
|
||||
self.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID)
|
||||
|
||||
def inject(self, unet: SparseControlNet):
|
||||
# inject input (down) blocks
|
||||
self._inject(unet.input_blocks, self.down_blocks)
|
||||
# inject mid block, if present
|
||||
if self.mid_block is not None:
|
||||
self._inject([unet.middle_block], [self.mid_block])
|
||||
unet.motion_holder = MotionWrapperHolder(self)
|
||||
|
||||
def _inject(self, unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList):
|
||||
# Rules for injection:
|
||||
# For each component list in a unet block:
|
||||
# if SpatialTransformer exists in list, place next block after last occurrence
|
||||
# elif ResBlock exists in list, place next block after first occurrence
|
||||
# else don't place block
|
||||
injection_count = 0
|
||||
unet_idx = 0
|
||||
# details about blocks passed in
|
||||
per_block = len(mm_blocks[0].motion_modules)
|
||||
injection_goal = len(mm_blocks) * per_block
|
||||
# only stop injecting when modules exhausted
|
||||
while injection_count < injection_goal:
|
||||
# figure out which VanillaTemporalModule from mm to inject
|
||||
mm_blk_idx, mm_vtm_idx = injection_count // per_block, injection_count % per_block
|
||||
# figure out layout of unet block components
|
||||
st_idx = -1 # SpatialTransformer index
|
||||
res_idx = -1 # first ResBlock index
|
||||
# first, figure out indeces of relevant blocks
|
||||
for idx, component in enumerate(unet_blocks[unet_idx]):
|
||||
if type(component) == SpatialTransformer:
|
||||
st_idx = idx
|
||||
elif type(component).__name__ == "ResBlock" and res_idx < 0:
|
||||
res_idx = idx
|
||||
# if SpatialTransformer exists, inject right after
|
||||
if st_idx >= 0:
|
||||
unet_blocks[unet_idx].insert(st_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx])
|
||||
injection_count += 1
|
||||
# otherwise, if only ResBlock exists, inject right after
|
||||
elif res_idx >= 0:
|
||||
unet_blocks[unet_idx].insert(res_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx])
|
||||
injection_count += 1
|
||||
# increment unet_idx
|
||||
unet_idx += 1
|
||||
|
||||
def eject(self, unet: SparseControlNet):
|
||||
# remove from input blocks (downblocks)
|
||||
self._eject(unet.input_blocks)
|
||||
# remove from middle block (encapsulate in list to make compatible)
|
||||
self._eject([unet.middle_block])
|
||||
del unet.motion_holder
|
||||
unet.motion_holder = None
|
||||
|
||||
def _eject(self, unet_blocks: nn.ModuleList):
|
||||
# eject all VanillaTemporalModule objects from all blocks
|
||||
for block in unet_blocks:
|
||||
idx_to_pop = []
|
||||
for idx, component in enumerate(block):
|
||||
if type(component) == VanillaTemporalModule:
|
||||
idx_to_pop.append(idx)
|
||||
# pop in backwards order, as to not disturb what the indeces refer to
|
||||
for idx in sorted(idx_to_pop, reverse=True):
|
||||
block.pop(idx)
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
self.AD_video_length = video_length
|
||||
if self.down_blocks is not None:
|
||||
for block in self.down_blocks:
|
||||
block.set_video_length(video_length, full_length)
|
||||
if self.up_blocks is not None:
|
||||
for block in self.up_blocks:
|
||||
block.set_video_length(video_length, full_length)
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.set_video_length(video_length, full_length)
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
if self.down_blocks is not None:
|
||||
for block in self.down_blocks:
|
||||
block.set_scale_multiplier(multiplier)
|
||||
if self.up_blocks is not None:
|
||||
for block in self.up_blocks:
|
||||
block.set_scale_multiplier(multiplier)
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_strength(self, strength: float):
|
||||
if self.down_blocks is not None:
|
||||
for block in self.down_blocks:
|
||||
block.set_strength(strength)
|
||||
if self.up_blocks is not None:
|
||||
for block in self.up_blocks:
|
||||
block.set_strength(strength)
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.set_strength(strength)
|
||||
|
||||
def reset_temp_vars(self):
|
||||
if self.down_blocks is not None:
|
||||
for block in self.down_blocks:
|
||||
block.reset_temp_vars()
|
||||
if self.up_blocks is not None:
|
||||
for block in self.up_blocks:
|
||||
block.reset_temp_vars()
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.reset_temp_vars()
|
||||
|
||||
def reset_scale_multiplier(self):
|
||||
self.set_scale_multiplier(None)
|
||||
|
||||
def reset(self):
|
||||
self.reset_scale_multiplier()
|
||||
self.reset_temp_vars()
|
||||
|
||||
|
||||
class MotionModule(nn.Module):
|
||||
def __init__(self, in_channels, temporal_position_encoding_max_len=24, block_type: str=BlockType.DOWN):
|
||||
super().__init__()
|
||||
if block_type == BlockType.MID:
|
||||
# mid blocks contain only a single VanillaTemporalModule
|
||||
self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len)])
|
||||
else:
|
||||
# down blocks contain two VanillaTemporalModules
|
||||
self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList(
|
||||
[
|
||||
get_motion_module(in_channels, temporal_position_encoding_max_len),
|
||||
get_motion_module(in_channels, temporal_position_encoding_max_len)
|
||||
]
|
||||
)
|
||||
# up blocks contain one additional VanillaTemporalModule
|
||||
if block_type == BlockType.UP:
|
||||
self.motion_modules.append(get_motion_module(in_channels, temporal_position_encoding_max_len))
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_video_length(video_length, full_length)
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_masks(self, masks: Tensor, min_val: float, max_val: float):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_masks(masks, min_val, max_val)
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_sub_idxs(sub_idxs)
|
||||
|
||||
def set_strength(self, strength: float):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_strength(strength)
|
||||
|
||||
def reset_temp_vars(self):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.reset_temp_vars()
|
||||
|
||||
|
||||
def get_motion_module(in_channels, temporal_position_encoding_max_len):
|
||||
# unlike normal AD, there is only one attention block expected in SparseCtrl models
|
||||
return VanillaTemporalModule(in_channels=in_channels, attention_block_types=("Temporal_Self",), temporal_position_encoding_max_len=temporal_position_encoding_max_len)
|
||||
|
||||
|
||||
class VanillaTemporalModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads=8,
|
||||
num_transformer_block=1,
|
||||
attention_block_types=("Temporal_Self", "Temporal_Self"),
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=True,
|
||||
temporal_position_encoding_max_len=24,
|
||||
temporal_attention_dim_div=1,
|
||||
zero_initialize=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.strength = 1.0
|
||||
self.temporal_transformer = TemporalTransformer3DModel(
|
||||
in_channels=in_channels,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=in_channels
|
||||
// num_attention_heads
|
||||
// temporal_attention_dim_div,
|
||||
num_layers=num_transformer_block,
|
||||
attention_block_types=attention_block_types,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
|
||||
if zero_initialize:
|
||||
self.temporal_transformer.proj_out = zero_module(
|
||||
self.temporal_transformer.proj_out
|
||||
)
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
self.temporal_transformer.set_video_length(video_length, full_length)
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
self.temporal_transformer.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_masks(self, masks: Tensor, min_val: float, max_val: float):
|
||||
self.temporal_transformer.set_masks(masks, min_val, max_val)
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
self.temporal_transformer.set_sub_idxs(sub_idxs)
|
||||
|
||||
def set_strength(self, strength: float):
|
||||
self.strength = strength
|
||||
|
||||
def reset_temp_vars(self):
|
||||
self.set_strength(1.0)
|
||||
self.temporal_transformer.reset_temp_vars()
|
||||
|
||||
def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None):
|
||||
if math.isclose(self.strength, 1.0):
|
||||
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)
|
||||
elif math.isclose(self.strength, 0.0):
|
||||
return input_tensor
|
||||
elif self.strength > 1.0:
|
||||
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength
|
||||
else:
|
||||
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + input_tensor*(1.0-self.strength)
|
||||
|
||||
|
||||
class TemporalTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_layers,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
self.video_length = 16
|
||||
self.full_length = 16
|
||||
self.scale_min = 1.0
|
||||
self.scale_max = 1.0
|
||||
self.raw_scale_mask: Union[Tensor, None] = None
|
||||
self.temp_scale_mask: Union[Tensor, None] = None
|
||||
self.sub_idxs: Union[list[int], None] = None
|
||||
self.prev_hidden_states_batch = 0
|
||||
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm = disable_weight_init_clean_groupnorm.GroupNorm(
|
||||
num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks: Iterable[TemporalTransformerBlock] = nn.ModuleList(
|
||||
[
|
||||
TemporalTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
attention_block_types=attention_block_types,
|
||||
dropout=dropout,
|
||||
norm_num_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
self.video_length = video_length
|
||||
self.full_length = full_length
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
for block in self.transformer_blocks:
|
||||
block.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_masks(self, masks: Tensor, min_val: float, max_val: float):
|
||||
self.scale_min = min_val
|
||||
self.scale_max = max_val
|
||||
self.raw_scale_mask = masks
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
self.sub_idxs = sub_idxs
|
||||
for block in self.transformer_blocks:
|
||||
block.set_sub_idxs(sub_idxs)
|
||||
|
||||
def reset_temp_vars(self):
|
||||
del self.temp_scale_mask
|
||||
self.temp_scale_mask = None
|
||||
self.prev_hidden_states_batch = 0
|
||||
|
||||
def get_scale_mask(self, hidden_states: Tensor) -> Union[Tensor, None]:
|
||||
# if no raw mask, return None
|
||||
if self.raw_scale_mask is None:
|
||||
return None
|
||||
shape = hidden_states.shape
|
||||
batch, channel, height, width = shape
|
||||
# if temp mask already calculated, return it
|
||||
if self.temp_scale_mask != None:
|
||||
# check if hidden_states batch matches
|
||||
if batch == self.prev_hidden_states_batch:
|
||||
if self.sub_idxs is not None:
|
||||
return self.temp_scale_mask[:, self.sub_idxs, :]
|
||||
return self.temp_scale_mask
|
||||
# if does not match, reset cached temp_scale_mask and recalculate it
|
||||
del self.temp_scale_mask
|
||||
self.temp_scale_mask = None
|
||||
# otherwise, calculate temp mask
|
||||
self.prev_hidden_states_batch = batch
|
||||
mask = prepare_mask_batch(self.raw_scale_mask, shape=(self.full_length, 1, height, width))
|
||||
mask = repeat_to_batch_size(mask, self.full_length)
|
||||
# if mask not the same amount length as full length, make it match
|
||||
if self.full_length != mask.shape[0]:
|
||||
mask = broadcast_image_to(mask, self.full_length, 1)
|
||||
# reshape mask to attention K shape (h*w, latent_count, 1)
|
||||
batch, channel, height, width = mask.shape
|
||||
# first, perform same operations as on hidden_states,
|
||||
# turning (b, c, h, w) -> (b, h*w, c)
|
||||
mask = mask.permute(0, 2, 3, 1).reshape(batch, height*width, channel)
|
||||
# then, make it the same shape as attention's k, (h*w, b, c)
|
||||
mask = mask.permute(1, 0, 2)
|
||||
# make masks match the expected length of h*w
|
||||
batched_number = shape[0] // self.video_length
|
||||
if batched_number > 1:
|
||||
mask = torch.cat([mask] * batched_number, dim=0)
|
||||
# cache mask and set to proper device
|
||||
self.temp_scale_mask = mask
|
||||
# move temp_scale_mask to proper dtype + device
|
||||
self.temp_scale_mask = self.temp_scale_mask.to(dtype=hidden_states.dtype, device=hidden_states.device)
|
||||
# return subset of masks, if needed
|
||||
if self.sub_idxs is not None:
|
||||
return self.temp_scale_mask[:, self.sub_idxs, :]
|
||||
return self.temp_scale_mask
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
batch, channel, height, width = hidden_states.shape
|
||||
residual = hidden_states
|
||||
scale_mask = self.get_scale_mask(hidden_states)
|
||||
# add some casts for fp8 purposes - does not affect speed otherwise
|
||||
hidden_states = self.norm(hidden_states).to(hidden_states.dtype)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(
|
||||
batch, height * width, inner_dim
|
||||
)
|
||||
hidden_states = self.proj_in(hidden_states).to(hidden_states.dtype)
|
||||
|
||||
# Transformer Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
video_length=self.video_length,
|
||||
scale_mask=scale_mask
|
||||
)
|
||||
|
||||
# output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = (
|
||||
hidden_states.reshape(batch, height, width, inner_dim)
|
||||
.permute(0, 3, 1, 2)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
attention_blocks = []
|
||||
norms = []
|
||||
|
||||
for block_name in attention_block_types:
|
||||
attention_blocks.append(
|
||||
VersatileAttention(
|
||||
attention_mode=block_name.split("_")[0],
|
||||
context_dim=cross_attention_dim # called context_dim for ComfyUI impl
|
||||
if block_name.endswith("_Cross")
|
||||
else None,
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
#bias=attention_bias, # remove for Comfy CrossAttention
|
||||
#upcast_attention=upcast_attention, # remove for Comfy CrossAttention
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
)
|
||||
norms.append(nn.LayerNorm(dim))
|
||||
|
||||
self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks)
|
||||
self.norms = nn.ModuleList(norms)
|
||||
|
||||
self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"))
|
||||
self.ff_norm = nn.LayerNorm(dim)
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
for block in self.attention_blocks:
|
||||
block.set_scale_multiplier(multiplier)
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
for block in self.attention_blocks:
|
||||
block.set_sub_idxs(sub_idxs)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
video_length=None,
|
||||
scale_mask=None
|
||||
):
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(hidden_states).to(hidden_states.dtype)
|
||||
hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states
|
||||
if attention_block.is_cross_attention
|
||||
else None,
|
||||
attention_mask=attention_mask,
|
||||
video_length=video_length,
|
||||
scale_mask=scale_mask
|
||||
)
|
||||
+ hidden_states
|
||||
)
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(self, d_model, dropout=0.0, max_len=24):
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
position = torch.arange(max_len).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)
|
||||
)
|
||||
pe = torch.zeros(1, max_len, d_model)
|
||||
pe[0, :, 0::2] = torch.sin(position * div_term)
|
||||
pe[0, :, 1::2] = torch.cos(position * div_term)
|
||||
self.register_buffer("pe", pe)
|
||||
self.sub_idxs = None
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
self.sub_idxs = sub_idxs
|
||||
|
||||
def forward(self, x):
|
||||
#if self.sub_idxs is not None:
|
||||
# x = x + self.pe[:, self.sub_idxs]
|
||||
#else:
|
||||
x = x + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class CrossAttentionMM(nn.Module):
|
||||
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None,
|
||||
operations=comfy.ops.disable_weight_init):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.scale = None
|
||||
|
||||
self.to_q = operations.Linear(query_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
||||
self.to_k = operations.Linear(context_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
||||
self.to_v = operations.Linear(context_dim, inner_dim, bias=False, dtype=dtype, device=device)
|
||||
|
||||
self.to_out = nn.Sequential(operations.Linear(inner_dim, query_dim, dtype=dtype, device=device), nn.Dropout(dropout))
|
||||
|
||||
def forward(self, x, context=None, value=None, mask=None, scale_mask=None):
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
k: Tensor = self.to_k(context)
|
||||
if value is not None:
|
||||
v = self.to_v(value)
|
||||
del value
|
||||
else:
|
||||
v = self.to_v(context)
|
||||
|
||||
# apply custom scale by multiplying k by scale factor
|
||||
if self.scale is not None:
|
||||
k *= self.scale
|
||||
|
||||
# apply scale mask, if present
|
||||
if scale_mask is not None:
|
||||
k *= scale_mask
|
||||
|
||||
out = optimized_attention(q, k, v, self.heads, mask)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class VersatileAttention(CrossAttentionMM):
|
||||
def __init__(
|
||||
self,
|
||||
attention_mode=None,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
assert attention_mode == "Temporal"
|
||||
|
||||
self.attention_mode = attention_mode
|
||||
self.is_cross_attention = kwargs["context_dim"] is not None
|
||||
|
||||
self.pos_encoder = (
|
||||
PositionalEncoding(
|
||||
kwargs["query_dim"],
|
||||
dropout=0.0,
|
||||
max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
if (temporal_position_encoding and attention_mode == "Temporal")
|
||||
else None
|
||||
)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
|
||||
|
||||
def set_scale_multiplier(self, multiplier: Union[float, None]):
|
||||
if multiplier is None or math.isclose(multiplier, 1.0):
|
||||
self.scale = None
|
||||
else:
|
||||
self.scale = multiplier
|
||||
|
||||
def set_sub_idxs(self, sub_idxs: list[int]):
|
||||
if self.pos_encoder != None:
|
||||
self.pos_encoder.set_sub_idxs(sub_idxs)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: Tensor,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
video_length=None,
|
||||
scale_mask=None,
|
||||
):
|
||||
if self.attention_mode != "Temporal":
|
||||
raise NotImplementedError
|
||||
|
||||
d = hidden_states.shape[1]
|
||||
hidden_states = rearrange(
|
||||
hidden_states, "(b f) d c -> (b d) f c", f=video_length
|
||||
)
|
||||
|
||||
if self.pos_encoder is not None:
|
||||
hidden_states = self.pos_encoder(hidden_states).to(hidden_states.dtype)
|
||||
|
||||
encoder_hidden_states = (
|
||||
repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d)
|
||||
if encoder_hidden_states is not None
|
||||
else encoder_hidden_states
|
||||
)
|
||||
|
||||
hidden_states = super().forward(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
value=None,
|
||||
mask=attention_mask,
|
||||
scale_mask=scale_mask,
|
||||
)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
|
||||
|
||||
return hidden_states
|
||||
+19
-11
@@ -3,13 +3,13 @@ from torch import Tensor
|
||||
|
||||
import folder_paths
|
||||
|
||||
from .control import load_controlnet, convert_to_advanced, ControlWeights, ControlWeightType,\
|
||||
LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet
|
||||
from .control import StrengthInterpolation as SI
|
||||
from .weight_nodes import DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights
|
||||
from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode
|
||||
from .deprecated_nodes import LoadImagesFromDirectory
|
||||
from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet
|
||||
from .utils import ControlWeights, ControlWeightType, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup
|
||||
from .utils import StrengthInterpolation as SI
|
||||
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights,
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights)
|
||||
from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode
|
||||
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -214,8 +214,12 @@ NODE_CLASS_MAPPINGS = {
|
||||
"SoftT2IAdapterWeights": SoftT2IAdapterWeights,
|
||||
"CustomT2IAdapterWeights": CustomT2IAdapterWeights,
|
||||
"ACN_DefaultUniversalWeights": DefaultWeights,
|
||||
# Image
|
||||
"LoadImagesFromDirectory": LoadImagesFromDirectory
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": RgbSparseCtrlPreprocessor,
|
||||
"ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced,
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced,
|
||||
"ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode,
|
||||
"ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -238,6 +242,10 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝",
|
||||
"CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝",
|
||||
"ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝",
|
||||
# Image
|
||||
"LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝"
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝",
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ import torch
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
from .control import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe
|
||||
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@ from typing import Union
|
||||
import numpy as np
|
||||
from collections.abc import Iterable
|
||||
|
||||
from .control import LatentKeyframe, LatentKeyframeGroup
|
||||
from .control import StrengthInterpolation as SI
|
||||
from .utils import LatentKeyframe, LatentKeyframeGroup
|
||||
from .utils import StrengthInterpolation as SI
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
from torch import Tensor
|
||||
|
||||
import folder_paths
|
||||
from nodes import VAEEncode
|
||||
import comfy.utils
|
||||
|
||||
from .utils import TimestepKeyframeGroup
|
||||
from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper
|
||||
from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced
|
||||
|
||||
|
||||
# node for SparseCtrl loading
|
||||
class SparseCtrlLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
"use_motion": ("BOOLEAN", {"default": True}, ),
|
||||
"motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"optional": {
|
||||
"sparse_method": ("SPARSE_METHOD", ),
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None):
|
||||
sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name)
|
||||
sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale)
|
||||
sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings)
|
||||
return (sparsectrl,)
|
||||
|
||||
|
||||
class SparseCtrlMergedLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
"use_motion": ("BOOLEAN", {"default": True}, ),
|
||||
"motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"optional": {
|
||||
"sparse_method": ("SPARSE_METHOD", ),
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/experimental"
|
||||
|
||||
def load_controlnet(self, sparsectrl_name: str, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None):
|
||||
sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name)
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale, merged=True)
|
||||
# first, load normal controlnet
|
||||
controlnet = load_controlnet(controlnet_path, timestep_keyframe=tk_optional)
|
||||
# confirm that controlnet is ControlNetAdvanced
|
||||
if controlnet is None or type(controlnet) != ControlNetAdvanced:
|
||||
raise ValueError(f"controlnet_path must point to a normal ControlNet, but instead: {type(controlnet).__name__}")
|
||||
# next, load sparsectrl, making sure to load motion portion
|
||||
sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=SparseSettings.default())
|
||||
# now, combine state dicts
|
||||
new_state_dict = controlnet.control_model.state_dict()
|
||||
for key, value in sparsectrl.control_model.motion_holder.motion_wrapper.state_dict().items():
|
||||
new_state_dict[key] = value
|
||||
# now, reload sparsectrl with real settings
|
||||
sparsectrl = load_sparsectrl(sparsectrl_path, controlnet_data=new_state_dict, timestep_keyframe=tk_optional, sparse_settings=sparse_settings)
|
||||
return (sparsectrl,)
|
||||
|
||||
|
||||
class SparseIndexMethodNode:
|
||||
@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 (SparseIndexMethod(idxs),)
|
||||
|
||||
|
||||
class SparseSpreadMethodNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"spread": (SparseSpreadMethod.LIST,),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SPARSE_METHOD",)
|
||||
FUNCTION = "get_method"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def get_method(self, spread: str):
|
||||
return (SparseSpreadMethod(spread=spread),)
|
||||
|
||||
|
||||
class RgbSparseCtrlPreprocessor:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"vae": ("VAE", ),
|
||||
"latent_size": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("proc_IMAGE",)
|
||||
FUNCTION = "preprocess_images"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/preprocess"
|
||||
|
||||
def preprocess_images(self, vae, image: Tensor, latent_size: Tensor):
|
||||
# first, resize image to match latents
|
||||
image = image.movedim(-1,1)
|
||||
image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center")
|
||||
image = image.movedim(1,-1)
|
||||
# then, vae encode
|
||||
image = VAEEncode.vae_encode_crop_pixels(image)
|
||||
encoded = vae.encode(image[:,:,:,:3])
|
||||
return (PreprocSparseRGBWrapper(condhint=encoded),)
|
||||
@@ -1,6 +1,6 @@
|
||||
from torch import Tensor
|
||||
import torch
|
||||
from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -1,12 +0,0 @@
|
||||
class AnimateDiffLoaderWithContext:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = ""
|
||||
@@ -0,0 +1,601 @@
|
||||
from typing import Callable, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
import torch.nn.functional as F
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
from comfy.controlnet import ControlBase, broadcast_image_to
|
||||
|
||||
|
||||
def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable):
|
||||
def load_torch_file_with_dict(*args, **kwargs):
|
||||
# immediately restore load_torch_file to original version
|
||||
comfy.utils.load_torch_file = orig_load_torch_file
|
||||
return controlnet_data
|
||||
return load_torch_file_with_dict
|
||||
|
||||
|
||||
def get_properly_arranged_t2i_weights(initial_weights: list[float]):
|
||||
new_weights = []
|
||||
new_weights.extend([initial_weights[0]]*3)
|
||||
new_weights.extend([initial_weights[1]]*3)
|
||||
new_weights.extend([initial_weights[2]]*3)
|
||||
new_weights.extend([initial_weights[3]]*3)
|
||||
return new_weights
|
||||
|
||||
|
||||
class ControlWeightType:
|
||||
DEFAULT = "default"
|
||||
UNIVERSAL = "universal"
|
||||
T2IADAPTER = "t2iadapter"
|
||||
CONTROLNET = "controlnet"
|
||||
CONTROLLORA = "controllora"
|
||||
CONTROLLLLITE = "controllllite"
|
||||
SPARSECTRL = "sparsectrl"
|
||||
|
||||
|
||||
class ControlWeights:
|
||||
def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None):
|
||||
self.weight_type = weight_type
|
||||
self.base_multiplier = base_multiplier
|
||||
self.flip_weights = flip_weights
|
||||
self.weights = weights
|
||||
if self.weights is not None and self.flip_weights:
|
||||
self.weights.reverse()
|
||||
self.weight_mask = weight_mask
|
||||
|
||||
def get(self, idx: int) -> Union[float, Tensor]:
|
||||
# if weights is not none, return index
|
||||
if self.weights is not None:
|
||||
return self.weights[idx]
|
||||
return 1.0
|
||||
|
||||
@classmethod
|
||||
def default(cls):
|
||||
return cls(ControlWeightType.DEFAULT)
|
||||
|
||||
@classmethod
|
||||
def universal(cls, base_multiplier: float, flip_weights: bool=False):
|
||||
return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def universal_mask(cls, weight_mask: Tensor):
|
||||
return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask)
|
||||
|
||||
@classmethod
|
||||
def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*12
|
||||
return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controlnet(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*13
|
||||
return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controllora(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
weights = [1.0]*10
|
||||
return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
@classmethod
|
||||
def controllllite(cls, weights: list[float]=None, flip_weights: bool=False):
|
||||
if weights is None:
|
||||
# TODO: make this have a real value
|
||||
weights = [1.0]*200
|
||||
return cls(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights)
|
||||
|
||||
|
||||
class StrengthInterpolation:
|
||||
LINEAR = "linear"
|
||||
EASE_IN = "ease-in"
|
||||
EASE_OUT = "ease-out"
|
||||
EASE_IN_OUT = "ease-in-out"
|
||||
NONE = "none"
|
||||
|
||||
|
||||
class LatentKeyframe:
|
||||
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 LatentKeyframeGroup:
|
||||
def __init__(self) -> None:
|
||||
self.keyframes: list[LatentKeyframe] = []
|
||||
|
||||
def add(self, keyframe: LatentKeyframe) -> 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[LatentKeyframe, None]:
|
||||
try:
|
||||
return self.keyframes[index]
|
||||
except IndexError:
|
||||
return None
|
||||
|
||||
def __getitem__(self, index) -> LatentKeyframe:
|
||||
return self.keyframes[index]
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self) -> 'LatentKeyframeGroup':
|
||||
cloned = LatentKeyframeGroup()
|
||||
for tk in self.keyframes:
|
||||
cloned.add(tk)
|
||||
return cloned
|
||||
|
||||
|
||||
class TimestepKeyframe:
|
||||
def __init__(self,
|
||||
start_percent: float = 0.0,
|
||||
strength: float = 1.0,
|
||||
interpolation: str = StrengthInterpolation.NONE,
|
||||
control_weights: ControlWeights = None,
|
||||
latent_keyframes: LatentKeyframeGroup = 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) -> 'TimestepKeyframe':
|
||||
return cls(0.0)
|
||||
|
||||
|
||||
# always maintain sorted state (by start_percent of TimestepKeyFrame)
|
||||
class TimestepKeyframeGroup:
|
||||
def __init__(self) -> None:
|
||||
self.keyframes: list[TimestepKeyframe] = []
|
||||
self.keyframes.append(TimestepKeyframe.default())
|
||||
|
||||
def add(self, keyframe: TimestepKeyframe) -> 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[TimestepKeyframe, 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) -> TimestepKeyframe:
|
||||
return self.keyframes[index]
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.keyframes)
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self) -> 'TimestepKeyframeGroup':
|
||||
cloned = TimestepKeyframeGroup()
|
||||
for tk in self.keyframes:
|
||||
cloned.add(tk)
|
||||
return cloned
|
||||
|
||||
@classmethod
|
||||
def default(cls, keyframe: TimestepKeyframe) -> 'TimestepKeyframeGroup':
|
||||
group = cls()
|
||||
group.keyframes[0] = keyframe
|
||||
return group
|
||||
|
||||
|
||||
# depending on model, AnimateDiff may inject into GroupNorm, so make sure GroupNorm will be clean
|
||||
class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init):
|
||||
class GroupNorm(comfy.ops.disable_weight_init.GroupNorm):
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
return F.group_norm(
|
||||
input, self.num_groups, self.weight, self.bias, self.eps)
|
||||
class manual_cast_clean_groupnorm(comfy.ops.manual_cast):
|
||||
class GroupNorm(comfy.ops.manual_cast.GroupNorm):
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
return F.group_norm(
|
||||
input, self.num_groups, self.weight, self.bias, self.eps)
|
||||
|
||||
|
||||
# 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 WeightTypeException(TypeError):
|
||||
"Raised when weight not compatible with AdvancedControlBase object"
|
||||
pass
|
||||
|
||||
|
||||
class AdvancedControlBase:
|
||||
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights):
|
||||
self.base = base
|
||||
self.compatible_weights = [ControlWeightType.UNIVERSAL]
|
||||
self.add_compatible_weight(weights_default.weight_type)
|
||||
# mask for which parts of controlnet output to keep
|
||||
self.mask_cond_hint_original = None
|
||||
self.mask_cond_hint = None
|
||||
self.tk_mask_cond_hint_original = None
|
||||
self.tk_mask_cond_hint = None
|
||||
self.weight_mask_cond_hint = None
|
||||
# actual index values
|
||||
self.sub_idxs = None
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
# timesteps
|
||||
self.t: Tensor = None
|
||||
self.batched_number: int = None
|
||||
# weights + override
|
||||
self.weights: ControlWeights = None
|
||||
self.weights_default: ControlWeights = weights_default
|
||||
self.weights_override: ControlWeights = None
|
||||
# latent keyframe + override
|
||||
self.latent_keyframes: LatentKeyframeGroup = None
|
||||
self.latent_keyframe_override: LatentKeyframeGroup = None
|
||||
# initialize timestep_keyframes
|
||||
self.set_timestep_keyframes(timestep_keyframes)
|
||||
# override some functions
|
||||
self.get_control = self.get_control_inject
|
||||
self.control_merge = self.control_merge_inject#.__get__(self, type(self))
|
||||
self.pre_run = self.pre_run_inject
|
||||
self.cleanup = self.cleanup_inject
|
||||
|
||||
def add_compatible_weight(self, control_weight_type: str):
|
||||
self.compatible_weights.append(control_weight_type)
|
||||
|
||||
def verify_all_weights(self, throw_error=True):
|
||||
# first, check if override exists - if so, only need to check the override
|
||||
if self.weights_override is not None:
|
||||
if self.weights_override.weight_type not in self.compatible_weights:
|
||||
msg = f"Weight override is type {self.weights_override.weight_type}, but loaded {type(self).__name__}" + \
|
||||
f"only supports {self.compatible_weights} weights."
|
||||
raise WeightTypeException(msg)
|
||||
# otherwise, check all timestep keyframe weights
|
||||
else:
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
if tk.has_control_weights() and tk.control_weights.weight_type not in self.compatible_weights:
|
||||
msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type" + \
|
||||
f"{tk.control_weights.weight_type}, but loaded {type(self).__name__} only supports {self.compatible_weights} weights."
|
||||
raise WeightTypeException(msg)
|
||||
|
||||
def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroup):
|
||||
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
|
||||
# prepare first timestep_keyframe related stuff
|
||||
self.current_timestep_keyframe = None
|
||||
self.current_timestep_index = -1
|
||||
self.next_timestep_keyframe = None
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
|
||||
def prepare_current_timestep(self, t: Tensor, batched_number: int):
|
||||
self.t = t
|
||||
self.batched_number = batched_number
|
||||
# get current step percent
|
||||
curr_t: float = t[0]
|
||||
prev_index = self.current_timestep_index
|
||||
# if has next index, loop through and see if need to switch
|
||||
if self.timestep_keyframes.has_index(self.current_timestep_index+1):
|
||||
for i in range(self.current_timestep_index+1, len(self.timestep_keyframes)):
|
||||
eval_tk = self.timestep_keyframes[i]
|
||||
# check if start percent is less or equal to curr_t
|
||||
if eval_tk.start_t >= curr_t:
|
||||
self.current_timestep_index = i
|
||||
self.current_timestep_keyframe = eval_tk
|
||||
# keep track of control weights, latent keyframes, and masks,
|
||||
# accounting for inherit_missing
|
||||
if self.current_timestep_keyframe.has_control_weights():
|
||||
self.weights = self.current_timestep_keyframe.control_weights
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
self.weights = self.weights_default
|
||||
if self.current_timestep_keyframe.has_latent_keyframes():
|
||||
self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
self.latent_keyframes = None
|
||||
if self.current_timestep_keyframe.has_mask_hint():
|
||||
self.tk_mask_cond_hint_original = self.current_timestep_keyframe.mask_hint_orig
|
||||
elif not self.current_timestep_keyframe.inherit_missing:
|
||||
del self.tk_mask_cond_hint_original
|
||||
self.tk_mask_cond_hint_original = None
|
||||
# if guarantee_usage, stop searching for other TKs
|
||||
if self.current_timestep_keyframe.guarantee_usage:
|
||||
break
|
||||
# if eval_tk is outside of percent range, stop looking further
|
||||
else:
|
||||
break
|
||||
|
||||
# if index changed, apply overrides
|
||||
if prev_index != self.current_timestep_index:
|
||||
if self.weights_override is not None:
|
||||
self.weights = self.weights_override
|
||||
if self.latent_keyframe_override is not None:
|
||||
self.latent_keyframes = self.latent_keyframe_override
|
||||
|
||||
# make sure weights and latent_keyframes are in a workable state
|
||||
# Note: each AdvancedControlBase should create their own get_universal_weights class
|
||||
self.prepare_weights()
|
||||
|
||||
def prepare_weights(self):
|
||||
if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT:
|
||||
self.weights = self.weights_default
|
||||
elif self.weights.weight_type == ControlWeightType.UNIVERSAL:
|
||||
# if universal and weight_mask present, no need to convert
|
||||
if self.weights.weight_mask is not None:
|
||||
return
|
||||
self.weights = self.get_universal_weights()
|
||||
|
||||
def get_universal_weights(self) -> ControlWeights:
|
||||
return self.weights
|
||||
|
||||
def set_cond_hint_mask(self, mask_hint):
|
||||
self.mask_cond_hint_original = mask_hint
|
||||
return self
|
||||
|
||||
def pre_run_inject(self, model, percent_to_timestep_function):
|
||||
self.base.pre_run(model, percent_to_timestep_function)
|
||||
self.pre_run_advanced(model, percent_to_timestep_function)
|
||||
|
||||
def pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
# for each timestep keyframe, calculate the start_t
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
tk.start_t = percent_to_timestep_function(tk.start_percent)
|
||||
# clear variables
|
||||
self.cleanup_advanced()
|
||||
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number):
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
# if should not perform any actions for the controlnet, exit without doing any work
|
||||
if self.strength == 0.0 or self.current_timestep_keyframe.strength == 0.0:
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
||||
if control_prev is not None:
|
||||
return control_prev
|
||||
else:
|
||||
return None
|
||||
# otherwise, perform normal function
|
||||
return self.get_control_advanced(x_noisy, t, cond, batched_number)
|
||||
|
||||
def get_control_advanced(self, x_noisy, t, cond, batched_number):
|
||||
pass
|
||||
|
||||
def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]:
|
||||
if self.weights.weight_mask is not None:
|
||||
# prepare weight mask
|
||||
self.prepare_weight_mask_cond_hint(x, self.batched_number)
|
||||
# adjust mask for current layer and return
|
||||
return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, layers=layers))
|
||||
return self.weights.get(idx=idx)
|
||||
|
||||
def get_calc_pow(self, idx: int, layers: int) -> int:
|
||||
return (layers-1)-idx
|
||||
|
||||
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int):
|
||||
# apply strengths, and get batch indeces to null out
|
||||
# AKA latents that should not be influenced by ControlNet
|
||||
if self.latent_keyframes is not None:
|
||||
latent_count = x.size(0)//batched_number
|
||||
indeces_to_null = set(range(latent_count))
|
||||
mapped_indeces = None
|
||||
# if expecting subdivision, will need to translate between subset and actual idx values
|
||||
if self.sub_idxs:
|
||||
mapped_indeces = {}
|
||||
for i, actual in enumerate(self.sub_idxs):
|
||||
mapped_indeces[actual] = i
|
||||
for keyframe in self.latent_keyframes:
|
||||
real_index = keyframe.batch_index
|
||||
# if negative, count from end
|
||||
if real_index < 0:
|
||||
real_index += latent_count if self.sub_idxs is None else self.full_latent_length
|
||||
|
||||
# if not mapping indeces, what you see is what you get
|
||||
if mapped_indeces is None:
|
||||
if real_index in indeces_to_null:
|
||||
indeces_to_null.remove(real_index)
|
||||
# otherwise, see if batch_index is even included in this set of latents
|
||||
else:
|
||||
real_index = mapped_indeces.get(real_index, None)
|
||||
if real_index is None:
|
||||
continue
|
||||
indeces_to_null.remove(real_index)
|
||||
|
||||
# if real_index is outside the bounds of latents, don't apply
|
||||
if real_index >= latent_count or real_index < 0:
|
||||
continue
|
||||
|
||||
# apply strength for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength
|
||||
|
||||
# null them out by multiplying by null_latent_kf_strength
|
||||
for batch_index in indeces_to_null:
|
||||
# apply null for each batched cond/uncond
|
||||
for b in range(batched_number):
|
||||
x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength
|
||||
# apply masks, resizing mask to required dims
|
||||
if self.mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
if self.tk_mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
# apply timestep keyframe strengths
|
||||
if self.current_timestep_keyframe.strength != 1.0:
|
||||
x[:] *= self.current_timestep_keyframe.strength
|
||||
|
||||
def control_merge_inject(self: 'AdvancedControlBase', control_input, control_output, control_prev, output_dtype):
|
||||
out = {'input':[], 'middle':[], 'output': []}
|
||||
|
||||
if control_input is not None:
|
||||
for i in range(len(control_input)):
|
||||
key = 'input'
|
||||
x = control_input[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_input))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
out[key].insert(0, x)
|
||||
|
||||
if control_output is not None:
|
||||
for i in range(len(control_output)):
|
||||
if i == (len(control_output) - 1):
|
||||
key = 'middle'
|
||||
index = 0
|
||||
else:
|
||||
key = 'output'
|
||||
index = i
|
||||
x = control_output[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
if self.global_average_pooling:
|
||||
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
|
||||
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_output))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
|
||||
out[key].append(x)
|
||||
if control_prev is not None:
|
||||
for x in ['input', 'middle', 'output']:
|
||||
o = out[x]
|
||||
for i in range(len(control_prev[x])):
|
||||
prev_val = control_prev[x][i]
|
||||
if i >= len(o):
|
||||
o.append(prev_val)
|
||||
elif prev_val is not None:
|
||||
if o[i] is None:
|
||||
o[i] = prev_val
|
||||
else:
|
||||
o[i] += prev_val
|
||||
return out
|
||||
|
||||
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype)
|
||||
self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype)
|
||||
|
||||
def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype)
|
||||
|
||||
def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None):
|
||||
return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True)
|
||||
|
||||
def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
|
||||
# make mask appropriate dimensions, if present
|
||||
if orig_mask is not None:
|
||||
out_mask = getattr(self, attr_name)
|
||||
if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]:
|
||||
self._reset_attr(attr_name)
|
||||
del out_mask
|
||||
# TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM
|
||||
# resize mask and match batch count
|
||||
multiplier = 1 if direct_attn else 8
|
||||
out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier)
|
||||
actual_latent_length = x_noisy.shape[0] // batched_number
|
||||
out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
|
||||
if self.sub_idxs is not None:
|
||||
out_mask = out_mask[self.sub_idxs]
|
||||
# make cond_hint_mask length match x_noise
|
||||
if x_noisy.shape[0] != out_mask.shape[0]:
|
||||
out_mask = broadcast_image_to(out_mask, x_noisy.shape[0], batched_number)
|
||||
# default dtype to be same as x_noisy
|
||||
if dtype is None:
|
||||
dtype = x_noisy.dtype
|
||||
setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device))
|
||||
del out_mask
|
||||
|
||||
def _reset_attr(self, attr_name, new_value=None):
|
||||
if hasattr(self, attr_name):
|
||||
delattr(self, attr_name)
|
||||
setattr(self, attr_name, new_value)
|
||||
|
||||
def cleanup_inject(self):
|
||||
self.base.cleanup()
|
||||
self.cleanup_advanced()
|
||||
|
||||
def cleanup_advanced(self):
|
||||
self.sub_idxs = None
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
self.t = None
|
||||
self.batched_number = None
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
# timestep stuff
|
||||
self.current_timestep_keyframe = None
|
||||
self.next_timestep_keyframe = None
|
||||
self.current_timestep_index = -1
|
||||
# clear mask hints
|
||||
if self.mask_cond_hint is not None:
|
||||
del self.mask_cond_hint
|
||||
self.mask_cond_hint = None
|
||||
if self.tk_mask_cond_hint_original is not None:
|
||||
del self.tk_mask_cond_hint_original
|
||||
self.tk_mask_cond_hint_original = None
|
||||
if self.tk_mask_cond_hint is not None:
|
||||
del self.tk_mask_cond_hint
|
||||
self.tk_mask_cond_hint = None
|
||||
if self.weight_mask_cond_hint is not None:
|
||||
del self.weight_mask_cond_hint
|
||||
self.weight_mask_cond_hint = None
|
||||
|
||||
def copy_to_advanced(self, copied: 'AdvancedControlBase'):
|
||||
copied.mask_cond_hint_original = self.mask_cond_hint_original
|
||||
copied.weights_override = self.weights_override
|
||||
copied.latent_keyframe_override = self.latent_keyframe_override
|
||||
Reference in New Issue
Block a user