From 88cc0ac149ee644927e2a87f08a0b0f3f9cd3731 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 18 Dec 2023 01:01:47 -0600 Subject: [PATCH] Progress on SparseCtrl support --- control/control.py | 413 +++---- control/control_sparsectrl.py | 1035 +++++++++++++++++ control/nodes.py | 14 +- ...eprecated_nodes.py => nodes_deprecated.py} | 2 +- ...rame_nodes.py => nodes_latent_keyframe.py} | 4 +- control/nodes_reference.py | 0 control/nodes_sparsectrl.py | 28 + control/{weight_nodes.py => nodes_weight.py} | 2 +- control/reference_nodes.py | 12 - control/utils.py | 260 +++++ 10 files changed, 1518 insertions(+), 252 deletions(-) create mode 100644 control/control_sparsectrl.py rename control/{deprecated_nodes.py => nodes_deprecated.py} (97%) rename control/{latent_keyframe_nodes.py => nodes_latent_keyframe.py} (99%) create mode 100644 control/nodes_reference.py create mode 100644 control/nodes_sparsectrl.py rename control/{weight_nodes.py => nodes_weight.py} (98%) delete mode 100644 control/reference_nodes.py create mode 100644 control/utils.py diff --git a/control/control.py b/control/control.py index f5e5495..7ed319f 100644 --- a/control/control.py +++ b/control/control.py @@ -1,223 +1,19 @@ -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 +from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper +from .utils import (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 -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): @@ -765,23 +561,56 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): self.already_patched = False +class SparseCtrlAdvanced(ControlNetAdvanced): + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, 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) + + def copy(self): + c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, 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): 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: + #raise NotImplementedError("SparseCtrl has not been fully implemented yet!") + control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model) # otherwise, load vanilla ControlNet else: - control = comfy_cn.load_controlnet(ckpt_path, model=model) + 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 # from pathlib import Path # with open(Path(__file__).parent.parent.parent / "controlnet_keys.txt", "w") as cfile: # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) @@ -811,25 +640,151 @@ 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, 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) + 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 + motion_wrapper.inject(control_model) + + control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + new_state_dict = control_model.state_dict() + from pathlib import Path + with open(Path(__file__).parent.parent.parent / "sparcectrlstatedict.txt", "w") as cfile: + for key in new_state_dict: + cfile.write(f"{key}\n") + return control diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py new file mode 100644 index 0000000..8e1e15c --- /dev/null +++ b/control/control_sparsectrl.py @@ -0,0 +1,1035 @@ +#taken from: https://github.com/lllyasviel/ControlNet +#and modified +#and then taken from comfy/cldm/cldm.py and modified again + +import math +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 ControlNet_cldm +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(ControlNet_cldm): + 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) + use_simplified_conditioning_embedding = kwargs.get("use_simplified_conditioning_embedding", False) + if use_simplified_conditioning_embedding: + self.input_hint_block = TimestepEmbedSequential( + operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device), + ) + + def forward(self, x, hint, 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) + + x = torch.zeros_like(x) + + conditioning_mask1 = torch.ones_like(hint[:, :1]) + conditioning_mask2 = torch.zeros_like(hint[:, :1]) + conditioning_mask = conditioning_mask2 + conditioning_mask[0] = conditioning_mask1[0] + conditioning_mask[16] = conditioning_mask1[16] + #conditioning_mask[15] = conditioning_mask1[15] + #conditioning_mask[31] = conditioning_mask1[31] + modified_hint = torch.zeros_like(hint) + modified_hint[0] = hint[0] + modified_hint[16] = hint[16] + #modified_hint[15] = hint[15] + #modified_hint[31] = hint[31] + hint = torch.cat([modified_hint, conditioning_mask], dim=1) + 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 + + + +# main class for holding SparseControlNet +class SparseControlNetOld(nn.Module): + def __init__( + self, + image_size, + in_channels, + model_channels, + hint_channels, + num_res_blocks, + dropout=0, + channel_mult=(1, 2, 4, 8), + conv_resample=True, + dims=2, + num_classes=None, + use_checkpoint=False, + dtype=torch.float32, + num_heads=-1, + num_head_channels=-1, + num_heads_upsample=-1, + use_scale_shift_norm=False, + resblock_updown=False, + use_new_attention_order=False, + use_spatial_transformer=False, # custom transformer support + transformer_depth=1, # custom transformer support + context_dim=None, # custom transformer support + n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model + legacy=True, + disable_self_attentions=None, + num_attention_blocks=None, + disable_middle_self_attn=False, + use_linear_in_transformer=False, + adm_in_channels=None, + transformer_depth_middle=None, + transformer_depth_output=None, + device=None, + operations=disable_weight_init_clean_groupnorm, + **kwargs, + ): + super().__init__() + assert use_spatial_transformer == True, "use_spatial_transformer has to be true" + if use_spatial_transformer: + assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...' + + if context_dim is not None: + assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...' + # from omegaconf.listconfig import ListConfig + # if type(context_dim) == ListConfig: + # context_dim = list(context_dim) + if num_heads_upsample == -1: + num_heads_upsample = num_heads + + if num_heads == -1: + assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set' + + if num_head_channels == -1: + assert num_heads != -1, 'Either num_heads or num_head_channels has to be set' + + self.dims = dims + self.image_size = image_size + self.in_channels = in_channels + self.model_channels = model_channels + + if isinstance(num_res_blocks, int): + self.num_res_blocks = len(channel_mult) * [num_res_blocks] + else: + if len(num_res_blocks) != len(channel_mult): + raise ValueError("provide num_res_blocks either as an int (globally constant) or " + "as a list/tuple (per-level) with the same length as channel_mult") + self.num_res_blocks = num_res_blocks + + if disable_self_attentions is not None: + # should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not + assert len(disable_self_attentions) == len(channel_mult) + if num_attention_blocks is not None: + assert len(num_attention_blocks) == len(self.num_res_blocks) + assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks)))) + + transformer_depth = transformer_depth[:] + + self.dropout = dropout + self.channel_mult = channel_mult + self.conv_resample = conv_resample + self.num_classes = num_classes + self.use_checkpoint = use_checkpoint + self.dtype = dtype + self.num_heads = num_heads + self.num_head_channels = num_head_channels + self.num_heads_upsample = num_heads_upsample + self.predict_codebook_ids = n_embed is not None + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + operations.Linear(model_channels, time_embed_dim, dtype=self.dtype, device=device), + nn.SiLU(), + operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), + ) + + if self.num_classes is not None: + if isinstance(self.num_classes, int): + self.label_emb = nn.Embedding(num_classes, time_embed_dim) + elif self.num_classes == "continuous": + print("setting up linear c_adm embedding layer") + self.label_emb = nn.Linear(1, time_embed_dim) + elif self.num_classes == "sequential": + assert adm_in_channels is not None + self.label_emb = nn.Sequential( + nn.Sequential( + operations.Linear(adm_in_channels, time_embed_dim, dtype=self.dtype, device=device), + nn.SiLU(), + operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), + ) + ) + else: + raise ValueError() + + self.input_blocks = nn.ModuleList( + [ + TimestepEmbedSequential( + operations.conv_nd(dims, in_channels, model_channels, 3, padding=1, dtype=self.dtype, device=device) + ) + ] + ) + self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels, operations=operations, dtype=self.dtype, device=device)]) + + self.input_hint_block = TimestepEmbedSequential( + operations.conv_nd(dims, hint_channels, 16, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 16, 16, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 16, 32, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 32, 32, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 32, 96, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 96, 96, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 96, 256, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 256, model_channels, 3, padding=1, dtype=self.dtype, device=device) + ) + + self._feature_size = model_channels + input_block_chans = [model_channels] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations, + ) + ] + ch = mult * model_channels + num_transformers = transformer_depth.pop(0) + if num_transformers > 0: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + #num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + if exists(disable_self_attentions): + disabled_sa = disable_self_attentions[level] + else: + disabled_sa = False + + if not exists(num_attention_blocks) or nr < num_attention_blocks[level]: + layers.append( + SpatialTransformer( + ch, num_heads, dim_head, depth=num_transformers, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations + ) + ) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) + self._feature_size += ch + input_block_chans.append(ch) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + dtype=self.dtype, + device=device, + operations=operations + ) + if resblock_updown + else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch, dtype=self.dtype, device=device, operations=operations + ) + ) + ) + ch = out_ch + input_block_chans.append(ch) + self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) + ds *= 2 + self._feature_size += ch + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + #num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + mid_block = [ + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations + )] + if transformer_depth_middle >= 0: + mid_block += [SpatialTransformer( # always uses a self-attn + ch, num_heads, dim_head, depth=transformer_depth_middle, context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations + ), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations + )] + self.middle_block = TimestepEmbedSequential(*mid_block) + self.middle_block_out = self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device) + self._feature_size += ch + + #self._motion_wrapper: SparseCtrlMotionWrapper = None + + def make_zero_conv(self, channels, operations=None, dtype=None, device=None): + return TimestepEmbedSequential(operations.conv_nd(self.dims, channels, channels, 1, padding=0, dtype=dtype, device=device)) + + def forward(self, x, hint, 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) + + x = torch.zeros_like(x) + + conditioning_mask1 = torch.ones_like(hint[:, :1]) + conditioning_mask2 = torch.zeros_like(hint[:, :1]) + conditioning_mask = conditioning_mask2 + conditioning_mask[0] = conditioning_mask1[0] + conditioning_mask[16] = conditioning_mask1[16] + #conditioning_mask[15] = conditioning_mask1[15] + #conditioning_mask[31] = conditioning_mask1[31] + modified_hint = torch.zeros_like(hint) + modified_hint[0] = hint[0] + modified_hint[16] = hint[16] + #modified_hint[15] = hint[15] + #modified_hint[31] = hint[31] + hint = torch.cat([modified_hint, conditioning_mask], dim=1) + 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 + + +# 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 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_wrapper = 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_wrapper + + 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 + for block in self.down_blocks: + block.set_video_length(video_length, full_length) + 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]): + for block in self.down_blocks: + block.set_scale_multiplier(multiplier) + 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 reset_temp_vars(self): + for block in self.down_blocks: + block.reset_temp_vars() + 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 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.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 reset_temp_vars(self): + self.temporal_transformer.reset_temp_vars() + + def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None): + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + + +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..dd9ecde 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_deprecated import LoadImagesFromDirectory from .logger import logger 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..2e62629 --- /dev/null +++ b/control/nodes_sparsectrl.py @@ -0,0 +1,28 @@ +import folder_paths + +from .utils import TimestepKeyframeGroup +from .control import load_sparsectrl + + +# node for SparseCtrl loading +class SparseCtrlLoaderAdvanced: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + }, + "optional": { + "tk_optional": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + + def load_controlnet(self, control_net_name: str, tk_optional: TimestepKeyframeGroup=None): + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional) + return controlnet 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..7772601 --- /dev/null +++ b/control/utils.py @@ -0,0 +1,260 @@ +from typing import Callable, Union +import torch +from torch import Tensor +import torch.nn.functional as F +import comfy.ops +import comfy.utils + + +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