From 2d2a5c39bb6dc7bedfce9ef1553c3267ce9118d8 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 10 Nov 2023 01:24:59 -0600 Subject: [PATCH] Added AnimateDiff-SDXL support - you need to use 'linear (AnimateDiff-SDXL)' beta_schedule --- animatediff/model_utils.py | 10 ++- animatediff/motion_module.py | 98 +++++++++++++++++++++++-- animatediff/motion_module_ad.py | 21 +++++- animatediff/motion_module_adxl.py | 118 ++++++++++++++++++++++++++++++ animatediff/motion_module_hsxl.py | 21 +++++- animatediff/motion_utils.py | 5 ++ animatediff/sampling.py | 5 +- 7 files changed, 265 insertions(+), 13 deletions(-) create mode 100644 animatediff/motion_module_adxl.py diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index d57688b..95233e2 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -29,15 +29,17 @@ class ModelSamplingConfig: class BetaSchedules: SQRT_LINEAR = "sqrt_linear (AnimateDiff)" + LINEAR_ADXL = "linear (AnimateDiff-SDXL)" LINEAR = "linear (HotshotXL/default)" SQRT = "sqrt" COSINE = "cosine" SQUAREDCOS_CAP_V2 = "squaredcos_cap_v2" - ALIAS_LIST = [SQRT_LINEAR, LINEAR, SQRT, COSINE, SQUAREDCOS_CAP_V2] + ALIAS_LIST = [SQRT_LINEAR, LINEAR_ADXL, LINEAR, SQRT, COSINE, SQUAREDCOS_CAP_V2] ALIAS_MAP = { SQRT_LINEAR: "sqrt_linear", + LINEAR_ADXL: "linear", # also linear, but has different linear_end (0.020) LINEAR: "linear", SQRT: "sqrt", COSINE: "cosine", @@ -54,7 +56,11 @@ class BetaSchedules: @classmethod def to_model_sampling(cls, alias: str, model: ModelPatcher): - return model_sampling(cls.to_config(alias), model_type=model.model.model_type) + ms_obj = model_sampling(cls.to_config(alias), model_type=model.model.model_type) + if alias == cls.LINEAR_ADXL: + # uses linear_end=0.020 + ms_obj._register_schedule(given_betas=None, beta_schedule=cls.to_name(alias), timesteps=1000, linear_start=0.00085, linear_end=0.020, cosine_s=8e-3) + return ms_obj @staticmethod def get_alias_list_with_first_element(first_element: str): diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index 077cd53..429c5ee 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -13,9 +13,10 @@ from .logger import logger from .model_utils import ModelTypesSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \ get_sd_model_type from .motion_lora import MotionLoRAList, MotionLoRAWrapper -from .motion_module_ad import AnimDiffMotionWrapper, has_mid_block +from .motion_module_ad import AnimDiffMotionWrapper, VanillaTemporalModule, has_mid_block +from .motion_module_adxl import AnimDiffSDXLMotionWrapper from .motion_module_hsxl import HotShotXLMotionWrapper, TransformerTemporal -from .motion_utils import GenericMotionWrapper, InjectorVersion, NoiseType, normalize_min_max +from .motion_utils import GenericMotionWrapper, InjectorVersion, MotionCompatibilityError, NoiseType, normalize_min_max # inject into ModelPatcher.clone to carry over injected params over to cloned ModelPatcher orig_modelpatcher_clone = comfy_model_patcher.ModelPatcher.clone @@ -225,13 +226,19 @@ def load_motion_module(model_name: str, motion_lora: MotionLoRAList = None, mode if sd_model_type == ModelTypesSD.SD1_5: try: motion_module = AnimDiffMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras) - except ValueError as e: + except MotionCompatibilityError as e: raise ValueError(f"Motion model {model_name} is not compatible with SD1.5-based model.", e) elif sd_model_type == ModelTypesSD.SDXL: + # determine if motion module is a AnimateDiffXL model or a HotshotXL model try: + # try to load as HotShotXL model first motion_module = HotShotXLMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras) - except ValueError as e: - raise ValueError(f"Motion model {model_name} is not compatible with SDXL-based model.", e) + except MotionCompatibilityError as e: + # if not compatible, try to load as AnimateDiff-SDXL + try: + motion_module = AnimDiffSDXLMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras) + except MotionCompatibilityError as e: + raise ValueError(f"Motion model {model_name} is not compatible with neither AnimateDiff-SDXL nor HotShotXL.", e) else: raise ValueError(f"SD model must be either SD1.5-based for AnimateDiff or SDXL-based for HotShotXL.") @@ -361,6 +368,85 @@ def _eject_motion_module_from_unet(model: ModelPatcher): ############################################################################################################ +############################################################################################################ +## AnimateDiff-SDXL +def _inject_adxl_motion_module_to_unet(model: ModelPatcher, motion_module: 'AnimDiffSDXLMotionWrapper'): + unet: openaimodel.UNetModel = model.model.diffusion_model + # inject input (down) blocks + # AnimateDiffSDXL mm contains 3 downblocks, each with 2 TransformerTemporals - 6 in total + # per_block is the amount of Temporal Blocks per down block + _perform_adxl_motion_module_injection(unet.input_blocks, motion_module.down_blocks, injection_goal=6, per_block=2) + + # inject output (up) blocks + # AnimateDiffSDXL mm contains 3 upblocks, each with 3 TransformerTemporals - 9 in total + _perform_adxl_motion_module_injection(unet.output_blocks, motion_module.up_blocks, injection_goal=9, per_block=3) + + # inject mid block, if needed (encapsulate in list to make structure compatible) + if motion_module.mid_block is not None: + _perform_adxl_motion_module_injection(unet.middle_block, [motion_module.mid_block], injection_goal=1, per_block=1) + + # keep track of if unet blocks actually affected + set_injected_unet_version(model, InjectorVersion.ADXL_V1_V2) + +def _perform_adxl_motion_module_injection(unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList, injection_goal: int, per_block: int): + # 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 + # 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) == ResBlock and res_idx < 0: + res_idx = idx + # if SpatialTransformer exists, inject right after + if st_idx >= 0: + #logger.info(f"ADXL: injecting after ST({st_idx})") + 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: + #logger.info(f"ADXL: injecting after Res({res_idx})") + 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_adxl_motion_module_from_unet(model: ModelPatcher): + unet: openaimodel.UNetModel = model.model.diffusion_model + # remove from input blocks + _perform_adxl_motion_module_ejection(unet.input_blocks) + # remove from output blocks + _perform_adxl_motion_module_ejection(unet.output_blocks) + # remove from middle block (encapsulate in list to make structure compatible) + _perform_adxl_motion_module_ejection([unet.middle_block]) + # remove attr; ejected + del_injected_unet_version(model) + +def _perform_adxl_motion_module_ejection(unet_blocks: nn.ModuleList): + # eject all TemporalTransformer3DModel 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) + #logger.info(f"ADXL: ejecting {idx_to_pop}") +############################################################################################################ + + ############################################################################################################ ## HotShot XL def _inject_hsxl_motion_module_to_unet(model: ModelPatcher, motion_module: 'HotShotXLMotionWrapper'): @@ -443,11 +529,13 @@ def _perform_hsxl_motion_module_ejection(unet_blocks: nn.ModuleList): injectors = { InjectorVersion.V1_V2: _inject_motion_module_to_unet, + InjectorVersion.ADXL_V1_V2: _inject_adxl_motion_module_to_unet, InjectorVersion.HOTSHOTXL_V1: _inject_hsxl_motion_module_to_unet, } ejectors = { InjectorVersion.V1_V2: _eject_motion_module_from_unet, + InjectorVersion.ADXL_V1_V2: _eject_adxl_motion_module_from_unet, InjectorVersion.HOTSHOTXL_V1: _eject_hsxl_motion_module_from_unet, } diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 895a210..320edf4 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -7,7 +7,7 @@ from torch import Tensor, nn from comfy.ldm.modules.attention import FeedForward from .motion_lora import MotionLoRAInfo -from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric, prepare_mask_batch +from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch def zero_module(module): @@ -22,7 +22,23 @@ def get_ad_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor], 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 mm_state_dict - {mm_type} is not a valid motion module!") + raise MotionCompatibilityError(f"No pos_encoder.pe found in mm_state_dict - {mm_type} is not a valid AnimateDiff-SD1.5 motion module!") + + +def validate_ad_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None: + # keep track of biggest down_block count in module + biggest_block = 0 + for key in mm_state_dict.keys(): + if "down_blocks" 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 + if biggest_block != 3: + raise MotionCompatibilityError(f"Expected biggest down_block to be 3, but was {biggest_block} - {mm_type} is not a valid AnimateDiff-SD1.5 motion module!") def has_mid_block(mm_state_dict: dict[str, Tensor]): @@ -40,6 +56,7 @@ class AnimDiffMotionWrapper(GenericMotionWrapper): self.up_blocks: Iterable[MotionModule] = nn.ModuleList([]) self.mid_block: Union[MotionModule, None] = None self.encoding_max_len = get_ad_temporal_position_encoding_max_len(mm_state_dict, mm_name) + validate_ad_block_count(mm_state_dict, mm_name) for c in (320, 640, 1280, 1280): self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN)) for c in (1280, 1280, 640, 320): diff --git a/animatediff/motion_module_adxl.py b/animatediff/motion_module_adxl.py new file mode 100644 index 0000000..d55613b --- /dev/null +++ b/animatediff/motion_module_adxl.py @@ -0,0 +1,118 @@ +import math +from typing import Callable, Iterable, Optional, Union + +import torch +from einops import rearrange, repeat +from torch import Tensor, nn + +from comfy.ldm.modules.attention import FeedForward +from .motion_lora import MotionLoRAInfo +from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch +from .motion_module_ad import MotionModule + + +def zero_module(module): + # Zero out the parameters of a module and return it. + for p in module.parameters(): + p.detach().zero_() + return module + + +def get_ad_sdxl_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_type: str) -> 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 MotionCompatibilityError(f"No pos_encoder.pe found in mm_state_dict - {mm_type} is not a valid AnimateDiff-SDXL motion module!") + + +def validate_ad_sdxl_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None: + # keep track of biggest down_block count in module + biggest_block = 0 + for key in mm_state_dict.keys(): + if "down_blocks" 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 + if biggest_block != 2: + raise MotionCompatibilityError(f"Expected biggest down_block to be 2, but was {biggest_block} - {mm_type} is not a valid AnimateDiff-SDXL motion module!") + + +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 + + +class AnimDiffSDXLMotionWrapper(GenericMotionWrapper): + def __init__(self, mm_state_dict: dict[str, Tensor], mm_hash: str, mm_name: str="mm_sd_v15.ckpt" , loras: list[MotionLoRAInfo]=None): + super().__init__(mm_hash, mm_name, loras) + self.down_blocks: Iterable[MotionModule] = nn.ModuleList([]) + self.up_blocks: Iterable[MotionModule] = nn.ModuleList([]) + self.mid_block: Union[MotionModule, None] = None + self.encoding_max_len = get_ad_sdxl_temporal_position_encoding_max_len(mm_state_dict, mm_name) + validate_ad_sdxl_block_count(mm_state_dict, mm_name) + for c in (320, 640, 1280): + self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN)) + for c in (1280, 640, 320): + 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) + self.mm_hash = mm_hash + self.mm_name = mm_name + self.version = "v1" if self.mid_block is None else "v2" + self.injector_version = InjectorVersion.ADXL_V1_V2 + self.AD_video_length: int = 24 + self.loras = loras + + def has_loras(self): + # TODO: fix this to return False if has an empty list as well + # but only after implementing a fix for lowvram loading + return self.loras is not None + + 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 set_masks(self, masks: Tensor, min_val: float, max_val: float): + for block in self.down_blocks: + block.set_masks(masks, min_val, max_val) + for block in self.up_blocks: + block.set_masks(masks, min_val, max_val) + if self.mid_block is not None: + self.mid_block.set_masks(masks, min_val, max_val) + + def set_sub_idxs(self, sub_idxs: list[int]): + for block in self.down_blocks: + block.set_sub_idxs(sub_idxs) + for block in self.up_blocks: + block.set_sub_idxs(sub_idxs) + if self.mid_block is not None: + self.mid_block.set_sub_idxs(sub_idxs) + + 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() diff --git a/animatediff/motion_module_hsxl.py b/animatediff/motion_module_hsxl.py index 2532fcf..f7cbb03 100644 --- a/animatediff/motion_module_hsxl.py +++ b/animatediff/motion_module_hsxl.py @@ -8,7 +8,7 @@ from torch import Tensor, nn from comfy.ldm.modules.attention import FeedForward from .motion_lora import MotionLoRAInfo -from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric +from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric def zero_module(module): @@ -23,7 +23,23 @@ def get_hsxl_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor] for key in mm_state_dict.keys(): if key.endswith("pos_encoder.positional_encoding"): return mm_state_dict[key].size(1) # get middle dim - raise ValueError(f"No pos_encoder.positional_encoding found in mm_state_dict - {mm_type} is not a valid HotShotXL motion module!") + raise MotionCompatibilityError(f"No pos_encoder.positional_encoding found in mm_state_dict - {mm_type} is not a valid HotShotXL motion module!") + + +def validate_hsxl_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None: + # keep track of biggest down_block count in module + biggest_block = 0 + for key in mm_state_dict.keys(): + if "down_blocks" 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 + if biggest_block != 2: + raise MotionCompatibilityError(f"Expected biggest down_block to be 2, but was {biggest_block} - {mm_type} is not a valid HotShotXL motion module!") def has_mid_block(mm_state_dict: dict[str, Tensor]): @@ -48,6 +64,7 @@ class HotShotXLMotionWrapper(GenericMotionWrapper): self.up_blocks: Iterable[HotShotXLMotionModule] = nn.ModuleList([]) self.mid_block: Union[HotShotXLMotionModule, None] = None self.encoding_max_len = get_hsxl_temporal_position_encoding_max_len(mm_state_dict, mm_name) + validate_hsxl_block_count(mm_state_dict, mm_name) for c in (320, 640, 1280): self.down_blocks.append(HotShotXLMotionModule(c, block_type=BlockType.DOWN, max_length=self.encoding_max_len)) for c in (1280, 640, 320): diff --git a/animatediff/motion_utils.py b/animatediff/motion_utils.py index 0b81399..ec4d566 100644 --- a/animatediff/motion_utils.py +++ b/animatediff/motion_utils.py @@ -139,6 +139,7 @@ class BlockType: class InjectorVersion: V1_V2 = "v1/v2" + ADXL_V1_V2 = "ADXL v1/v2" HOTSHOTXL_V1 = "HSXL v1" @@ -271,3 +272,7 @@ class NoiseType: generator = torch.manual_seed(seed+i) all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu")) return torch.cat(all_noises, dim=0) + + +class MotionCompatibilityError(ValueError): + pass diff --git a/animatediff/sampling.py b/animatediff/sampling.py index b588a97..87eea42 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -21,6 +21,7 @@ from .motion_module import InjectionParams, eject_motion_module, inject_motion_m load_motion_module, unload_motion_module from .motion_module import is_injected_mm_params, get_injected_mm_params from .motion_module_ad import AnimDiffMotionWrapper, VanillaTemporalModule +from .motion_module_adxl import AnimDiffSDXLMotionWrapper from .motion_utils import GenericMotionWrapper, GroupNormAD, NoiseType @@ -155,8 +156,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: openaimodel.forward_timestep_embed = forward_timestep_embed if params.unlimited_area_hack: model_management.maximum_batch_area = unlimited_batch_area - # only apply groupnorm hack if not v2 and should not apply v2 properly - if not (isinstance(motion_module, AnimDiffMotionWrapper) and motion_module.version == "v2" and params.apply_v2_models_properly): + # only apply groupnorm hack if not [AnimateDiff and v2 and should apply v2 properly] + if not ((isinstance(motion_module, AnimDiffMotionWrapper) and motion_module.version == "v2" and params.apply_v2_models_properly)): torch.nn.GroupNorm.forward = groupnorm_mm_factory(params) if params.apply_mm_groupnorm_hack: GroupNormAD.forward = groupnorm_mm_factory(params)