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

This commit is contained in:
Jedrzej Kosinski
2024-03-15 19:20:59 -05:00
parent 70ebc0a396
commit 295e70e076
5 changed files with 329 additions and 86 deletions
+22 -2
View File
@@ -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)
+159 -42
View File
@@ -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
############################################################################
+6 -2
View File
@@ -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
+106 -40
View File
@@ -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):
+36
View File
@@ -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