diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index e358d6a..2c080fa 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -89,13 +89,26 @@ def load_motion_lora(lora_name: str) -> MotionLoRAWrapper: return lora +def interpolate_pe_to_length(model_dict: dict[str, Tensor], key: str, new_length: int): + pe_shape = model_dict[key].shape + temp_pe = rearrange(model_dict[key], "(t b) f d -> t b f d", t=1) + temp_pe = F.interpolate(temp_pe, size=(new_length, pe_shape[-1]), mode="bilinear") + temp_pe = rearrange(temp_pe, "t b f d -> (t b) f d", t=1) + model_dict[key] = temp_pe + del temp_pe + + def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSettings') -> dict[str, Tensor]: if not mm_settings.has_anything_to_apply(): return model_dict for key in model_dict: if "attention_blocks" in key: - # apply pe_strength, if needed if "pos_encoder" in key: + # apply simple motion pe stretch, if needed + if mm_settings.has_motion_pe_stretch(): + new_pe_length = model_dict[key].shape[1] + mm_settings.motion_pe_stretch + interpolate_pe_to_length(model_dict, key, new_length=new_pe_length) + # apply pe_strength, if needed if mm_settings.has_pe_strength(): model_dict[key] *= mm_settings.pe_strength # apply pe_idx_offset, if needed @@ -106,12 +119,7 @@ def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSe model_dict[key] = model_dict[key][:, :mm_settings.cap_initial_pe_length] # apply interpolate_pe_to_length, if needed if mm_settings.has_interpolate_pe_to_length(): - pe_shape = model_dict[key].shape - temp_pe = rearrange(model_dict[key], "(t b) f d -> t b f d", t=1) - temp_pe = F.interpolate(temp_pe, size=(mm_settings.interpolate_pe_to_length, pe_shape[-1]), mode="bilinear") - temp_pe = rearrange(temp_pe, "t b f d -> (t b) f d", t=1) - model_dict[key] = temp_pe - del temp_pe + interpolate_pe_to_length(model_dict, key, new_length=mm_settings.interpolate_pe_to_length) # apply final_pe_idx_offset, if needed if mm_settings.has_final_pe_idx_offset(): model_dict[key] = model_dict[key][:, mm_settings.final_pe_idx_offset:] @@ -501,6 +509,7 @@ class MotionModelSettings: pe_strength: float=1.0, attn_strength: float=1.0, other_strength: float=1.0, cap_initial_pe_length: int=0, interpolate_pe_to_length: int=0, initial_pe_idx_offset: int=0, final_pe_idx_offset: int=0, + motion_pe_stretch: int=0, attn_scale: float=1.0, ): # PE-interpolation settings @@ -511,6 +520,7 @@ class MotionModelSettings: self.interpolate_pe_to_length = interpolate_pe_to_length self.initial_pe_idx_offset = initial_pe_idx_offset self.final_pe_idx_offset = final_pe_idx_offset + self.motion_pe_stretch = motion_pe_stretch # attention scale settings self.attn_scale = attn_scale @@ -535,6 +545,9 @@ class MotionModelSettings: def has_final_pe_idx_offset(self) -> bool: return self.final_pe_idx_offset > 0 + def has_motion_pe_stretch(self) -> bool: + return self.motion_pe_stretch > 0 + def has_anything_to_apply(self) -> bool: return self.has_pe_strength() \ or self.has_attn_strength() \ @@ -542,4 +555,5 @@ class MotionModelSettings: or self.has_cap_initial_pe_length() \ or self.has_interpolate_pe_to_length() \ or self.has_initial_pe_idx_offset() \ - or self.has_final_pe_idx_offset() + or self.has_final_pe_idx_offset() \ + or self.has_motion_pe_stretch() diff --git a/animatediff/nodes.py b/animatediff/nodes.py index c2888dc..05d55a5 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -57,6 +57,27 @@ class AnimateDiffLoRALoader: return (prev_motion_lora,) +class AnimateDiffModelSettingsSimple: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "motion_pe_stretch": ("INT", {"default": 0, "min": 0, "step": 1}), + }, + } + + RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings" + FUNCTION = "get_motion_model_settings" + + def get_motion_model_settings(self, motion_pe_stretch: int): + motion_model_settings = MotionModelSettings( + motion_pe_stretch=motion_pe_stretch + ) + + return (motion_model_settings,) + + class AnimateDiffModelSettingsAdvanced: @classmethod def INPUT_TYPES(s): @@ -65,6 +86,7 @@ class AnimateDiffModelSettingsAdvanced: "pe_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), "attn_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), "other_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), + "motion_pe_stretch": ("INT", {"default": 0, "min": 0, "step": 1}), "cap_initial_pe_length": ("INT", {"default": 0, "min": 0, "step": 1}), "interpolate_pe_to_length": ("INT", {"default": 0, "min": 0, "step": 1}), "initial_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}), @@ -73,10 +95,11 @@ class AnimateDiffModelSettingsAdvanced: } RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) - CATEGORY = "Animate Diff 🎭🅐🅓" + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings" FUNCTION = "get_motion_model_settings" def get_motion_model_settings(self, pe_strength: float, attn_strength: float, other_strength: float, + motion_pe_stretch: int, cap_initial_pe_length: int, interpolate_pe_to_length: int, initial_pe_idx_offset: int, final_pe_idx_offset: int): motion_model_settings = MotionModelSettings( @@ -87,6 +110,7 @@ class AnimateDiffModelSettingsAdvanced: interpolate_pe_to_length=interpolate_pe_to_length, initial_pe_idx_offset=initial_pe_idx_offset, final_pe_idx_offset=final_pe_idx_offset, + motion_pe_stretch=motion_pe_stretch ) return (motion_model_settings,) @@ -485,6 +509,7 @@ NODE_CLASS_MAPPINGS = { "ADE_AnimateDiffUniformContextOptions": AnimateDiffUniformContextOptions, "ADE_AnimateDiffLoaderWithContext": AnimateDiffLoaderWithContext, "ADE_AnimateDiffLoRALoader": AnimateDiffLoRALoader, + "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, @@ -497,6 +522,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_AnimateDiffUniformContextOptions": "Uniform Context Options 🎭🅐🅓", "ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader 🎭🅐🅓", "ADE_AnimateDiffLoRALoader": "AnimateDiff LoRA Loader 🎭🅐🅓", + "ADE_AnimateDiffModelSettingsSimple": "Motion Model Settings (Simple) 🎭🅐🅓", "ADE_AnimateDiffModelSettings": "Motion Model Settings (Advanced) 🎭🅐🅓", "ADE_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓",