From 37b318debc6fc98731edc7ed1135b7424b1a3a7a Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 1 Dec 2024 19:46:47 -0600 Subject: [PATCH] Refacotr SparseCtrl to depend on AnimateDiff-Evolved for AnimateDiffModel definition/creation; allows for code to not be duplicated and any feature that works for AD can be resued/exposed for SparseCtrl --- adv_control/control.py | 47 +- adv_control/control_sparsectrl.py | 820 ++---------------------------- adv_control/dinklink.py | 52 +- 3 files changed, 118 insertions(+), 801 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 20b4663..a6714e8 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -11,7 +11,7 @@ import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlNetSD35, ControlLora, T2IAdapter, StrengthType from comfy.model_patcher import ModelPatcher -from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseSettings, SparseConst, create_sparse_modelpatcher +from .control_sparsectrl import SparseControlNet, SparseSettings, SparseConst, InterfaceAnimateDiffModel, create_sparse_modelpatcher, load_sparsectrl_motionmodel from .control_lllite import LLLiteModule, LLLitePatch, load_controllllite from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, AbstractPreprocWrapper, ControlWeightType, ControlWeights, WeightTypeException, Extras, @@ -343,11 +343,13 @@ class SVDControlNetAdvanced(ControlNetAdvanced): class SparseCtrlAdvanced(ControlNetAdvanced): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, load_device=None, manual_cast_dtype=None): - super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) - self.control_model_wrapped = create_sparse_modelpatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) + def __init__(self, control_model: SparseControlNet, motion_model: InterfaceAnimateDiffModel, + timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, load_device=None, manual_cast_dtype=None): + super().__init__(control_model=None, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + self.control_model = control_model + self.motion_model = motion_model + self.control_model_wrapped: ModelPatcher = create_sparse_modelpatcher(self.control_model, self.motion_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.add_compatible_weight(ControlWeightType.SPARSECTRL) - self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints self.postpone_condhint_latents_check = True if self.control_model.use_simplified_conditioning_embedding: # TODO: allow vae_optional to be used instead of preprocessor @@ -356,7 +358,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() self.model_latent_format = None # latent format for active SD model, NOT controlnet self.preprocessed = False - + + def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int, transformer_options): # normal ControlNet stuff control_prev = None @@ -377,7 +380,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): # 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) + if self.motion_model is not None: + self.motion_model.set_video_length(video_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]: @@ -465,10 +469,10 @@ class SparseCtrlAdvanced(ControlNetAdvanced): 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.model_latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint - if self.control_model.motion_wrapper is not None: - self.control_model.motion_wrapper.reset() - self.control_model.motion_wrapper.set_strength(self.sparse_settings.motion_strength) - self.control_model.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) + if self.motion_model is not None: + self.motion_model.cleanup() + self.motion_model.set_effect(self.sparse_settings.motion_strength) + self.motion_model.set_scale(self.sparse_settings.motion_scale) def cleanup_advanced(self): super().cleanup_advanced() @@ -479,11 +483,16 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.local_sparse_idxs_inverse = None def copy(self): - c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.load_device, self.manual_cast_dtype) + c = SparseCtrlAdvanced(self.control_model, self.motion_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.load_device, self.manual_cast_dtype) self.copy_to(c) self.copy_to_advanced(c) return c + def get_models(self): + to_return = super().get_models() + to_return.extend(self.control_model_wrapped.get_additional_models()) + return to_return + def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) @@ -828,16 +837,12 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim global_average_pooling = True # actually load motion portion of model now - motion_wrapper: SparseCtrlMotionWrapper = SparseCtrlMotionWrapper(motion_data, ops=controlnet_config.get("operations", None)).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}") + motion_model = load_sparsectrl_motionmodel(ckpt_path=ckpt_path, motion_data=motion_data, ops=controlnet_config.get("operations", None)).to(comfy.model_management.unet_dtype()) + # both motion portion and controlnet portions are loaded; ignore motion_model if shouldn't use motion portion + if not sparse_settings.use_motion: + motion_model = None - # 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) + control = SparseCtrlAdvanced(control_model, motion_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/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 5b5a53d..a102baa 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -3,58 +3,30 @@ #and then taken from comfy/cldm/cldm.py and modified again from abc import ABC, abstractmethod -import copy -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.cli_args import args from comfy.cldm.cldm import ControlNet as ControlNetCLDM -from comfy.ldm.modules.attention import SpatialTransformer -from comfy.ldm.modules.attention import attention_basic, attention_pytorch, attention_split, attention_sub_quad, default -from comfy.ldm.modules.attention import FeedForward, SpatialTransformer from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential from comfy.model_patcher import ModelPatcher -from comfy.patcher_extension import CallbacksMP -import comfy.ops -import comfy.model_management -import comfy.utils +from comfy.patcher_extension import PatcherInjection -from .dinklink import get_AnimateDiffModel, get_AnimateDiffInfo +from .dinklink import (InterfaceAnimateDiffInfo, InterfaceAnimateDiffModel, + get_CreateMotionModelPatcher, get_AnimateDiffModel, get_AnimateDiffInfo) from .logger import logger -from .utils import (BIGMAX, AbstractPreprocWrapper, disable_weight_init_clean_groupnorm, - prepare_mask_batch, broadcast_image_to_extend, extend_to_batch_size) +from .utils import (BIGMAX, AbstractPreprocWrapper, disable_weight_init_clean_groupnorm, WrapperConsts) -# until xformers bug is fixed, do not use xformers for VersatileAttention! TODO: change this when fix is out -# logic for choosing optimized_attention method taken from comfy/ldm/modules/attention.py -# a fallback_attention_mm is selected to avoid CUDA configuration limitation with pytorch's scaled_dot_product -optimized_attention_mm = attention_basic -fallback_attention_mm = attention_basic -if comfy.model_management.xformers_enabled(): - pass - #optimized_attention_mm = attention_xformers -if comfy.model_management.pytorch_attention_enabled(): - optimized_attention_mm = attention_pytorch - if args.use_split_cross_attention: - fallback_attention_mm = attention_split - else: - fallback_attention_mm = attention_sub_quad -else: - if args.use_split_cross_attention: - optimized_attention_mm = attention_split - else: - optimized_attention_mm = attention_sub_quad +class SparseMotionModelPatcher(ModelPatcher): + '''Class only used for IDE type hints.''' + def __init__(self, *args, **kwargs): + self.model = InterfaceAnimateDiffModel class SparseConst: @@ -74,11 +46,6 @@ class SparseControlNet(ControlNetCLDM): 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)), ) - self.motion_wrapper: SparseCtrlMotionWrapper = None - - def set_actual_length(self, actual_length: int, full_length: int): - if self.motion_wrapper is not None: - self.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) @@ -112,47 +79,44 @@ class SparseControlNet(ControlNetCLDM): return {"middle": out_middle, "output": out_output} -def create_sparse_modelpatcher(model, load_device, offload_device): +def load_sparsectrl_motionmodel(ckpt_path: str, motion_data: dict[str, Tensor], ops=None) -> InterfaceAnimateDiffModel: + mm_info: InterfaceAnimateDiffInfo = get_AnimateDiffInfo()("SD1.5", "AnimateDiff", "v3", ckpt_path) + init_kwargs = { + "ops": ops, + "get_unet_func": _get_unet_func, + } + motion_model: InterfaceAnimateDiffModel = get_AnimateDiffModel()(mm_state_dict=motion_data, mm_info=mm_info, init_kwargs=init_kwargs) + missing, unexpected = motion_model.load_state_dict(motion_data) + if len(missing) > 0 or len(unexpected) > 0: + logger.info(f"SparseCtrl MotionModel: {missing}, {unexpected}") + return motion_model + + +def create_sparse_modelpatcher(model, motion_model, load_device, offload_device): patcher = ModelPatcher(model, load_device=load_device, offload_device=offload_device) - acn = "ACN" - patcher.add_callback_with_key(CallbacksMP.ON_LOAD, acn, _patch_lowvram_extras) - patcher.add_callback_with_key(CallbacksMP.ON_LOAD, acn, _handle_float8_pe_tensors) + if motion_model is not None: + _motionpatcher = _create_sparse_motionmodelpatcher(motion_model, load_device, offload_device) + patcher.set_additional_models(WrapperConsts.ACN, [_motionpatcher]) + patcher.set_injections(WrapperConsts.ACN, + [PatcherInjection(inject=_inject_motion_models, eject=_eject_motion_models)]) return patcher - -def _patch_lowvram_extras(self: ModelPatcher, device_to, lowvram_model_memory, force_patch_weights, full_load, *args, **kwargs): - if lowvram_model_memory > 0: - motion_wrapper: SparseCtrlMotionWrapper = self.model.motion_wrapper - if motion_wrapper is not None: - # figure out the tensors (likely pe's) that should be cast to device besides just the named_modules - remaining_tensors = list(motion_wrapper.state_dict().keys()) - named_modules = [] - for n, _ in motion_wrapper.named_modules(): - named_modules.append(n) - named_modules.append(f"{n}.weight") - named_modules.append(f"{n}.bias") - for name in named_modules: - if name in remaining_tensors: - remaining_tensors.remove(name) - - for key in remaining_tensors: - self.patch_weight_to_device(key, device_to) - if device_to is not None: - comfy.utils.set_attr(motion_wrapper, key, comfy.utils.get_attr(motion_wrapper, key).to(device_to)) +def _create_sparse_motionmodelpatcher(motion_model, load_device, offload_device) -> SparseMotionModelPatcher: + return get_CreateMotionModelPatcher()(motion_model, load_device, offload_device) -def _handle_float8_pe_tensors(self: ModelPatcher, *args, **kwargs): - motion_wrapper: SparseCtrlMotionWrapper = self.model.motion_wrapper - if motion_wrapper is not None: - remaining_tensors = list(motion_wrapper.state_dict().keys()) - pe_tensors = [x for x in remaining_tensors if '.pe' in x] - is_first = True - for key in pe_tensors: - if is_first: - is_first = False - if comfy.utils.get_attr(motion_wrapper, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]: - break - comfy.utils.set_attr(motion_wrapper, key, comfy.utils.get_attr(motion_wrapper, key).half()) +def _inject_motion_models(patcher: ModelPatcher): + motion_models: list[SparseMotionModelPatcher] = patcher.get_additional_models_with_key(WrapperConsts.ACN) + for mm in motion_models: + mm.model.inject(patcher) + +def _eject_motion_models(patcher: ModelPatcher): + motion_models: list[SparseMotionModelPatcher] = patcher.get_additional_models_with_key(WrapperConsts.ACN) + for mm in motion_models: + mm.model.eject(patcher) + +def _get_unet_func(wrapper, model: ModelPatcher): + return model.model class PreprocSparseRGBWrapper(AbstractPreprocWrapper): @@ -353,705 +317,3 @@ def get_idx_list_from_str(indexes: str) -> list[int]: if len(idxs) == 0: raise ValueError(f"No indexes were listed in Sparse Index Method.") return 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!") - - -# TODO: replace with DinkLink reference from ADE -class SparseCtrlMotionWrapper(nn.Module): - def __init__(self, mm_state_dict: dict[str, Tensor], ops=disable_weight_init_clean_groupnorm): - 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, ops=ops)) - 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, ops=ops)) - 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, ops=ops) - - 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 - unet.motion_wrapper = 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, ops=disable_weight_init_clean_groupnorm): - 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, ops=ops)]) - else: - # down blocks contain two VanillaTemporalModules - self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList( - [ - get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops), - get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops) - ] - ) - # 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, ops=ops)) - - 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, ops=disable_weight_init_clean_groupnorm): - # 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, ops=ops) - - -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, - ops=disable_weight_init_clean_groupnorm, - ): - 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, - ops=ops, - ) - - 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, - ops=disable_weight_init_clean_groupnorm, - ): - 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 = ops.GroupNorm( - num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True - ) - self.proj_in = ops.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, - ops=ops, - ) - for d in range(num_layers) - ] - ) - self.proj_out = ops.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 - for block in self.transformer_blocks: - block.reset_temp_vars() - - 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 = extend_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_extend(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, - ops=disable_weight_init_clean_groupnorm, - ): - 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, - ops=ops, - ) - ) - norms.append(ops.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"), operations=ops) - self.ff_norm = ops.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 reset_temp_vars(self): - for block in self.attention_blocks: - block.reset_temp_vars() - - 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 CrossAttentionMMSparse(nn.Module): - def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None, - operations=disable_weight_init_clean_groupnorm): - super().__init__() - inner_dim = dim_head * heads - context_dim = default(context_dim, query_dim) - - self.actual_attention = optimized_attention_mm - 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 reset_attention_type(self): - self.actual_attention = optimized_attention_mm - - 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 - - try: - out = self.actual_attention(q, k, v, self.heads, mask) - except RuntimeError as e: - if str(e).startswith("CUDA error: invalid configuration argument"): - self.actual_attention = fallback_attention_mm - out = self.actual_attention(q, k, v, self.heads, mask) - else: - raise - return self.to_out(out) - - -class VersatileAttention(CrossAttentionMMSparse): - def __init__( - self, - attention_mode=None, - cross_frame_attention_mode=None, - temporal_position_encoding=False, - temporal_position_encoding_max_len=24, - ops=disable_weight_init_clean_groupnorm, - *args, - **kwargs, - ): - super().__init__(operations=ops, *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 reset_temp_vars(self): - self.reset_attention_type() - - 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/adv_control/dinklink.py b/adv_control/dinklink.py index d591d3b..c100971 100644 --- a/adv_control/dinklink.py +++ b/adv_control/dinklink.py @@ -11,6 +11,10 @@ # purposely exposing node pack classes/functions with other node packs. #################################################################################################### from __future__ import annotations +from typing import Union +from torch import Tensor, nn + +from comfy.model_patcher import ModelPatcher import comfy.hooks DINKLINK = "__DINKLINK" @@ -35,10 +39,56 @@ class DinkLinkConst: ADE = "ADE" ADE_ANIMATEDIFFMODEL = "AnimateDiffModel" ADE_ANIMATEDIFFINFO = "AnimateDiffInfo" + ADE_CREATE_MOTIONMODELPATCHER = "create_MotionModelPatcher" def prepare_dinklink(): pass + +class InterfaceAnimateDiffInfo: + '''Class only used for IDE type hints; interface of ADE's AnimateDiffInfo''' + def __init__(self, sd_type: str, mm_format: str, mm_version: str, mm_name: str): + self.sd_type = sd_type + self.mm_format = mm_format + self.mm_version = mm_version + self.mm_name = mm_name + + +class InterfaceAnimateDiffModel(nn.Module): + '''Class only used for IDE type hints; interface of ADE's AnimateDiffModel''' + def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: InterfaceAnimateDiffInfo, init_kwargs: dict[str]={}): + pass + + def set_video_length(self, video_length: int, full_length: int) -> None: + raise NotImplemented() + + def set_scale(self, scale: Union[float, Tensor, None], per_block_list: Union[list, None]=None) -> None: + raise NotImplemented() + + def set_effect(self, multival: Union[float, Tensor, None], per_block_list: Union[list, None]=None) -> None: + raise NotImplemented() + + def cleanup(self): + raise NotImplemented() + + def inject(self, model: ModelPatcher): + pass + + def eject(self, model: ModelPatcher): + pass + + +def get_CreateMotionModelPatcher(throw_exception=True): + d = get_dinklink() + try: + link_ade = d[DinkLinkConst.ADE] + return link_ade[DinkLinkConst.ADE_CREATE_MOTIONMODELPATCHER] + except KeyError: + if throw_exception: + raise Exception("Could not get create_MotionModelPatcher function. AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \ + "they are either not installed or are of an insufficient version.") + return None + def get_AnimateDiffModel(throw_exception=True): d = get_dinklink() try: @@ -50,7 +100,7 @@ def get_AnimateDiffModel(throw_exception=True): "they are either not installed or are of an insufficient version.") return None -def get_AnimateDiffInfo(throw_exception=True): +def get_AnimateDiffInfo(throw_exception=True) -> InterfaceAnimateDiffInfo: d = get_dinklink() try: link_ade = d[DinkLinkConst.ADE]