Add Motion Module Settings (Simple) node, added motion_pe_stretch setting

This commit is contained in:
Jedrzej Kosinski
2023-10-30 09:20:53 -05:00
parent 8be9133ce5
commit e078dcfd69
2 changed files with 49 additions and 9 deletions
+22 -8
View File
@@ -89,13 +89,26 @@ def load_motion_lora(lora_name: str) -> MotionLoRAWrapper:
return lora
def interpolate_pe_to_length(model_dict: dict[str, Tensor], key: str, new_length: int):
pe_shape = model_dict[key].shape
temp_pe = rearrange(model_dict[key], "(t b) f d -> t b f d", t=1)
temp_pe = F.interpolate(temp_pe, size=(new_length, pe_shape[-1]), mode="bilinear")
temp_pe = rearrange(temp_pe, "t b f d -> (t b) f d", t=1)
model_dict[key] = temp_pe
del temp_pe
def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSettings') -> dict[str, Tensor]:
if not mm_settings.has_anything_to_apply():
return model_dict
for key in model_dict:
if "attention_blocks" in key:
# apply pe_strength, if needed
if "pos_encoder" in key:
# apply simple motion pe stretch, if needed
if mm_settings.has_motion_pe_stretch():
new_pe_length = model_dict[key].shape[1] + mm_settings.motion_pe_stretch
interpolate_pe_to_length(model_dict, key, new_length=new_pe_length)
# apply pe_strength, if needed
if mm_settings.has_pe_strength():
model_dict[key] *= mm_settings.pe_strength
# apply pe_idx_offset, if needed
@@ -106,12 +119,7 @@ def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSe
model_dict[key] = model_dict[key][:, :mm_settings.cap_initial_pe_length]
# apply interpolate_pe_to_length, if needed
if mm_settings.has_interpolate_pe_to_length():
pe_shape = model_dict[key].shape
temp_pe = rearrange(model_dict[key], "(t b) f d -> t b f d", t=1)
temp_pe = F.interpolate(temp_pe, size=(mm_settings.interpolate_pe_to_length, pe_shape[-1]), mode="bilinear")
temp_pe = rearrange(temp_pe, "t b f d -> (t b) f d", t=1)
model_dict[key] = temp_pe
del temp_pe
interpolate_pe_to_length(model_dict, key, new_length=mm_settings.interpolate_pe_to_length)
# apply final_pe_idx_offset, if needed
if mm_settings.has_final_pe_idx_offset():
model_dict[key] = model_dict[key][:, mm_settings.final_pe_idx_offset:]
@@ -501,6 +509,7 @@ class MotionModelSettings:
pe_strength: float=1.0, attn_strength: float=1.0, other_strength: float=1.0,
cap_initial_pe_length: int=0, interpolate_pe_to_length: int=0,
initial_pe_idx_offset: int=0, final_pe_idx_offset: int=0,
motion_pe_stretch: int=0,
attn_scale: float=1.0,
):
# PE-interpolation settings
@@ -511,6 +520,7 @@ class MotionModelSettings:
self.interpolate_pe_to_length = interpolate_pe_to_length
self.initial_pe_idx_offset = initial_pe_idx_offset
self.final_pe_idx_offset = final_pe_idx_offset
self.motion_pe_stretch = motion_pe_stretch
# attention scale settings
self.attn_scale = attn_scale
@@ -535,6 +545,9 @@ class MotionModelSettings:
def has_final_pe_idx_offset(self) -> bool:
return self.final_pe_idx_offset > 0
def has_motion_pe_stretch(self) -> bool:
return self.motion_pe_stretch > 0
def has_anything_to_apply(self) -> bool:
return self.has_pe_strength() \
or self.has_attn_strength() \
@@ -542,4 +555,5 @@ class MotionModelSettings:
or self.has_cap_initial_pe_length() \
or self.has_interpolate_pe_to_length() \
or self.has_initial_pe_idx_offset() \
or self.has_final_pe_idx_offset()
or self.has_final_pe_idx_offset() \
or self.has_motion_pe_stretch()
+27 -1
View File
@@ -57,6 +57,27 @@ class AnimateDiffLoRALoader:
return (prev_motion_lora,)
class AnimateDiffModelSettingsSimple:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"motion_pe_stretch": ("INT", {"default": 0, "min": 0, "step": 1}),
},
}
RETURN_TYPES = ("MOTION_MODEL_SETTINGS",)
CATEGORY = "Animate Diff 🎭🅐🅓/motion settings"
FUNCTION = "get_motion_model_settings"
def get_motion_model_settings(self, motion_pe_stretch: int):
motion_model_settings = MotionModelSettings(
motion_pe_stretch=motion_pe_stretch
)
return (motion_model_settings,)
class AnimateDiffModelSettingsAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -65,6 +86,7 @@ class AnimateDiffModelSettingsAdvanced:
"pe_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"attn_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"other_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"motion_pe_stretch": ("INT", {"default": 0, "min": 0, "step": 1}),
"cap_initial_pe_length": ("INT", {"default": 0, "min": 0, "step": 1}),
"interpolate_pe_to_length": ("INT", {"default": 0, "min": 0, "step": 1}),
"initial_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}),
@@ -73,10 +95,11 @@ class AnimateDiffModelSettingsAdvanced:
}
RETURN_TYPES = ("MOTION_MODEL_SETTINGS",)
CATEGORY = "Animate Diff 🎭🅐🅓"
CATEGORY = "Animate Diff 🎭🅐🅓/motion settings"
FUNCTION = "get_motion_model_settings"
def get_motion_model_settings(self, pe_strength: float, attn_strength: float, other_strength: float,
motion_pe_stretch: int,
cap_initial_pe_length: int, interpolate_pe_to_length: int,
initial_pe_idx_offset: int, final_pe_idx_offset: int):
motion_model_settings = MotionModelSettings(
@@ -87,6 +110,7 @@ class AnimateDiffModelSettingsAdvanced:
interpolate_pe_to_length=interpolate_pe_to_length,
initial_pe_idx_offset=initial_pe_idx_offset,
final_pe_idx_offset=final_pe_idx_offset,
motion_pe_stretch=motion_pe_stretch
)
return (motion_model_settings,)
@@ -485,6 +509,7 @@ NODE_CLASS_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": AnimateDiffUniformContextOptions,
"ADE_AnimateDiffLoaderWithContext": AnimateDiffLoaderWithContext,
"ADE_AnimateDiffLoRALoader": AnimateDiffLoRALoader,
"ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple,
"ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced,
"ADE_AnimateDiffUnload": AnimateDiffUnload,
"ADE_EmptyLatentImageLarge": EmptyLatentImageLarge,
@@ -497,6 +522,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": "Uniform Context Options 🎭🅐🅓",
"ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader 🎭🅐🅓",
"ADE_AnimateDiffLoRALoader": "AnimateDiff LoRA Loader 🎭🅐🅓",
"ADE_AnimateDiffModelSettingsSimple": "Motion Model Settings (Simple) 🎭🅐🅓",
"ADE_AnimateDiffModelSettings": "Motion Model Settings (Advanced) 🎭🅐🅓",
"ADE_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓",
"ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓",