diff --git a/control/control.py b/control/control.py index e46cd6a..cf66378 100644 --- a/control/control.py +++ b/control/control.py @@ -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 diff --git a/control/control_lllite.py b/control/control_lllite.py index 8b13789..ba327c6 100644 --- a/control/control_lllite.py +++ b/control/control_lllite.py @@ -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 \ No newline at end of file diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py new file mode 100644 index 0000000..d423872 --- /dev/null +++ b/control/control_sparsectrl.py @@ -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 diff --git a/control/nodes.py b/control/nodes.py index 3794ac7..3a071c9 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -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 🛂🅐🅒🅝", } diff --git a/control/deprecated_nodes.py b/control/nodes_deprecated.py similarity index 97% rename from control/deprecated_nodes.py rename to control/nodes_deprecated.py index a64ac9b..93ef08f 100644 --- a/control/deprecated_nodes.py +++ b/control/nodes_deprecated.py @@ -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 diff --git a/control/latent_keyframe_nodes.py b/control/nodes_latent_keyframe.py similarity index 99% rename from control/latent_keyframe_nodes.py rename to control/nodes_latent_keyframe.py index 2fde61e..2716295 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/nodes_latent_keyframe.py @@ -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 diff --git a/control/nodes_reference.py b/control/nodes_reference.py new file mode 100644 index 0000000..e69de29 diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py new file mode 100644 index 0000000..dca2e3f --- /dev/null +++ b/control/nodes_sparsectrl.py @@ -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),) diff --git a/control/weight_nodes.py b/control/nodes_weight.py similarity index 98% rename from control/weight_nodes.py rename to control/nodes_weight.py index f80d607..35d0ffb 100644 --- a/control/weight_nodes.py +++ b/control/nodes_weight.py @@ -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 diff --git a/control/reference_nodes.py b/control/reference_nodes.py deleted file mode 100644 index 6879f97..0000000 --- a/control/reference_nodes.py +++ /dev/null @@ -1,12 +0,0 @@ -class AnimateDiffLoaderWithContext: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "image": ("IMAGE",), - }, - } - - RETURN_TYPES = ("MODEL",) - CATEGORY = "" \ No newline at end of file diff --git a/control/utils.py b/control/utils.py new file mode 100644 index 0000000..edd45a8 --- /dev/null +++ b/control/utils.py @@ -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