From b459e3e0ef10873c1f5ce99427f9febf9adc2eef Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 17 Jan 2024 09:36:10 -0600 Subject: [PATCH] Fixed same Load AnimateDiff Model output linked to two connected Apply AD Model nodes not working as intended --- animatediff/model_injection.py | 9 +++++++++ animatediff/nodes_gen2.py | 7 ++++++- 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index bf63e2b..0d840f3 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -396,6 +396,15 @@ def load_motion_module_gen2(model_name: str, motion_model_settings: 'MotionModel return motion_model +def create_fresh_motion_module(motion_model: MotionModelPatcher) -> MotionModelPatcher: + ad_wrapper = AnimateDiffModel(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()) + return MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(), + offload_device=comfy.model_management.unet_offload_device()) + + 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/nodes_gen2.py b/animatediff/nodes_gen2.py index 8e79d8c..2eedfa7 100644 --- a/animatediff/nodes_gen2.py +++ b/animatediff/nodes_gen2.py @@ -9,7 +9,7 @@ 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, MotionModelSettings, +from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, MotionModelSettings, create_fresh_motion_module, load_motion_module, load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2) from .sample_settings import SampleSettings, SeedNoiseGeneration from .sampling import motion_sample_factory @@ -103,6 +103,11 @@ class ApplyAnimateDiffModelNode: prev_m_models = MotionModelGroup() prev_m_models = prev_m_models.clone() motion_model = motion_model.clone() + # check if internal motion model already present in previous model - create new if so + for prev_model in prev_m_models.models: + if motion_model.model is prev_model.model: + # need to create new internal model based on same state_dict + motion_model = create_fresh_motion_module(motion_model) # apply motion model to loaded_mm if motion_lora is not None: for lora in motion_lora.loras: