Added motion_scale parameter to Motion Model Settings (Advanced) node, support scale_multiplier in AnimateDiff and HotshotXL motion models
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 🎭🅐🅓",
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user