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:
@@ -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
@@ -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
|
||||
############################################################################
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user