diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 83dd265..b0b232f 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -1,10 +1,5 @@ -from pathlib import Path - import comfy.sample as comfy_sample -from .logger import logger -from .utils_model import get_available_motion_loras, get_motion_lora_path -from .motion_lora import MotionLoraInfo, MotionLoraList from .sampling import motion_sample_factory from .nodes_gen1 import (AnimateDiffLoaderGen1, LegacyAnimateDiffLoaderWithContext, AnimateDiffModelSettings, @@ -17,46 +12,15 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform from .nodes_ad_settings import AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect from .nodes_deprecated import AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated -from .nodes_lora import MaskedLoraLoader +from .nodes_lora import AnimateDiffLoraLoader, MaskedLoraLoader + +from .logger import logger # override comfy_sample.sample with animatediff-support version comfy_sample.sample = motion_sample_factory(comfy_sample.sample) comfy_sample.sample_custom = motion_sample_factory(comfy_sample.sample_custom, is_custom=True) -class AnimateDiffLoraLoader: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "lora_name": (get_available_motion_loras(),), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}), - }, - "optional": { - "prev_motion_lora": ("MOTION_LORA",), - } - } - - RETURN_TYPES = ("MOTION_LORA",) - CATEGORY = "Animate Diff 🎭🅐🅓" - FUNCTION = "load_motion_lora" - - def load_motion_lora(self, lora_name: str, strength: float, prev_motion_lora: MotionLoraList=None): - if prev_motion_lora is None: - prev_motion_lora = MotionLoraList() - else: - prev_motion_lora = prev_motion_lora.clone() - # check if motion lora with name exists - lora_path = get_motion_lora_path(lora_name) - if not Path(lora_path).is_file(): - raise FileNotFoundError(f"Motion lora with name '{lora_name}' not found.") - # create motion lora info to be loaded in AnimateDiff Loader - lora_info = MotionLoraInfo(name=lora_name, strength=strength) - prev_motion_lora.add_lora(lora_info) - - return (prev_motion_lora,) - - NODE_CLASS_MAPPINGS = { # Unencapsulated "ADE_AnimateDiffLoRALoader": AnimateDiffLoraLoader, @@ -105,7 +69,7 @@ NODE_CLASS_MAPPINGS = { "ADE_ApplyAnimateDiffModel": ApplyAnimateDiffModelNode, "ADE_LoadAnimateDiffModel": LoadAnimateDiffModelNode, # MaskedLoraLoader - "ADE_MaskedLoadLora": MaskedLoraLoader, + #"ADE_MaskedLoadLora": MaskedLoraLoader, # Deprecated Nodes "AnimateDiffLoaderV1": AnimateDiffLoader_Deprecated, "ADE_AnimateDiffLoaderV1Advanced": AnimateDiffLoaderAdvanced_Deprecated, @@ -158,6 +122,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_ApplyAnimateDiffModelSimple": "Apply AnimateDiff Model 🎭🅐🅓②", "ADE_ApplyAnimateDiffModel": "Apply AnimateDiff Model (Adv.) 🎭🅐🅓②", "ADE_LoadAnimateDiffModel": "Load AnimateDiff Model 🎭🅐🅓②", + # MaskedLoraLoader + #"ADE_MaskedLoadLora": "Load LoRA (Masked) 🎭🅐🅓", # Deprecated Nodes "AnimateDiffLoaderV1": "AnimateDiff Loader [DEPRECATED] 🎭🅐🅓", "ADE_AnimateDiffLoaderV1Advanced": "AnimateDiff Loader (Advanced) [DEPRECATED] 🎭🅐🅓", diff --git a/animatediff/nodes_lora.py b/animatediff/nodes_lora.py index 5300551..5cc3ba3 100644 --- a/animatediff/nodes_lora.py +++ b/animatediff/nodes_lora.py @@ -1,7 +1,46 @@ +from pathlib import Path + import folder_paths import comfy.utils import comfy.sd +from .logger import logger +from .utils_model import get_available_motion_loras, get_motion_lora_path +from .motion_lora import MotionLoraInfo, MotionLoraList + + +class AnimateDiffLoraLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora_name": (get_available_motion_loras(),), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}), + }, + "optional": { + "prev_motion_lora": ("MOTION_LORA",), + } + } + + RETURN_TYPES = ("MOTION_LORA",) + CATEGORY = "Animate Diff 🎭🅐🅓" + FUNCTION = "load_motion_lora" + + def load_motion_lora(self, lora_name: str, strength: float, prev_motion_lora: MotionLoraList=None): + if prev_motion_lora is None: + prev_motion_lora = MotionLoraList() + else: + prev_motion_lora = prev_motion_lora.clone() + # check if motion lora with name exists + lora_path = get_motion_lora_path(lora_name) + if not Path(lora_path).is_file(): + raise FileNotFoundError(f"Motion lora with name '{lora_name}' not found.") + # create motion lora info to be loaded in AnimateDiff Loader + lora_info = MotionLoraInfo(name=lora_name, strength=strength) + prev_motion_lora.add_lora(lora_info) + + return (prev_motion_lora,) + class MaskedLoraLoader: def __init__(self): diff --git a/animatediff/nodes_multival.py b/animatediff/nodes_multival.py index 70d0e7e..61d243f 100644 --- a/animatediff/nodes_multival.py +++ b/animatediff/nodes_multival.py @@ -4,7 +4,7 @@ from typing import Union import torch from torch import Tensor -from .utils_motion import linear_conversion, normalize_min_max +from .utils_motion import linear_conversion, normalize_min_max, extend_to_batch_size class ScaleType: @@ -41,13 +41,15 @@ class MultivalDynamicNode: if len(float_val) < mask_optional.shape[0]: # copies last entry enough times to match mask shape float_val = float_val + float_val[-1]*(mask_optional.shape[0]-len(float_val)) + if mask_optional.shape[0] < len(float_val): + mask_optional = extend_to_batch_size(mask_optional, len(float_val)) float_val = float_val[:mask_optional.shape[0]] float_val: Tensor = torch.tensor(float_val).unsqueeze(-1).unsqueeze(-1) # now that inputs are normalized, figure out what value to actually return if mask_optional is not None: mask_optional = mask_optional.clone() if float_is_iterable: - mask_optional = mask_optional * float_val.to(mask_optional.dtype).to(mask_optional.device) + mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device) else: mask_optional = mask_optional * float_val return (mask_optional,)