From 295e70e076f8682e4dea453efacbdd64664ad1bf Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 15 Mar 2024 19:20:59 -0500 Subject: [PATCH] Fixed bug with ref in AnimateLCM-I2V continued to be used, exposed a proper way of replicating the bug or using just the img_encoder without the motion model, AnimateDiffModel changes to facilitate these features --- animatediff/model_injection.py | 24 +++- animatediff/motion_module_ad.py | 201 +++++++++++++++++++++++++------- animatediff/nodes.py | 8 +- animatediff/nodes_gen2.py | 146 ++++++++++++++++------- animatediff/utils_motion.py | 36 ++++++ 5 files changed, 329 insertions(+), 86 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index c40cb69..baa35f6 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -13,7 +13,7 @@ from comfy.model_base import BaseModel from .ad_settings import AnimateDiffSettings from .context import ContextOptions, ContextOptions, ContextOptionsGroup -from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, has_mid_block, normalize_ad_state_dict +from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, has_mid_block, normalize_ad_state_dict from .logger import logger from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max from .motion_lora import MotionLoraInfo, MotionLoraList @@ -103,6 +103,7 @@ class MotionModelPatcher(ModelPatcher): # AnimateLCM-I2V self.orig_ref_drift: float = None self.orig_insertion_weights: list[float] = None + self.orig_apply_ref_when_disabled = False self.orig_img_latents: Tensor = None self.img_features: list[int, Tensor] = None # temporary self.img_latents_shape: tuple = None @@ -210,7 +211,7 @@ class MotionModelPatcher(ModelPatcher): if goal_length != img_latents.shape[0]: img_latents = ade_broadcast_image_to(img_latents, goal_length, batched_number) img_features = self.model.img_encoder(img_latents, goal_length, batched_number) - self.model.set_img_features(img_features=img_features) + self.model.set_img_features(img_features=img_features, apply_ref_when_disabled=self.orig_apply_ref_when_disabled) # cache values for next step self.img_latents_shape = img_latents.shape self.prev_sub_idxs = sub_idxs @@ -252,6 +253,9 @@ class MotionModelPatcher(ModelPatcher): n.scale_multival = self.scale_multival n.effect_multival = self.effect_multival n.orig_img_latents = self.orig_img_latents + n.orig_ref_drift = self.orig_ref_drift + n.orig_insertion_weights = self.orig_insertion_weights.copy() if self.orig_insertion_weights is not None else self.orig_insertion_weights + n.orig_apply_ref_when_disabled = self.orig_apply_ref_when_disabled return n @@ -452,6 +456,22 @@ def create_fresh_motion_module(motion_model: MotionModelPatcher) -> MotionModelP offload_device=comfy.model_management.unet_offload_device()) +def create_fresh_encoder_only_model(motion_model: MotionModelPatcher) -> MotionModelPatcher: + ad_wrapper = EncoderOnlyAnimateDiffModel(mm_state_dict=motion_model.model.state_dict(), mm_info=motion_model.model.mm_info) + ad_wrapper.to(comfy.model_management.unet_dtype()) + ad_wrapper.to(comfy.model_management.unet_offload_device()) + ad_wrapper.load_state_dict(motion_model.model.state_dict(), strict=False) + return MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(), + offload_device=comfy.model_management.unet_offload_device()) + + +def inject_img_encoder_into_model(motion_model: MotionModelPatcher, w_encoder: MotionModelPatcher): + motion_model.model.init_img_encoder() + motion_model.model.img_encoder.to(comfy.model_management.unet_dtype()) + motion_model.model.img_encoder.to(comfy.model_management.unet_offload_device()) + motion_model.model.img_encoder.load_state_dict(w_encoder.model.img_encoder.state_dict()) + + def validate_model_compatibility_gen2(model: ModelPatcher, motion_model: MotionModelPatcher): # check that motion model is compatible with sd model model_sd_type = get_sd_model_type(model) diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 62ba097..4ef7a0a 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -17,7 +17,7 @@ import comfy.model_management from .context import ContextFuseMethod, ContextOptions, get_context_weights, get_context_windows from .animatelcm_i2v_adapter import AdapterEmbed -from .utils_motion import CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch +from .utils_motion import CrossAttentionMM, MotionCompatibilityError, DummyNNModule, extend_to_batch_size, prepare_mask_batch from .utils_model import BetaSchedules, ModelTypeSD from .logger import logger @@ -195,11 +195,13 @@ class AnimateDiffModel(nn.Module): ops = comfy.ops.disable_weight_init else: ops = comfy.ops.manual_cast + self.ops = ops # SDXL has 3 up/down blocks, SD1.5 has 4 up/down blocks if mm_info.sd_type == ModelTypeSD.SDXL: layer_channels = (320, 640, 1280) else: layer_channels = (320, 640, 1280, 1280) + self.layer_channels = layer_channels # fill out down/up blocks and middle block, if present for idx, c in enumerate(layer_channels): self.down_blocks.append(MotionModule(c, temporal_pe=self.has_position_encoding, @@ -214,7 +216,11 @@ class AnimateDiffModel(nn.Module): # create AdapterEmbed if keys present for it self.img_encoder: AdapterEmbed = None if has_img_encoder(mm_state_dict): - self.img_encoder = AdapterEmbed(cin=4, channels=layer_channels, nums_rb=2, ksize=1, sk=True, use_conv=False, ops=ops) + self.init_img_encoder() + + def init_img_encoder(self): + del self.img_encoder + self.img_encoder = AdapterEmbed(cin=4, channels=self.layer_channels, nums_rb=2, ksize=1, sk=True, use_conv=False, ops=self.ops) def get_device_debug(self): return self.down_blocks[0].motion_modules[0].temporal_transformer.proj_in.weight.device @@ -255,11 +261,13 @@ class AnimateDiffModel(nn.Module): # inject input (down) blocks # SD15 mm contains 4 downblocks, each with 2 TemporalTransformers - 8 in total # SDXL mm contains 3 downblocks, each with 2 TemporalTransformers - 6 in total - self._inject(unet.input_blocks, self.down_blocks) + if self.down_blocks is not None: + self._inject(unet.input_blocks, self.down_blocks) # inject output (up) blocks # SD15 mm contains 4 upblocks, each with 3 TemporalTransformers - 12 in total # SDXL mm contains 3 upblocks, each with 3 TemporalTransformers - 9 in total - self._inject(unet.output_blocks, self.up_blocks) + if self.up_blocks is not None: + self._inject(unet.output_blocks, self.up_blocks) # inject mid block, if needed (encapsulate in list to make structure compatible) if self.mid_block is not None: self._inject([unet.middle_block], [self.mid_block]) @@ -325,10 +333,12 @@ class AnimateDiffModel(nn.Module): 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.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) @@ -343,55 +353,68 @@ class AnimateDiffModel(nn.Module): self._set_scale_mask(None) def set_effect(self, multival: Union[float, Tensor]): - for block in self.down_blocks: - block.set_effect(multival) - for block in self.up_blocks: - block.set_effect(multival) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_effect(multival) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_effect(multival) if self.mid_block is not None: self.mid_block.set_effect(multival) 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.down_blocks is not None: + for block in self.down_blocks: + block.set_sub_idxs(sub_idxs) + if self.up_blocks is not None: + 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 set_view_options(self, view_options: ContextOptions): - for block in self.down_blocks: - block.set_view_options(view_options) - for block in self.up_blocks: - block.set_view_options(view_options) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_view_options(view_options) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_view_options(view_options) if self.mid_block is not None: self.mid_block.set_view_options(view_options) - def set_img_features(self, img_features: list[Tensor]): + def set_img_features(self, img_features: list[Tensor], apply_ref_when_disabled=False): # img_features should only impact downblocks - for block in self.down_blocks: - block.set_img_features(img_features=img_features) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_img_features(img_features=img_features, apply_ref_when_disabled=apply_ref_when_disabled) 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.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_scale_mask(self, mask: Tensor): - for block in self.down_blocks: - block.set_scale_mask(mask) - for block in self.up_blocks: - block.set_scale_mask(mask) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_scale_mask(mask) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_scale_mask(mask) if self.mid_block is not None: self.mid_block.set_scale_mask(mask) 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.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() @@ -451,9 +474,9 @@ class MotionModule(nn.Module): for motion_module in self.motion_modules: motion_module.set_view_options(view_options=view_options) - def set_img_features(self, img_features: list[Tensor]): + def set_img_features(self, img_features: list[Tensor], apply_ref_when_disabled=False): for motion_module in self.motion_modules: - motion_module.set_img_features(img_features=img_features) + motion_module.set_img_features(img_features=img_features, apply_ref_when_disabled=apply_ref_when_disabled) def reset_temp_vars(self): for motion_module in self.motion_modules: @@ -499,6 +522,7 @@ class VanillaTemporalModule(nn.Module): self.prev_input_tensor_batch = 0 # AnimateLCM-I2V vars self.img_features: list[Tensor] = None + self.apply_ref_when_disabled = False self.temporal_transformer = TemporalTransformer3DModel( in_channels=in_channels, @@ -546,9 +570,10 @@ class VanillaTemporalModule(nn.Module): def set_view_options(self, view_options: ContextOptions): self.view_options = view_options - def set_img_features(self, img_features: list[Tensor]): + def set_img_features(self, img_features: list[Tensor], apply_ref_when_disabled=False): del self.img_features self.img_features = img_features + self.apply_ref_when_disabled = apply_ref_when_disabled def reset_temp_vars(self): self.set_effect(None) @@ -580,21 +605,29 @@ class VanillaTemporalModule(nn.Module): return self.temp_effect_mask[self.sub_idxs*batched_number] return self.temp_effect_mask[full_batched_idxs] + def should_handle_img_features(self): + return self.img_features is not None and self.block_type == BlockType.DOWN and self.module_idx == 1 + def forward(self, input_tensor: Tensor, encoder_hidden_states=None, attention_mask=None): - # do AnimateLCM-I2V stuff if needed - if self.img_features is not None: - if self.block_type == BlockType.DOWN and self.module_idx == 1: - input_tensor += self.img_features[self.block_idx] if self.effect is None: + # do AnimateLCM-I2V stuff if needed + if self.should_handle_img_features(): + input_tensor += self.img_features[self.block_idx] return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options) # return weighted average of input_tensor and AD output if type(self.effect) != Tensor: effect = self.effect # do nothing if effect is 0 if math.isclose(effect, 0.0): + # do AnimateLCM-I2V stuff if needed + if self.apply_ref_when_disabled and self.should_handle_img_features(): + input_tensor += self.img_features[self.block_idx] return input_tensor else: effect = self.get_effect_mask(input_tensor) + # do AnimateLCM-I2V stuff if needed + if self.should_handle_img_features(): + return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor+self.img_features[self.block_idx], encoder_hidden_states, attention_mask, self.view_options)*effect return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options)*effect @@ -1014,3 +1047,87 @@ class VersatileAttention(CrossAttentionMM): hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) return hidden_states + +############################################################################ +### EncoderOnly Version +############################################################################ +class EncoderOnlyAnimateDiffModel(AnimateDiffModel): + def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: AnimateDiffInfo): + super().__init__(mm_state_dict=mm_state_dict, mm_info=mm_info) + self.down_blocks: Iterable[EncoderOnlyMotionModule] = nn.ModuleList([]) + self.up_blocks = None + self.mid_block = None + # fill out down/up blocks and middle block, if present + for idx, c in enumerate(self.layer_channels): + self.down_blocks.append(EncoderOnlyMotionModule(c, block_type=BlockType.DOWN, block_idx=idx, ops=self.ops)) + + +class EncoderOnlyMotionModule(MotionModule): + ''' + MotionModule that will store EncoderOnlyTemporalModule objects instead of VanillaTemporalModules + ''' + def __init__( + self, + in_channels, + block_type: str=BlockType.DOWN, + block_idx: int=0, + ops=comfy.ops.disable_weight_init + ): + super().__init__(in_channels=in_channels, block_type=block_type, block_idx=block_idx, ops=ops) + if block_type == BlockType.MID: + # mid blocks contain only a single VanillaTemporalModule + self.motion_modules: Iterable[EncoderOnlyTemporalModule] = nn.ModuleList([EncoderOnlyTemporalModule.create(in_channels, block_type, block_idx, module_idx=0, ops=ops)]) + else: + # down blocks contain two VanillaTemporalModules + self.motion_modules: Iterable[EncoderOnlyTemporalModule] = nn.ModuleList( + [ + EncoderOnlyTemporalModule.create(in_channels, block_type, block_idx, module_idx=0, ops=ops), + EncoderOnlyTemporalModule.create(in_channels, block_type, block_idx, module_idx=1, ops=ops) + ] + ) + # up blocks contain one additional VanillaTemporalModule + if block_type == BlockType.UP: + self.motion_modules.append(EncoderOnlyTemporalModule.create(in_channels, block_type, block_idx, module_idx=2, ops=ops)) + + +class EncoderOnlyTemporalModule(VanillaTemporalModule): + ''' + VanillaTemporalModule that will only add img_features to input_tensor while respecting effect_multival + ''' + def __init__( + self, + in_channels, + block_type: str, + block_idx: int, + module_idx: int, + ops=comfy.ops.disable_weight_init, + ): + super().__init__(in_channels=in_channels, block_type=block_type, block_idx=block_idx, module_idx=module_idx, zero_initialize=False, ops=ops) + # make temporal_transformer a dummy class that does nothing, but will allow inherited VanillaTemporalModule code to work + self.temporal_transformer = DummyNNModule() + + @classmethod + def create(cls, in_channels, block_type: str, block_idx: int, module_idx: int, ops=comfy.ops.disable_weight_init): + return cls(in_channels=in_channels, block_type=block_type, block_idx=block_idx, module_idx=module_idx, ops=ops) + + def forward(self, input_tensor: Tensor, encoder_hidden_states=None, attention_mask=None): + if self.effect is None: + # do AnimateLCM-I2V stuff if needed + if self.should_handle_img_features(): + input_tensor += self.img_features[self.block_idx] + return input_tensor + # handle effect + if type(self.effect) != Tensor: + effect = self.effect + # do nothing if effect is 0 + if math.isclose(effect, 0.0): + # do AnimateLCM-I2V stuff if needed + if self.apply_ref_when_disabled and self.should_handle_img_features(): + input_tensor += self.img_features[self.block_idx] + return input_tensor + else: + effect = self.get_effect_mask(input_tensor) + if self.should_handle_img_features(): + return input_tensor*(1.0-effect) + (input_tensor+self.img_features[self.block_idx])*effect + return input_tensor # since no img_features to apply, no need for weighted average +############################################################################ diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 86029eb..2f1ae6d 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -4,8 +4,8 @@ from .sampling import motion_sample_factory from .nodes_gen1 import (AnimateDiffLoaderGen1, LegacyAnimateDiffLoaderWithContext, AnimateDiffModelSettings, AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths) -from .nodes_gen2 import (UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, LoadAnimateDiffModelNode, ADKeyframeNode, - ApplyAnimateLCMI2VModel) +from .nodes_gen2 import (UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, ApplyAnimateLCMI2VModel, ADKeyframeNode, + LoadAnimateDiffModelNode, LoadAnimateLCMI2VModelNode, LoadAnimateDiffAndInjectI2VNode) from .nodes_multival import MultivalDynamicNode, MultivalScaledMaskNode from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) @@ -81,6 +81,8 @@ NODE_CLASS_MAPPINGS = { "ADE_LoadAnimateDiffModel": LoadAnimateDiffModelNode, # AnimateLCM-I2V Nodes "ADE_ApplyAnimateLCMI2VModel": ApplyAnimateLCMI2VModel, + "ADE_LoadAnimateLCMI2VModel": LoadAnimateLCMI2VModelNode, + "ADE_InjectI2VIntoAnimateDiffModel": LoadAnimateDiffAndInjectI2VNode, # MaskedLoraLoader #"ADE_MaskedLoadLora": MaskedLoraLoader, # Deprecated Nodes @@ -145,6 +147,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_LoadAnimateDiffModel": "Load AnimateDiff Model πŸŽ­πŸ…πŸ…“β‘‘", # AnimateLCM-I2V Nodes "ADE_ApplyAnimateLCMI2VModel": "Apply AnimateLCM-I2V Model πŸŽ­πŸ…πŸ…“β‘‘", + "ADE_LoadAnimateLCMI2VModel": "Load AnimateLCM-I2V Model πŸŽ­πŸ…πŸ…“β‘‘", + "ADE_InjectI2VIntoAnimateDiffModel": "πŸ§ͺInject I2V into AnimateDiff Model πŸŽ­πŸ…πŸ…“β‘‘", # MaskedLoraLoader #"ADE_MaskedLoadLora": "Load LoRA (Masked) πŸŽ­πŸ…πŸ…“", # Deprecated Nodes diff --git a/animatediff/nodes_gen2.py b/animatediff/nodes_gen2.py index 2640bdd..4339e6e 100644 --- a/animatediff/nodes_gen2.py +++ b/animatediff/nodes_gen2.py @@ -6,13 +6,15 @@ import comfy.sample as comfy_sample from comfy.model_patcher import ModelPatcher from .ad_settings import AnimateDiffSettings +from .animatelcm_i2v_adapter import AdapterEmbed from .context import ContextOptions, ContextOptionsGroup, ContextSchedules from .logger import logger from .utils_model import BIGMAX, BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path from .utils_motion import ADKeyframeGroup, ADKeyframe from .motion_lora import MotionLoraInfo, MotionLoraList -from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, create_fresh_motion_module, - load_motion_module_gen1, load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2) +from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, create_fresh_motion_module, create_fresh_encoder_only_model, + load_motion_module_gen1, load_motion_module_gen2, load_motion_lora_as_patches, inject_img_encoder_into_model, validate_model_compatibility_gen2) +from .motion_module_ad import AnimateDiffFormat from .sample_settings import SampleSettings, SeedNoiseGeneration from .sampling import motion_sample_factory @@ -159,44 +161,6 @@ class ApplyAnimateDiffModelBasicNode: ad_keyframes=ad_keyframes) -class ApplyAnimateLCMI2VModel: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "motion_model": ("MOTION_MODEL_ADE",), - "ref_latent": ("LATENT",), - "ref_drift": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 10.0, "step": 0.001}), - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), - "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), - }, - "optional": { - "motion_lora": ("MOTION_LORA",), - "scale_multival": ("MULTIVAL",), - "effect_multival": ("MULTIVAL",), - "ad_keyframes": ("AD_KEYFRAMES",), - "prev_m_models": ("M_MODELS",), - } - } - - RETURN_TYPES = ("M_MODELS",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/β‘‘ Gen2 nodes β‘‘" - FUNCTION = "apply_motion_model" - - def apply_motion_model(self, motion_model: MotionModelPatcher, ref_latent: dict, ref_drift: float=0.5, start_percent: float=0.0, end_percent: float=1.0, - motion_lora: MotionLoraList=None, ad_keyframes: ADKeyframeGroup=None, - scale_multival=None, effect_multival=None, - prev_m_models: MotionModelGroup=None,): - new_m_models = ApplyAnimateDiffModelNode.apply_motion_model(self, motion_model, start_percent=start_percent, end_percent=end_percent, - motion_lora=motion_lora, ad_keyframes=ad_keyframes, - scale_multival=scale_multival, effect_multival=effect_multival, prev_m_models=prev_m_models) - # most recent added model will always be first in list; - curr_model = new_m_models[0].models[0] - curr_model.orig_img_latents = ref_latent["samples"] - curr_model.orig_ref_drift = ref_drift - return new_m_models - - class LoadAnimateDiffModelNode: @classmethod def INPUT_TYPES(s): @@ -220,6 +184,108 @@ class LoadAnimateDiffModelNode: return (motion_model,) +class ApplyAnimateLCMI2VModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "motion_model": ("MOTION_MODEL_ADE",), + "ref_latent": ("LATENT",), + "ref_drift": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}), + "apply_ref_when_disabled": ("BOOLEAN", {"default": False}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + }, + "optional": { + "motion_lora": ("MOTION_LORA",), + "scale_multival": ("MULTIVAL",), + "effect_multival": ("MULTIVAL",), + "ad_keyframes": ("AD_KEYFRAMES",), + "prev_m_models": ("M_MODELS",), + } + } + + RETURN_TYPES = ("M_MODELS",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/β‘‘ Gen2 nodes β‘‘/AnimateLCM-I2V" + FUNCTION = "apply_motion_model" + + def apply_motion_model(self, motion_model: MotionModelPatcher, ref_latent: dict, ref_drift: float=0.0, apply_ref_when_disabled=False, start_percent: float=0.0, end_percent: float=1.0, + motion_lora: MotionLoraList=None, ad_keyframes: ADKeyframeGroup=None, + scale_multival=None, effect_multival=None, + prev_m_models: MotionModelGroup=None,): + new_m_models = ApplyAnimateDiffModelNode.apply_motion_model(self, motion_model, start_percent=start_percent, end_percent=end_percent, + motion_lora=motion_lora, ad_keyframes=ad_keyframes, + scale_multival=scale_multival, effect_multival=effect_multival, prev_m_models=prev_m_models) + # most recent added model will always be first in list; + curr_model = new_m_models[0].models[0] + # confirm that model contains img_encoder + if curr_model.model.img_encoder is None: + raise Exception(f"Motion model '{curr_model.model.mm_info.mm_name}' does not contain an img_encoder; cannot be used with Apply AnimateLCM-I2V Model node.") + curr_model.orig_img_latents = ref_latent["samples"] + curr_model.orig_ref_drift = ref_drift + curr_model.orig_apply_ref_when_disabled = apply_ref_when_disabled + return new_m_models + + +class LoadAnimateLCMI2VModelNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (get_available_motion_models(),), + }, + "optional": { + "ad_settings": ("AD_SETTINGS",), + } + } + + RETURN_TYPES = ("MOTION_MODEL_ADE", "MOTION_MODEL_ADE") + RETURN_NAMES = ("MOTION_MODEL", "encoder_only") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/β‘‘ Gen2 nodes β‘‘/AnimateLCM-I2V" + FUNCTION = "load_motion_model" + + def load_motion_model(self, model_name: str, ad_settings: AnimateDiffSettings=None): + # load motion module and motion settings, if included + motion_model = load_motion_module_gen2(model_name=model_name, motion_model_settings=ad_settings) + # make sure model is an AnimateLCM-I2V model + if motion_model.model.mm_info.mm_format != AnimateDiffFormat.ANIMATELCM: + raise Exception(f"Motion model '{motion_model.model.mm_info.mm_name}' is not an AnimateLCM-I2V model; selected model is not AnimateLCM, and does not contain an img_encoder.") + if motion_model.model.img_encoder is None: + raise Exception(f"Motion model '{motion_model.model.mm_info.mm_name}' is not an AnimateLCM-I2V model; selected model IS AnimateLCM, but does NOT contain an img_encoder.") + # create encoder-only motion model + encoder_only_motion_model = create_fresh_encoder_only_model(motion_model=motion_model) + return (motion_model, encoder_only_motion_model) + + +class LoadAnimateDiffAndInjectI2VNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (get_available_motion_models(),), + "motion_model": ("MOTION_MODEL_ADE",), + }, + "optional": { + "ad_settings": ("AD_SETTINGS",), + } + } + + RETURN_TYPES = ("MOTION_MODEL_ADE",) + RETURN_NAMES = ("MOTION_MODEL",) + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/β‘‘ Gen2 nodes β‘‘/AnimateLCM-I2V/πŸ§ͺexperimental" + FUNCTION = "load_motion_model" + + def load_motion_model(self, model_name: str, motion_model: MotionModelPatcher, ad_settings: AnimateDiffSettings=None): + # make sure model w/ encoder actually has encoder + if motion_model.model.img_encoder is None: + raise Exception("Passed-in motion model was expected to have an img_encoder, but did not.") + # load motion module and motion settings, if included + loaded_motion_model = load_motion_module_gen2(model_name=model_name, motion_model_settings=ad_settings) + inject_img_encoder_into_model(motion_model=loaded_motion_model, w_encoder=motion_model) + return (loaded_motion_model,) + + class ADKeyframeNode: @classmethod def INPUT_TYPES(s): diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index d5ebb5d..0aaac63 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -248,3 +248,39 @@ class ADKeyframeGroup: if not tk.default: cloned.add(tk) return cloned + + +class DummyNNModule(nn.Module): + class DoNothingWhenCalled: + def __call__(self, *args, **kwargs): + return + + ''' + Class that does not throw exceptions for almost anything you throw at it. As name implies, does nothing. + ''' + def __init__(self): + super().__init__() + + def __getattr__(self, *args, **kwargs): + return self.DoNothingWhenCalled() + + def __setattr__(self, name, value): + pass + + def __iter__(self, *args, **kwargs): + pass + + def __next__(self, *args, **kwargs): + pass + + def __len__(self, *args, **kwargs): + pass + + def __getitem__(self, *args, **kwargs): + pass + + def __setitem__(self, *args, **kwargs): + pass + + def __call__(self, *args, **kwargs): + pass