Added ad_keyframes input to Gen1 loader, fixed typo + small cleanup

This commit is contained in:
Jedrzej Kosinski
2024-01-11 04:28:13 -06:00
parent fd3bde20fe
commit 71f976bc6f
4 changed files with 13 additions and 15 deletions
+2 -2
View File
@@ -14,7 +14,7 @@ from comfy.model_base import BaseModel
from .context import ContextOptions, ContextOptions
from .motion_module_ad import AnimateDiffModel, has_mid_block, normalize_ad_state_dict
from .logger import logger
from .motion_utils import ADKeyframe, ADKeygrameGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max
from .motion_utils import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max
from .motion_lora import MotionLoraInfo, MotionLoraList
from .model_utils import get_motion_lora_path, get_motion_model_path, get_sd_model_type
from .sample_settings import SampleSettings, SeedNoiseGeneration
@@ -94,7 +94,7 @@ class MotionModelPatcher(ModelPatcher):
self.model: AnimateDiffModel = self.model
self.timestep_percent_range = (0.0, 1.0)
self.timestep_range: tuple[float, float] = None
self.keyframes: ADKeygrameGroup = ADKeygrameGroup()
self.keyframes: ADKeyframeGroup = ADKeyframeGroup()
self.scale_multival = None
self.effect_multival = None
+3 -3
View File
@@ -162,7 +162,7 @@ class ADKeyframe:
return self.effect_multival is not None
class ADKeygrameGroup:
class ADKeyframeGroup:
def __init__(self):
self.keyframes: list[ADKeyframe] = []
self.keyframes.append(ADKeyframe())
@@ -197,8 +197,8 @@ class ADKeygrameGroup:
def is_empty(self) -> bool:
return len(self.keyframes) == 0
def clone(self) -> 'ADKeygrameGroup':
cloned = ADKeygrameGroup()
def clone(self) -> 'ADKeyframeGroup':
cloned = ADKeyframeGroup()
for tk in self.keyframes:
cloned.add(tk)
return cloned
+4 -6
View File
@@ -7,6 +7,7 @@ from comfy.model_patcher import ModelPatcher
from .context import ContextOptions, ContextSchedules
from .logger import logger
from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
from .motion_utils import ADKeyframeGroup
from .motion_lora import MotionLoraInfo, MotionLoraList
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelSettings, load_motion_module
from .sample_settings import SampleSettings, SeedNoiseGeneration
@@ -30,6 +31,7 @@ class AnimateDiffLoaderWithContext:
"sample_settings": ("SAMPLE_SETTINGS",),
"motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}),
"apply_v2_models_properly": ("BOOLEAN", {"default": True}),
"ad_keyframes": ("AD_KEYFRAMES",),
}
}
@@ -42,7 +44,7 @@ class AnimateDiffLoaderWithContext:
model: ModelPatcher,
model_name: str, beta_schedule: str,# apply_mm_groupnorm_hack: bool,
context_options: ContextOptions=None, motion_lora: MotionLoraList=None, motion_model_settings: MotionModelSettings=None,
sample_settings: SampleSettings=None, motion_scale: float=1.0, apply_v2_models_properly: bool=False,
sample_settings: SampleSettings=None, motion_scale: float=1.0, apply_v2_models_properly: bool=False, ad_keyframes: ADKeyframeGroup=None,
):
# load motion module
motion_model = load_motion_module(model_name, model, motion_lora=motion_lora, motion_model_settings=motion_model_settings)
@@ -65,11 +67,7 @@ class AnimateDiffLoaderWithContext:
else:
motion_model.scale_multival = params.motion_model_settings.attn_scale
# apply scale multiplier, if needed
#motion_model.model._set_scale_multiplier(params.motion_model_settings.attn_scale)
# apply scale mask, if needed
#motion_model.model._set_scale_mask(mask=params.motion_model_settings.mask_attn_scale)
motion_model.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeyframeGroup()
model = ModelPatcherAndInjector(model)
model.motion_models = MotionModelGroup(motion_model)
+4 -4
View File
@@ -7,7 +7,7 @@ from comfy.model_patcher import ModelPatcher
from .context import ContextOptions, ContextSchedules
from .logger import logger
from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
from .motion_utils import ADKeygrameGroup, ADKeyframe
from .motion_utils import ADKeyframeGroup, ADKeyframe
from .motion_lora import MotionLoraInfo, MotionLoraList
from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, MotionModelSettings,
load_motion_module, load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2)
@@ -91,7 +91,7 @@ class ApplyAnimateDiffModelNode:
FUNCTION = "apply_motion_model"
def apply_motion_model(self, motion_model: MotionModelPatcher, start_percent: float=0.0, end_percent: float=1.0,
motion_lora: MotionLoraList=None, ad_keyframes: ADKeygrameGroup=None,
motion_lora: MotionLoraList=None, ad_keyframes: ADKeyframeGroup=None,
scale_multival=None, effect_multival=None,
prev_m_models: MotionModelGroup=None,):
# set up motion models list
@@ -105,7 +105,7 @@ class ApplyAnimateDiffModelNode:
load_motion_lora_as_patches(motion_model, lora)
motion_model.scale_multival = scale_multival
motion_model.effect_multival = effect_multival
motion_model.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeygrameGroup()
motion_model.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeyframeGroup()
motion_model.timestep_percent_range = (start_percent, end_percent)
# add to beginning, so that after injection, it will be the earliest of prev_m_models to be run
prev_m_models.add_to_start(mm=motion_model)
@@ -189,7 +189,7 @@ class ADKeyframeNode:
scale_multival: [float, torch.Tensor]=None, effect_multival: [float, torch.Tensor]=None,
inherit_missing: bool=True, guarantee_usage: bool=True):
if not prev_ad_keyframes:
prev_ad_keyframes = ADKeygrameGroup()
prev_ad_keyframes = ADKeyframeGroup()
prev_ad_keyframes.clone()
keyframe = ADKeyframe(start_percent=start_percent, scale_multival=scale_multival, effect_multival=effect_multival,
inherit_missing=inherit_missing, guarantee_usage=guarantee_usage)