Made float_val length reflect on length of mask_optional, reorganized motion lora code
This commit is contained in:
+6
-40
@@ -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] 🎭🅐🅓",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,)
|
||||
|
||||
Reference in New Issue
Block a user