From 63e5ee26aedd76e56a777fc5c5202754ec42404f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 6 Dec 2023 07:49:19 -0600 Subject: [PATCH] Readded motion LoRA support --- animatediff/model_injection.py | 12 +++++++----- animatediff/model_utils.py | 12 ++++++------ animatediff/motion_lora.py | 13 ------------- animatediff/motion_module_ad.py | 8 ++++---- animatediff/nodes.py | 14 +++++++++----- animatediff/sampling_motion.py | 4 ++-- 6 files changed, 28 insertions(+), 35 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 8acce4e..ec04cc6 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -14,9 +14,9 @@ from .motion_module_ad import AnimateDiffModelWrapper, has_mid_block, normalize_ from .logger import logger from .motion_utils import MotionCompatibilityError, NoiseType, normalize_min_max -from .motion_lora import MotionLoraInfo, MotionLoraList, MotionLoraWrapper +from .motion_lora import MotionLoraInfo, MotionLoraList -from .model_utils import ModelTypesSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \ +from .model_utils import ModelTypeSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \ get_sd_model_type @@ -55,15 +55,15 @@ class ModelPatcherAndInjector(ModelPatcher): return super().unpatch_model(device_to) def inject_model(self, device_to=None): - # motion_model.model is AnimateDiffModelWrapper if self.motion_model is not None: self.motion_model.model.eject(self) self.motion_model.model.inject(self) + self.motion_model.model.to(device_to) def eject_model(self, device_to=None): - # motion_model.model is AnimateDiffModelWrapper if self.motion_model is not None: self.motion_model.model.eject(self) + self.motion_model.model.to(device_to) def clone(self): cloned = ModelPatcherAndInjector(self) @@ -73,6 +73,7 @@ class ModelPatcherAndInjector(ModelPatcher): class MotionModelPatcher(ModelPatcher): + # Mostly here so that type hints work in IDEs def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.model: AnimateDiffModelWrapper = self.model @@ -118,7 +119,8 @@ def load_motion_lora_as_patches(motion_model: MotionModelPatcher, lora: MotionLo weight_down = state_dict[key] weight_up = state_dict[up_key] # actual weights obtained by matrix multiplication of up and down weights - patches[model_key] += torch.mm(weight_up, weight_down) + # save as a tuple, so that (Motion)ModelPatcher's calculate_weight function detects len==1, applying it correctly + patches[model_key] = (torch.mm(weight_up, weight_down),) del state_dict # add patches to motion ModelPatcher motion_model.add_patches(patches=patches, strength_patch=lora.strength) diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index 8d15267..5c55de8 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -154,7 +154,7 @@ def calculate_model_hash(model: ModelPatcher): return m.hexdigest() -class ModelTypesSD: +class ModelTypeSD: SD1_5 = "SD1.5" SD2_1 = "SD2.1" SDXL = "SDXL" @@ -166,15 +166,15 @@ def get_sd_model_type(model: ModelPatcher) -> str: if model is None: return None elif type(model.model) == BaseModel: - return ModelTypesSD.SD1_5 + return ModelTypeSD.SD1_5 elif type(model.model) == SDXL: - return ModelTypesSD.SDXL + return ModelTypeSD.SDXL elif type(model.model) == SD21UNCLIP: - return ModelTypesSD.SD2_1 + return ModelTypeSD.SD2_1 elif type(model.model) == SDXLRefiner: - return ModelTypesSD.SDXL_REFINER + return ModelTypeSD.SDXL_REFINER elif type(model.model) == SVD_img2vid: - return ModelTypesSD.SVD + return ModelTypeSD.SVD else: return str(type(model.model).__name__) diff --git a/animatediff/motion_lora.py b/animatediff/motion_lora.py index 7a42784..d96259d 100644 --- a/animatediff/motion_lora.py +++ b/animatediff/motion_lora.py @@ -1,6 +1,3 @@ -from torch import Tensor - - class MotionLoraInfo: def __init__(self, name: str, strength: float = 1.0, hash: str=""): self.name = name @@ -14,16 +11,6 @@ class MotionLoraInfo: return MotionLoraInfo(self.name, self.strength, self.hash) -class MotionLoraWrapper: - def __init__(self, state_dict: dict[str, Tensor], hash: str=""): - self.state_dict = state_dict - self.hash = hash - self.info: MotionLoraInfo = None - - def set_info(self, info: MotionLoraInfo): - self.info = info - - class MotionLoraList: def __init__(self): self.loras: list[MotionLoraInfo] = [] diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 1903d09..bd8d168 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -12,7 +12,7 @@ from ldm.modules.diffusionmodules import openaimodel from comfy.ldm.modules.diffusionmodules.openaimodel import ResBlock, SpatialTransformer from .motion_lora import MotionLoraInfo from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch -from .model_utils import ModelTypesSD +from .model_utils import ModelTypeSD def zero_module(module): @@ -96,9 +96,9 @@ def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> T sd_type: str = None down_block_max = get_down_block_max(mm_state_dict) if down_block_max == 3: - sd_type = ModelTypesSD.SD1_5 + sd_type = ModelTypeSD.SD1_5 elif down_block_max == 2: - sd_type = ModelTypesSD.SDXL + sd_type = ModelTypeSD.SDXL else: raise ValueError(f"'{mm_name}' is not a valid SD1.5 nor SDXL motion module - contained {down_block_max} downblocks.") # determine the model's format @@ -141,7 +141,7 @@ class AnimateDiffModelWrapper(nn.Module): self.mid_block: Union[MotionModule, None] = None self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, mm_info) # SDXL has 3 up/down blocks, SD1.5 has 4 up/down blocks - if mm_info.sd_type == ModelTypesSD.SDXL: + if mm_info.sd_type == ModelTypeSD.SDXL: layer_channels = (320, 640, 1280) else: layer_channels = (320, 640, 1280, 1280) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 30459e0..41c79a2 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -1,3 +1,4 @@ +from pathlib import Path import torch import comfy.sample as comfy_sample @@ -5,7 +6,7 @@ from comfy.model_patcher import ModelPatcher from .context import ContextOptions, ContextSchedules, UniformContextOptions from .logger import logger -from .model_utils import get_available_motion_loras, get_available_motion_models, BetaSchedules +from .model_utils import get_available_motion_loras, get_available_motion_models, BetaSchedules, get_motion_lora_path from .motion_utils import NoiseType from .motion_lora import MotionLoraInfo, MotionLoraList from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelSettings, load_motion_module @@ -69,10 +70,13 @@ class AnimateDiffLoraLoader: prev_motion_lora = MotionLoraList() else: prev_motion_lora = prev_motion_lora.clone() - # load lora - # lora = load_motion_lora(lora_name) - # lora_info = MotionLoraInfo(name=lora_name, strength=strength, hash=lora.hash) - # prev_motion_lora.add_lora(lora_info) + # check if motion lora with name exists + lora_path = get_motion_lora_path(lora_name) + if not Path(lora_path).is_file(): + raise FileNotFoundError(f"Motion lora with name '{lora_name}' not found.") + # create motion lora info to be loaded in AnimateDiff Loader + lora_info = MotionLoraInfo(name=lora_name, strength=strength) + prev_motion_lora.add_lora(lora_info) return (prev_motion_lora,) diff --git a/animatediff/sampling_motion.py b/animatediff/sampling_motion.py index b1798d4..28927a3 100644 --- a/animatediff/sampling_motion.py +++ b/animatediff/sampling_motion.py @@ -16,7 +16,7 @@ from comfy.controlnet import ControlBase from .context import get_context_scheduler from .motion_utils import GroupNormAD, NoiseType -from .model_utils import BetaScheduleCache, BetaSchedules, ModelTypesSD, wrap_function_to_inject_xformers_bug_info +from .model_utils import BetaScheduleCache, BetaSchedules, ModelTypeSD, wrap_function_to_inject_xformers_bug_info from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelPatcher from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffModelWrapper, VanillaTemporalModule from .logger import logger @@ -219,7 +219,7 @@ def motion_sample_factory(orig_comfy_sample: Callable) -> Callable: model.model.memory_required = unlimited_memory_required # only apply groupnorm hack if not [AnimateDiff SD1.5 and v2 and should apply v2 properly] info: AnimateDiffInfo = model.motion_model.model.mm_info - if not (info.mm_format == AnimateDiffFormat.ANIMATEDIFF and info.sd_type == ModelTypesSD.SD1_5 and \ + if not (info.mm_format == AnimateDiffFormat.ANIMATEDIFF and info.sd_type == ModelTypeSD.SD1_5 and \ info.mm_version == "v2" and params.apply_v2_models_properly): torch.nn.GroupNorm.forward = groupnorm_mm_factory(params) if params.apply_mm_groupnorm_hack: