Added motion_scale parameter to Motion Model Settings (Advanced) node, support scale_multiplier in AnimateDiff and HotshotXL motion models

This commit is contained in:
Jedrzej Kosinski
2023-10-29 21:45:40 -05:00
parent d2f624eca4
commit 5bbcae4d22
6 changed files with 101 additions and 18 deletions
+6 -1
View File
@@ -500,7 +500,10 @@ class MotionModelSettings:
def __init__(self,
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):
initial_pe_idx_offset: int=0, final_pe_idx_offset: int=0,
attn_scale: float=1.0,
):
# PE-interpolation settings
self.pe_strength = pe_strength
self.attn_strength = attn_strength
self.other_strength = other_strength
@@ -508,6 +511,8 @@ 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
# attention scale settings
self.attn_scale = attn_scale
def has_pe_strength(self) -> bool:
return self.pe_strength != 1.0
+36 -6
View File
@@ -1,3 +1,4 @@
from typing import Iterable, Union
import torch
from torch import Tensor, nn
@@ -37,9 +38,9 @@ def has_mid_block(mm_state_dict: dict[str, Tensor]):
class AnimDiffMotionWrapper(GenericMotionWrapper):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_hash: str, mm_name: str="mm_sd_v15.ckpt" , loras: list[MotionLoRAInfo]=None):
super().__init__(mm_hash, mm_name, loras)
self.down_blocks = nn.ModuleList([])
self.up_blocks = nn.ModuleList([])
self.mid_block = None
self.down_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.up_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.mid_block: Union[MotionModule, None] = None
self.encoding_max_len = get_ad_temporal_position_encoding_max_len(mm_state_dict, mm_name)
for c in (320, 640, 1280, 1280):
self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN))
@@ -68,6 +69,14 @@ class AnimDiffMotionWrapper(GenericMotionWrapper):
if self.mid_block is not None:
self.mid_block.set_video_length(video_length)
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.mid_block is not None:
self.mid_block.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
for block in self.down_blocks:
block.set_sub_idxs(sub_idxs)
@@ -82,7 +91,7 @@ class MotionModule(nn.Module):
super().__init__()
if block_type == BlockType.MID:
# mid blocks contain only a single VanillaTemporalModule
self.motion_modules = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len)])
self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len)])
else:
# down blocks contain two VanillaTemporalModules
self.motion_modules = nn.ModuleList(
@@ -99,6 +108,10 @@ class MotionModule(nn.Module):
for motion_module in self.motion_modules:
motion_module.set_video_length(video_length)
def set_scale_multiplier(self, multiplier: Union[float, None]):
for motion_module in self.motion_modules:
motion_module.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
for motion_module in self.motion_modules:
motion_module.set_sub_idxs(sub_idxs)
@@ -144,6 +157,9 @@ class VanillaTemporalModule(nn.Module):
def set_video_length(self, video_length: int):
self.temporal_transformer.set_video_length(video_length)
def set_scale_multiplier(self, multiplier: Union[float, None]):
self.temporal_transformer.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
self.temporal_transformer.set_sub_idxs(sub_idxs)
@@ -181,7 +197,7 @@ class TemporalTransformer3DModel(nn.Module):
)
self.proj_in = nn.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList(
self.transformer_blocks: Iterable[TemporalTransformerBlock] = nn.ModuleList(
[
TemporalTransformerBlock(
dim=inner_dim,
@@ -207,6 +223,10 @@ class TemporalTransformer3DModel(nn.Module):
def set_video_length(self, video_length: int):
self.video_length = video_length
def set_scale_multiplier(self, multiplier: Union[float, None]):
for block in self.transformer_blocks:
block.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
for block in self.transformer_blocks:
block.set_sub_idxs(sub_idxs)
@@ -289,12 +309,16 @@ class TemporalTransformerBlock(nn.Module):
)
norms.append(nn.LayerNorm(dim))
self.attention_blocks = nn.ModuleList(attention_blocks)
self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks)
self.norms = nn.ModuleList(norms)
self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"))
self.ff_norm = nn.LayerNorm(dim)
def set_scale_multiplier(self, multiplier: Union[float, None]):
for block in self.attention_blocks:
block.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
for block in self.attention_blocks:
block.set_sub_idxs(sub_idxs)
@@ -380,6 +404,12 @@ class VersatileAttention(CrossAttentionMM):
def extra_repr(self):
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
def set_scale_multiplier(self, multiplier: Union[float, None]):
if multiplier is None or math.isclose(multiplier, 1.0):
self.scale = None
else:
self.scale = self.default_scale * multiplier
def set_sub_idxs(self, sub_idxs: list[int]):
if self.pos_encoder != None:
self.pos_encoder.set_sub_idxs(sub_idxs)
+33 -7
View File
@@ -1,5 +1,5 @@
# original HotShotXL components adapted from https://github.com/hotshotco/Hotshot-XL/blob/main/hotshot_xl/models/transformer_temporal.py
from typing import Optional
from typing import Iterable, Optional, Union
import torch
from torch import Tensor, nn
@@ -45,9 +45,9 @@ def has_mid_block(mm_state_dict: dict[str, Tensor]):
class HotShotXLMotionWrapper(GenericMotionWrapper):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_hash: str, mm_name: str="mm_sd_v15.ckpt", loras: list[MotionLoRAInfo]=None):
super().__init__(mm_hash, mm_name, loras)
self.down_blocks = nn.ModuleList([])
self.up_blocks = nn.ModuleList([])
self.mid_block = None
self.down_blocks: Iterable[HotShotXLMotionModule] = nn.ModuleList([])
self.up_blocks: Iterable[HotShotXLMotionModule] = nn.ModuleList([])
self.mid_block: Union[HotShotXLMotionModule, None] = None
self.encoding_max_len = get_hsxl_temporal_position_encoding_max_len(mm_state_dict, mm_name)
for c in (320, 640, 1280):
self.down_blocks.append(HotShotXLMotionModule(c, block_type=BlockType.DOWN, max_length=self.encoding_max_len))
@@ -76,6 +76,14 @@ class HotShotXLMotionWrapper(GenericMotionWrapper):
if self.mid_block is not None:
self.mid_block.set_video_length(video_length)
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.mid_block is not None:
self.mid_block.set_scale_multiplier(multiplier)
def set_sub_idxs(self, sub_idxs: list[int]):
pass
@@ -85,7 +93,7 @@ class HotShotXLMotionModule(nn.Module):
super().__init__()
if block_type == BlockType.MID:
# mid blocks contain only a single TransformerTemporal
self.temporal_attentions = nn.ModuleList([get_transformer_temporal(in_channels, max_length)])
self.temporal_attentions: Iterable[TransformerTemporal] = nn.ModuleList([get_transformer_temporal(in_channels, max_length)])
else:
# down blocks contain two TransformerTemporals
self.temporal_attentions = nn.ModuleList(
@@ -102,6 +110,10 @@ class HotShotXLMotionModule(nn.Module):
for tt in self.temporal_attentions:
tt.set_video_length(video_length)
def set_scale_multiplier(self, multiplier: Union[float, None]):
for tt in self.temporal_attentions:
tt.set_scale_multiplier(multiplier)
def get_transformer_temporal(in_channels, max_length) -> 'TransformerTemporal':
num_attention_heads = 8
@@ -135,7 +147,7 @@ class TransformerTemporal(nn.Module):
self.norm = GroupNormAD(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
self.proj_in = nn.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList(
self.transformer_blocks: Iterable[TransformerBlock] = nn.ModuleList(
[
TransformerBlock(
dim=inner_dim,
@@ -157,6 +169,10 @@ class TransformerTemporal(nn.Module):
def set_video_length(self, video_length: int):
self.video_length = video_length
def set_scale_multiplier(self, multiplier: Union[float, None]):
for block in self.transformer_blocks:
block.set_scale_multiplier(multiplier)
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
batch, channel, height, weight = hidden_states.shape
residual = hidden_states
@@ -225,12 +241,16 @@ class TransformerBlock(nn.Module):
)
norms.append(nn.LayerNorm(dim))
self.attention_blocks = nn.ModuleList(attention_blocks)
self.attention_blocks: Iterable[TemporalAttention] = nn.ModuleList(attention_blocks)
self.norms = nn.ModuleList(norms)
self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"))
self.ff_norm = nn.LayerNorm(dim)
def set_scale_multiplier(self, multiplier: Union[float, None]):
for block in self.attention_blocks:
block.set_scale_multiplier(multiplier)
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, number_of_frames=None):
if not self.is_cross:
@@ -289,6 +309,12 @@ class TemporalAttention(CrossAttentionMM):
super().__init__(*args, **kwargs)
self.pos_encoder = PositionalEncoding(kwargs["query_dim"], dropout=0, max_length=max_length)
def set_scale_multiplier(self, multiplier: Union[float, None]):
if multiplier is None or math.isclose(multiplier, 1.0):
self.scale = None
else:
self.scale = self.default_scale * multiplier
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, number_of_frames=8):
sequence_length = hidden_states.shape[1]
hidden_states = rearrange(hidden_states, "(b f) s c -> (b s) f c", f=number_of_frames)
+14
View File
@@ -1,3 +1,4 @@
from typing import Union
import torch
from torch import Tensor, nn
import torch.nn.functional as F
@@ -34,6 +35,8 @@ class CrossAttentionMM(nn.Module):
self.heads = heads
self.dim_head = dim_head
self.scale = None
self.default_scale = dim_head ** -0.5
self.to_q = operations.Linear(query_dim, inner_dim, bias=False, dtype=dtype, device=device)
self.to_k = operations.Linear(context_dim, inner_dim, bias=False, dtype=dtype, device=device)
@@ -51,6 +54,10 @@ class CrossAttentionMM(nn.Module):
else:
v = self.to_v(context)
# apply custom scale by multiplying k by scale factor;
# division by default_scale is needed to account for internal attn code multiplying by default_scale
if self.scale is not None:
k *= (self.scale / self.default_scale)
out = optimized_attention_mm(q, k, v, self.heads, mask)
return self.to_out(out)
@@ -88,6 +95,13 @@ class GenericMotionWrapper(nn.Module, ABC):
def set_video_length(self, video_length: int):
pass
@abstractmethod
def set_scale_multiplier(self, multiplier: Union[float, None]):
pass
def reset_scale_multiplier(self):
self.set_scale_multiplier(None)
@abstractmethod
def set_sub_idxs(self, sub_idxs: list[int]):
pass
+6 -4
View File
@@ -57,7 +57,7 @@ class AnimateDiffLoRALoader:
return (prev_motion_lora,)
class AnimateDiffModelSettings:
class AnimateDiffModelSettingsAdvanced:
@classmethod
def INPUT_TYPES(s):
return {
@@ -69,6 +69,7 @@ class AnimateDiffModelSettings:
"interpolate_pe_to_length": ("INT", {"default": 0, "min": 0, "step": 1}),
"initial_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}),
"final_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}),
"motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}),
},
}
@@ -78,7 +79,7 @@ class AnimateDiffModelSettings:
def get_motion_model_settings(self, pe_strength: float, attn_strength: float, other_strength: float,
cap_initial_pe_length: int, interpolate_pe_to_length: int,
initial_pe_idx_offset: int, final_pe_idx_offset: int):
initial_pe_idx_offset: int, final_pe_idx_offset: int, motion_scale: float):
motion_model_settings = MotionModelSettings(
pe_strength=pe_strength,
attn_strength=attn_strength,
@@ -87,6 +88,7 @@ class AnimateDiffModelSettings:
interpolate_pe_to_length=interpolate_pe_to_length,
initial_pe_idx_offset=initial_pe_idx_offset,
final_pe_idx_offset=final_pe_idx_offset,
attn_scale=motion_scale,
)
return (motion_model_settings,)
@@ -479,7 +481,7 @@ NODE_CLASS_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": AnimateDiffUniformContextOptions,
"ADE_AnimateDiffLoaderWithContext": AnimateDiffLoaderWithContext,
"ADE_AnimateDiffLoRALoader": AnimateDiffLoRALoader,
"ADE_AnimateDiffModelSettings": AnimateDiffModelSettings,
"ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced,
"ADE_AnimateDiffUnload": AnimateDiffUnload,
"ADE_EmptyLatentImageLarge": EmptyLatentImageLarge,
"CheckpointLoaderSimpleWithNoiseSelect": CheckpointLoaderSimpleWithNoiseSelect,
@@ -491,7 +493,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": "Uniform Context Options 🎭🅐🅓",
"ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader 🎭🅐🅓",
"ADE_AnimateDiffLoRALoader": "AnimateDiff LoRA Loader 🎭🅐🅓",
"ADE_AnimateDiffModelSettings": "AnimateDiff Motion Model Settings 🎭🅐🅓",
"ADE_AnimateDiffModelSettings": "Motion Model Settings (Advanced) 🎭🅐🅓",
"ADE_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓",
"ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓",
"CheckpointLoaderSimpleWithNoiseSelect": "Load Checkpoint w/ Noise Select 🎭🅐🅓",
+6
View File
@@ -156,6 +156,9 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable:
beta_schedule = BetaSchedules.to_name(params.beta_schedule)
model.model.register_schedule(given_betas=None, beta_schedule=beta_schedule, timesteps=1000, linear_start=0.00085, linear_end=0.012, cosine_s=8e-3)
# apply scale multiplier, if needed
motion_module.set_scale_multiplier(params.motion_model_settings.attn_scale)
# handle GLOBALSTATE vars and step tally
ADGS.motion_module = motion_module
ADGS.update_with_inject_params(params)
@@ -175,6 +178,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable:
finally:
# attempt to eject motion module
eject_motion_module(model=model)
# reset motion module scale multiplier
motion_module.reset_scale_multiplier()
# reset motion module sub_idxs
motion_module.set_sub_idxs(None)
# if loras are present, remove model so it can be re-loaded next time with fresh weights
@@ -474,6 +479,7 @@ def sliding_sampling_function(model_function, x, timestep, uncond, cond, cond_sc
raise ValueError(f"Control type {type(control_item).__name__} may not support required features for sliding context window; \
use Control objects from Kosinkadink/Advanced-ControlNet nodes, or make sure Advanced-ControlNet is updated.")
resized_actual_cond[key] = control_item
del control_item
elif isinstance(cond_item, dict):
new_cond_item = cond_item.copy()
# when in dictionary, look for tensors and CONDCrossAttn [comfy/conds.py] (has cond attr that is a tensor)