From 0186edee1e97ffb1c9ba0f7fa841b1c481a2bc84 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 7 Nov 2023 22:27:05 -0600 Subject: [PATCH] Implemented mask_motion_scale into Motion Model Settings. Refactored a lot of code, moved extra/experimental/deprecated nodes into separate files, --- animatediff/deprecated_nodes.py | 0 animatediff/motion_module.py | 36 ++- animatediff/motion_module_ad.py | 84 +++++- animatediff/motion_module_hsxl.py | 93 ++++-- animatediff/motion_utils.py | 107 ++++++- animatediff/nodes.py | 475 +++--------------------------- animatediff/nodes_deprecated.py | 254 ++++++++++++++++ animatediff/nodes_experimental.py | 122 ++++++++ animatediff/nodes_extras.py | 69 +++++ animatediff/sampling.py | 20 +- 10 files changed, 776 insertions(+), 484 deletions(-) delete mode 100644 animatediff/deprecated_nodes.py create mode 100644 animatediff/nodes_deprecated.py create mode 100644 animatediff/nodes_experimental.py create mode 100644 animatediff/nodes_extras.py diff --git a/animatediff/deprecated_nodes.py b/animatediff/deprecated_nodes.py deleted file mode 100644 index e69de29..0000000 diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index b520574..a5764a3 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -15,7 +15,7 @@ from .model_utils import ModelTypesSD, calculate_file_hash, get_motion_lora_path from .motion_lora import MotionLoRAList, MotionLoRAWrapper from .motion_module_ad import AnimDiffMotionWrapper, has_mid_block from .motion_module_hsxl import HotShotXLMotionWrapper, TransformerTemporal -from .motion_utils import GenericMotionWrapper, InjectorVersion +from .motion_utils import GenericMotionWrapper, InjectorVersion, normalize_min_max # inject into ModelPatcher.clone to carry over injected params over to cloned ModelPatcher orig_modelpatcher_clone = comfy_model_patcher.ModelPatcher.clone @@ -122,6 +122,18 @@ def interpolate_pe_to_length_pingpong(model_dict: dict[str, Tensor], key: str, n model_dict[key] = model_dict[key][:, :new_length] +def freeze_mask_of_pe(model_dict: dict[str, Tensor], key: str): + pe_portion = model_dict[key].shape[2] // 64 + first_pe = model_dict[key][:,:1,:] + model_dict[key][:,:,pe_portion:] = first_pe[:,:,pe_portion:] + del first_pe + + +def freeze_mask_of_attn(model_dict: dict[str, Tensor], key: str): + attn_portion = model_dict[key].shape[0] // 2 + model_dict[key][:attn_portion,:attn_portion] *= 1.5 + + 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 @@ -168,7 +180,6 @@ def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSe elif mm_settings.has_other_strength(): model_dict[key] *= mm_settings.other_strength return model_dict - #cond_or_uncond = inspect.currentframe().f_back.f_locals["transformer_options"]["cond_or_uncond"] def load_motion_module(model_name: str, motion_lora: MotionLoRAList = None, model: ModelPatcher = None, motion_model_settings = None) -> GenericMotionWrapper: # if already loaded, return it @@ -274,12 +285,12 @@ def inject_motion_module(model: ModelPatcher, motion_module: GenericMotionWrappe if not params.context_length: if params.video_length > motion_module.encoding_max_len: raise ValueError(f"Without a context window, AnimateDiff model {motion_module.mm_name} has upper limit of {motion_module.encoding_max_len} frames, but received {params.video_length} latents.") - motion_module.set_video_length(params.video_length) + motion_module.set_video_length(params.video_length, params.full_length) # otherwise, treat context_length as intended AD frame window else: if params.context_length > motion_module.encoding_max_len: raise ValueError(f"AnimateDiff model {motion_module.mm_name} has upper limit of {motion_module.encoding_max_len} frames for a context window, but received context length of {params.context_length}.") - motion_module.set_video_length(params.context_length) + motion_module.set_video_length(params.context_length, params.full_length) # inject model params.set_version(motion_module) logger.info(f"Injecting motion module {motion_module.mm_name} version {motion_module.version}.") @@ -448,6 +459,7 @@ class InjectionParams: def __init__(self, video_length: int, unlimited_area_hack: bool, apply_mm_groupnorm_hack: bool, beta_schedule: str, injector: str, model_name: str, apply_v2_models_properly: bool=False) -> None: self.video_length = video_length + self.full_length = None self.unlimited_area_hack = unlimited_area_hack self.apply_mm_groupnorm_hack = apply_mm_groupnorm_hack self.beta_schedule = beta_schedule @@ -496,6 +508,7 @@ class InjectionParams: self.video_length, self.unlimited_area_hack, self.apply_mm_groupnorm_hack, self.beta_schedule, self.injector, self.model_name, apply_v2_models_properly=self.apply_v2_models_properly, ) + new_params.full_length = self.full_length new_params.version = self.version new_params.set_context( context_length=self.context_length, context_stride=self.context_stride, @@ -558,6 +571,9 @@ class MotionModelSettings: initial_pe_idx_offset: int=0, final_pe_idx_offset: int=0, motion_pe_stretch: int=0, attn_scale: float=1.0, + mask_attn_scale: Tensor=None, + mask_attn_scale_min: float=1.0, + mask_attn_scale_max: float=1.0, ): # general strengths self.pe_strength = pe_strength @@ -577,6 +593,18 @@ class MotionModelSettings: self.motion_pe_stretch = motion_pe_stretch # attention scale settings self.attn_scale = attn_scale + # attention scale mask settings + self.mask_attn_scale = mask_attn_scale.clone() if mask_attn_scale is not None else mask_attn_scale + self.mask_attn_scale_min = mask_attn_scale_min + self.mask_attn_scale_max = mask_attn_scale_max + self._prepare_mask_attn_scale() + + def _prepare_mask_attn_scale(self): + if self.mask_attn_scale is not None: + self.mask_attn_scale = normalize_min_max(self.mask_attn_scale, self.mask_attn_scale_min, self.mask_attn_scale_max) + + def has_mask_attn_scale(self) -> bool: + return self.mask_attn_scale is not None def has_pe_strength(self) -> bool: return self.pe_strength != 1.0 diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index b50bc45..26548f8 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -5,9 +5,11 @@ import torch from einops import rearrange, repeat from torch import Tensor, nn +from comfy.utils import repeat_to_batch_size from comfy.ldm.modules.attention import FeedForward +from controlnet import broadcast_image_to from .motion_lora import MotionLoRAInfo -from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM +from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric, prepare_mask_batch def zero_module(module): @@ -58,14 +60,14 @@ class AnimDiffMotionWrapper(GenericMotionWrapper): # but only after implementing a fix for lowvram loading return self.loras is not None - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): self.AD_video_length = video_length for block in self.down_blocks: - block.set_video_length(video_length) + block.set_video_length(video_length, full_length) for block in self.up_blocks: - block.set_video_length(video_length) + block.set_video_length(video_length, full_length) if self.mid_block is not None: - self.mid_block.set_video_length(video_length) + self.mid_block.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.down_blocks: @@ -75,6 +77,14 @@ class AnimDiffMotionWrapper(GenericMotionWrapper): if self.mid_block is not None: self.mid_block.set_scale_multiplier(multiplier) + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + for block in self.down_blocks: + block.set_masks(masks, min_val, max_val) + for block in self.up_blocks: + block.set_masks(masks, min_val, max_val) + if self.mid_block is not None: + self.mid_block.set_masks(masks, min_val, max_val) + def set_sub_idxs(self, sub_idxs: list[int]): for block in self.down_blocks: block.set_sub_idxs(sub_idxs) @@ -82,6 +92,14 @@ class AnimDiffMotionWrapper(GenericMotionWrapper): block.set_sub_idxs(sub_idxs) if self.mid_block is not None: self.mid_block.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + for block in self.down_blocks: + block.reset_temp_vars() + for block in self.up_blocks: + block.reset_temp_vars() + if self.mid_block is not None: + self.mid_block.reset_temp_vars() class MotionModule(nn.Module): @@ -102,18 +120,26 @@ class MotionModule(nn.Module): if block_type == BlockType.UP: self.motion_modules.append(get_motion_module(in_channels, temporal_position_encoding_max_len)) - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): for motion_module in self.motion_modules: - motion_module.set_video_length(video_length) + motion_module.set_video_length(video_length, full_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_masks(self, masks: Tensor, min_val: float, max_val: float): + for motion_module in self.motion_modules: + motion_module.set_masks(masks, min_val, max_val) + def set_sub_idxs(self, sub_idxs: list[int]): for motion_module in self.motion_modules: motion_module.set_sub_idxs(sub_idxs) + def reset_temp_vars(self): + for motion_module in self.motion_modules: + motion_module.reset_temp_vars() + def get_motion_module(in_channels, temporal_position_encoding_max_len): return VanillaTemporalModule(in_channels=in_channels, temporal_position_encoding_max_len=temporal_position_encoding_max_len) @@ -152,20 +178,32 @@ class VanillaTemporalModule(nn.Module): self.temporal_transformer.proj_out ) - def set_video_length(self, video_length: int): - self.temporal_transformer.set_video_length(video_length) + def set_video_length(self, video_length: int, full_length: int): + self.temporal_transformer.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): self.temporal_transformer.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + self.temporal_transformer.set_masks(masks, min_val, max_val) def set_sub_idxs(self, sub_idxs: list[int]): self.temporal_transformer.set_sub_idxs(sub_idxs) + def reset_temp_vars(self): + self.temporal_transformer.reset_temp_vars() + def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None): return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + #portion = output_tensor.shape[2] // 4 + output_tensor.shape[2] // 2 + portion = output_tensor.shape[2] // 2 + ad_effect = 0.7 + #output_tensor[:,:,portion:] = input_tensor[:,:,portion:] * (1-ad_effect) + output_tensor[:,:,portion:] * ad_effect + #output_tensor[:,:,portion:] = input_tensor[:,:,portion:] #* 0.5 + return output_tensor -class TemporalTransformer3DModel(nn.Module): +class TemporalTransformer3DModel(nn.Module, TemporalTransformerGeneric): def __init__( self, in_channels, @@ -187,6 +225,7 @@ class TemporalTransformer3DModel(nn.Module): temporal_position_encoding_max_len=24, ): super().__init__() + super().temporal_transformer_init(default_length=16) inner_dim = num_attention_heads * attention_head_dim @@ -216,27 +255,35 @@ class TemporalTransformer3DModel(nn.Module): ] ) self.proj_out = nn.Linear(inner_dim, in_channels) - self.video_length = 16 - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): self.video_length = video_length + self.full_length = full_length def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.transformer_blocks: block.set_scale_multiplier(multiplier) + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + self.scale_min = min_val + self.scale_max = max_val + self.raw_scale_mask = masks + def set_sub_idxs(self, sub_idxs: list[int]): + self.sub_idxs = sub_idxs for block in self.transformer_blocks: block.set_sub_idxs(sub_idxs) def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None): - batch, channel, height, weight = hidden_states.shape + batch, channel, height, width = hidden_states.shape residual = hidden_states + scale_mask = self.get_scale_mask(hidden_states) + hidden_states = self.norm(hidden_states) inner_dim = hidden_states.shape[1] hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( - batch, height * weight, inner_dim + batch, height * width, inner_dim ) hidden_states = self.proj_in(hidden_states) @@ -247,12 +294,13 @@ class TemporalTransformer3DModel(nn.Module): encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask, video_length=self.video_length, + scale_mask=scale_mask ) # output hidden_states = self.proj_out(hidden_states) hidden_states = ( - hidden_states.reshape(batch, height, weight, inner_dim) + hidden_states.reshape(batch, height, width, inner_dim) .permute(0, 3, 1, 2) .contiguous() ) @@ -327,6 +375,7 @@ class TemporalTransformerBlock(nn.Module): encoder_hidden_states=None, attention_mask=None, video_length=None, + scale_mask=None ): for attention_block, norm in zip(self.attention_blocks, self.norms): norm_hidden_states = norm(hidden_states) @@ -338,6 +387,7 @@ class TemporalTransformerBlock(nn.Module): else None, attention_mask=attention_mask, video_length=video_length, + scale_mask=scale_mask ) + hidden_states ) @@ -406,7 +456,7 @@ class VersatileAttention(CrossAttentionMM): if multiplier is None or math.isclose(multiplier, 1.0): self.scale = None else: - self.scale = self.default_scale * multiplier + self.scale = multiplier def set_sub_idxs(self, sub_idxs: list[int]): if self.pos_encoder != None: @@ -418,6 +468,7 @@ class VersatileAttention(CrossAttentionMM): encoder_hidden_states=None, attention_mask=None, video_length=None, + scale_mask=None, ): if self.attention_mode != "Temporal": raise NotImplementedError @@ -441,6 +492,7 @@ class VersatileAttention(CrossAttentionMM): encoder_hidden_states, value=None, mask=attention_mask, + scale_mask=scale_mask, ) hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) diff --git a/animatediff/motion_module_hsxl.py b/animatediff/motion_module_hsxl.py index cfe7c5d..2532fcf 100644 --- a/animatediff/motion_module_hsxl.py +++ b/animatediff/motion_module_hsxl.py @@ -8,7 +8,7 @@ from torch import Tensor, nn from comfy.ldm.modules.attention import FeedForward from .motion_lora import MotionLoRAInfo -from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM +from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric def zero_module(module): @@ -66,14 +66,14 @@ class HotShotXLMotionWrapper(GenericMotionWrapper): # but only after implementing a fix for lowvram loading return self.loras is not None - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): self.AD_video_length = video_length for block in self.down_blocks: - block.set_video_length(video_length) + block.set_video_length(video_length, full_length) for block in self.up_blocks: - block.set_video_length(video_length) + block.set_video_length(video_length, full_length) if self.mid_block is not None: - self.mid_block.set_video_length(video_length) + self.mid_block.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.down_blocks: @@ -83,8 +83,29 @@ class HotShotXLMotionWrapper(GenericMotionWrapper): if self.mid_block is not None: self.mid_block.set_scale_multiplier(multiplier) + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + for block in self.down_blocks: + block.set_masks(masks, min_val, max_val) + for block in self.up_blocks: + block.set_masks(masks, min_val, max_val) + if self.mid_block is not None: + self.mid_block.set_masks(masks, min_val, max_val) + def set_sub_idxs(self, sub_idxs: list[int]): - pass + for block in self.down_blocks: + block.set_sub_idxs(sub_idxs) + for block in self.up_blocks: + block.set_sub_idxs(sub_idxs) + if self.mid_block is not None: + self.mid_block.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + for block in self.down_blocks: + block.reset_temp_vars() + for block in self.up_blocks: + block.reset_temp_vars() + if self.mid_block is not None: + self.mid_block.reset_temp_vars() class HotShotXLMotionModule(nn.Module): @@ -105,13 +126,25 @@ class HotShotXLMotionModule(nn.Module): if block_type == BlockType.UP: self.temporal_attentions.append(get_transformer_temporal(in_channels, max_length)) - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): for tt in self.temporal_attentions: - tt.set_video_length(video_length) + tt.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): for tt in self.temporal_attentions: tt.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + for tt in self.temporal_attentions: + tt.set_masks(masks, min_val, max_val) + + def set_sub_idxs(self, sub_idxs: list[int]): + for tt in self.temporal_attentions: + tt.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + for tt in self.temporal_attentions: + tt.reset_temp_vars() def get_transformer_temporal(in_channels, max_length) -> 'TransformerTemporal': @@ -124,7 +157,7 @@ def get_transformer_temporal(in_channels, max_length) -> 'TransformerTemporal': ) -class TransformerTemporal(nn.Module): +class TransformerTemporal(nn.Module, TemporalTransformerGeneric): def __init__( self, num_attention_heads: int, @@ -140,6 +173,7 @@ class TransformerTemporal(nn.Module): max_length = 24, ): super().__init__() + super().temporal_transformer_init(default_length=8) inner_dim = num_attention_heads * attention_head_dim @@ -163,23 +197,33 @@ class TransformerTemporal(nn.Module): ] ) self.proj_out = nn.Linear(inner_dim, in_channels) - self.video_length = 8 - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_length: int): self.video_length = video_length + self.full_length = full_length def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.transformer_blocks: block.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + self.scale_min = min_val + self.scale_max = max_val + self.raw_scale_mask = masks + + def set_sub_idxs(self, sub_idxs: list[int]): + self.sub_idxs = sub_idxs def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None): - batch, channel, height, weight = hidden_states.shape + batch, channel, height, width = hidden_states.shape residual = hidden_states + scale_mask = self.get_scale_mask(hidden_states) + hidden_states = self.norm(hidden_states) inner_dim = hidden_states.shape[1] hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( - batch, height * weight, inner_dim + batch, height * width, inner_dim ) hidden_states = self.proj_in(hidden_states) @@ -189,12 +233,14 @@ class TransformerTemporal(nn.Module): hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask, - number_of_frames=self.video_length) + number_of_frames=self.video_length, + scale_mask=scale_mask + ) # output hidden_states = self.proj_out(hidden_states) hidden_states = ( - hidden_states.reshape(batch, height, weight, inner_dim) + hidden_states.reshape(batch, height, width, inner_dim) .permute(0, 3, 1, 2) .contiguous() ) @@ -250,8 +296,7 @@ class TransformerBlock(nn.Module): 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): - + def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, number_of_frames=None, scale_mask=None): if not self.is_cross: encoder_hidden_states = None @@ -261,7 +306,8 @@ class TransformerBlock(nn.Module): norm_hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask, - number_of_frames=number_of_frames + number_of_frames=number_of_frames, + scale_mask=scale_mask ) + hidden_states norm_hidden_states = self.ff_norm(hidden_states) @@ -312,9 +358,9 @@ class TemporalAttention(CrossAttentionMM): if multiplier is None or math.isclose(multiplier, 1.0): self.scale = None else: - self.scale = self.default_scale * multiplier + self.scale = multiplier - def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, number_of_frames=8): + def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, number_of_frames=8, scale_mask=None): sequence_length = hidden_states.shape[1] hidden_states = rearrange(hidden_states, "(b f) s c -> (b s) f c", f=number_of_frames) hidden_states = self.pos_encoder(hidden_states, length=number_of_frames) @@ -322,6 +368,11 @@ class TemporalAttention(CrossAttentionMM): if encoder_hidden_states: encoder_hidden_states = repeat(encoder_hidden_states, "b n c -> (b s) n c", s=sequence_length) - hidden_states = super().forward(hidden_states, encoder_hidden_states, mask=attention_mask) + hidden_states = super().forward( + hidden_states, + encoder_hidden_states, + mask=attention_mask, + scale_mask=scale_mask + ) return rearrange(hidden_states, "(b s) f c -> (b f) s c", s=sequence_length) diff --git a/animatediff/motion_utils.py b/animatediff/motion_utils.py index ebbed07..988a3eb 100644 --- a/animatediff/motion_utils.py +++ b/animatediff/motion_utils.py @@ -1,5 +1,6 @@ from abc import ABC, abstractmethod from typing import Union +from einops import rearrange import torch import torch.nn.functional as F @@ -9,6 +10,8 @@ import comfy.model_management as model_management import comfy.ops from comfy.cli_args import args from comfy.ldm.modules.attention import attention_basic, attention_pytorch, attention_split, attention_sub_quad, default +from controlnet import broadcast_image_to +from utils import repeat_to_batch_size from .motion_lora import MotionLoRAInfo # until xformers bug is fixed, do not use xformers for VersatileAttention! TODO: change this when fix is out @@ -43,24 +46,80 @@ class CrossAttentionMM(nn.Module): self.to_out = nn.Sequential(operations.Linear(inner_dim, query_dim, dtype=dtype, device=device), nn.Dropout(dropout)) - def forward(self, x, context=None, value=None, mask=None): + def forward(self, x, context=None, value=None, mask=None, scale_mask=None): q = self.to_q(x) context = default(context, x) - k = self.to_k(context) + k: Tensor = self.to_k(context) if value is not None: v = self.to_v(value) del value 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 + # apply custom scale by multiplying k by scale factor if self.scale is not None: - k *= (self.scale / self.default_scale) + k *= self.scale + + # apply scale mask, if present + if scale_mask is not None: + k *= scale_mask + out = optimized_attention_mm(q, k, v, self.heads, mask) return self.to_out(out) +# super class to TemporalTransformer-like classes +class TemporalTransformerGeneric: + def temporal_transformer_init(self, default_length: int): + self.video_length = default_length + self.full_length = default_length + self.scale_min = 1.0 + self.scale_max = 1.0 + self.raw_scale_mask: Union[Tensor, None] = None + self.temp_scale_mask: Union[Tensor, None] = None + self.sub_idxs: Union[list[int], None] = None + + def reset_temp_vars(self): + self.temp_scale_mask = None + + def get_scale_mask(self, hidden_states: Tensor) -> Union[Tensor, None]: + # if no raw mask, return None + if self.raw_scale_mask is None: + return None + # if temp mask already calculated, return it + if self.temp_scale_mask != None: + if self.sub_idxs is not None: + return self.temp_scale_mask[:, self.sub_idxs, :] + return self.temp_scale_mask + # otherwise, calculate temp mask + shape = hidden_states.shape + batch, channel, height, width = shape + mask = prepare_mask_batch(self.raw_scale_mask, shape=(self.full_length, 1, height, width)) + mask = repeat_to_batch_size(mask, self.full_length) + # if mask not the same amount length as full length, make it match + if self.full_length != mask.shape[0]: + mask = broadcast_image_to(mask, self.full_length, 1) + # reshape mask to attention K shape (h*w, latent_count, 1) + batch, channel, height, width = mask.shape + # first, perform same operations as on hidden_states, + # turning (b, c, h, w) -> (b, h*w, c) + mask = mask.permute(0, 2, 3, 1).reshape(batch, height*width, channel) + # then, make it the same shape as attention's k, (h*w, b, c) + mask = mask.permute(1, 0, 2) + # make masks match the expected length of h*w + batched_number = shape[0] // self.video_length + if batched_number > 1: + mask = torch.cat([mask] * batched_number, dim=0) + # cache mask and set to proper device + self.temp_scale_mask = mask + # move temp_scale_mask to proper dtype + device + self.temp_scale_mask = self.temp_scale_mask.to(dtype=hidden_states.dtype, device=hidden_states.device) + # return subset of masks, if needed + if self.sub_idxs is not None: + return self.temp_scale_mask[:, self.sub_idxs, :] + return self.temp_scale_mask + + class BlockType: UP = "up" DOWN = "down" @@ -91,19 +150,35 @@ class GenericMotionWrapper(nn.Module, ABC): return self.loras is not None @abstractmethod - def set_video_length(self, video_length: int): + def set_video_length(self, video_length: int, full_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_masks(self, masks: Tensor, min_val: float, max_val: float): + pass @abstractmethod def set_sub_idxs(self, sub_idxs: list[int]): pass + + @abstractmethod + def reset_temp_vars(self): + pass + + def reset_scale_multiplier(self): + self.set_scale_multiplier(None) + + def reset_sub_idxs(self): + self.set_sub_idxs(None) + + def reset(self): + self.reset_sub_idxs() + self.reset_scale_multiplier() + self.reset_temp_vars() class GroupNormAD(torch.nn.GroupNorm): @@ -114,3 +189,19 @@ class GroupNormAD(torch.nn.GroupNorm): def forward(self, input: Tensor) -> Tensor: return F.group_norm( input, self.num_groups, self.weight, self.bias, self.eps) + + +# applies min-max normalization, from: +# https://stackoverflow.com/questions/68791508/min-max-normalization-of-a-tensor-in-pytorch +def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): + x_min, x_max = x.min(), x.max() + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + + +# adapted from comfy/sample.py +def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): + mask = mask.clone() + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") + if match_dim1: + mask = torch.cat([mask] * shape[1], dim=1) + return mask diff --git a/animatediff/nodes.py b/animatediff/nodes.py index c3305e3..26ef883 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -1,32 +1,52 @@ -import json -import os -import shutil -import subprocess -from typing import Dict, List - -import numpy as np import torch -from PIL import Image -from PIL.PngImagePlugin import PngInfo import comfy.sample as comfy_sample -import folder_paths -import nodes as comfy_nodes from comfy.model_patcher import ModelPatcher -from comfy.sd import load_checkpoint_guess_config + from .context import ContextOptions, ContextSchedules, UniformContextOptions from .logger import logger -from .model_utils import IsChangedHelper, get_available_motion_loras, get_available_motion_models, BetaSchedules, \ - raise_if_not_checkpoint_sd1_5 +from .model_utils import get_available_motion_loras, get_available_motion_models, BetaSchedules from .motion_lora import MotionLoRAInfo, MotionLoRAList -from .motion_module import InjectorVersion, InjectionParams, MotionModelSettings -from .motion_module import eject_params_from_model, inject_params_into_model, load_motion_lora, load_motion_module +from .motion_module import InjectionParams, MotionModelSettings +from .motion_module import inject_params_into_model, load_motion_lora, load_motion_module from .sampling import animatediff_sample_factory +from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect +from .nodes_experimental import AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths +from .nodes_deprecated import AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated + # override comfy_sample.sample with animatediff-support version comfy_sample.sample = animatediff_sample_factory(comfy_sample.sample) +class AnimateDiffModelSettings: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "min_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), + "max_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), + }, + "optional": { + "mask_motion_scale": ("MASK",), + } + } + + RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings" + FUNCTION = "get_motion_model_settings" + + def get_motion_model_settings(self, mask_motion_scale: torch.Tensor=None, min_motion_scale: float=1.0, max_motion_scale: float=1.0): + motion_model_settings = MotionModelSettings( + mask_attn_scale=mask_motion_scale, + mask_attn_scale_min=min_motion_scale, + mask_attn_scale_max=max_motion_scale, + ) + + return (motion_model_settings,) + + + class AnimateDiffLoRALoader: @classmethod def INPUT_TYPES(s): @@ -57,119 +77,6 @@ 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): - return { - "required": { - "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}), - "final_pe_idx_offset": ("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, 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( - pe_strength=pe_strength, - attn_strength=attn_strength, - other_strength=other_strength, - cap_initial_pe_length=cap_initial_pe_length, - 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,) - - -class AnimateDiffModelSettingsAdvancedAttnStrengths: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "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}), - "attn_q_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), - "attn_k_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), - "attn_v_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), - "attn_out_weight_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), - "attn_out_bias_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}), - "final_pe_idx_offset": ("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, pe_strength: float, attn_strength: float, - attn_q_strength: float, - attn_k_strength: float, - attn_v_strength: float, - attn_out_weight_strength: float, - attn_out_bias_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( - pe_strength=pe_strength, - attn_strength=attn_strength, - attn_q_strength=attn_q_strength, - attn_k_strength=attn_k_strength, - attn_v_strength=attn_v_strength, - attn_out_weight_strength=attn_out_weight_strength, - attn_out_bias_strength=attn_out_bias_strength, - other_strength=other_strength, - cap_initial_pe_length=cap_initial_pe_length, - 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,) - - class AnimateDiffLoaderWithContext: @classmethod def INPUT_TYPES(s): @@ -266,310 +173,20 @@ class AnimateDiffUniformContextOptions: return (context_options,) - -class AnimateDiffLoader_Deprecated: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "latents": ("LATENT",), - "model_name": (get_available_motion_models(),), - "unlimited_area_hack": ("BOOLEAN", {"default": False},), - "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.SQRT_LINEAR),), - }, - } - - RETURN_TYPES = ("MODEL", "LATENT") - CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" - FUNCTION = "load_mm_and_inject_params" - - def load_mm_and_inject_params( - self, - model: ModelPatcher, - latents: Dict[str, torch.Tensor], - model_name: str, unlimited_area_hack: bool, beta_schedule: str, - ): - raise_if_not_checkpoint_sd1_5(model) - # load motion module - load_motion_module(model_name) - # get total frames - init_frames_len = len(latents["samples"]) - # set injection params - injection_params = InjectionParams( - video_length=init_frames_len, - unlimited_area_hack=unlimited_area_hack, - apply_mm_groupnorm_hack=True, - beta_schedule=beta_schedule, - injector=InjectorVersion.V1_V2, - model_name=model_name, - ) - # inject for use in sampling code - model = inject_params_into_model(model, injection_params) - - return (model, latents) - - -class AnimateDiffLoaderAdvanced_Deprecated: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "latents": ("LATENT",), - "model_name": (get_available_motion_models(),), - "unlimited_area_hack": ("BOOLEAN", {"default": False},), - "context_length": ("INT", {"default": 16, "min": 0, "max": 1000}), - "context_stride": ("INT", {"default": 1, "min": 1, "max": 1000}), - "context_overlap": ("INT", {"default": 4, "min": 0, "max": 1000}), - "context_schedule": (ContextSchedules.CONTEXT_SCHEDULE_LIST,), - "closed_loop": ("BOOLEAN", {"default": False},), - "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.SQRT_LINEAR),), - }, - } - - RETURN_TYPES = ("MODEL", "LATENT") - CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" - FUNCTION = "load_mm_and_inject_params" - - def load_mm_and_inject_params(self, - model: ModelPatcher, - latents: Dict[str, torch.Tensor], - model_name: str, unlimited_area_hack: bool, - context_length: int, context_stride: int, context_overlap: int, context_schedule: str, closed_loop: bool, - beta_schedule: str, - ): - raise_if_not_checkpoint_sd1_5(model) - # load motion module - load_motion_module(model_name) - # get total frames - init_frames_len = len(latents["samples"]) - # set injection params - injection_params = InjectionParams( - video_length=init_frames_len, - unlimited_area_hack=unlimited_area_hack, - apply_mm_groupnorm_hack=True, - beta_schedule=beta_schedule, - injector=InjectorVersion.V1_V2, - model_name=model_name, - ) - # set context settings - injection_params.set_context( - context_length=context_length, - context_stride=context_stride, - context_overlap=context_overlap, - context_schedule=context_schedule, - closed_loop=closed_loop - ) - # inject for use in sampling code - model = inject_params_into_model(model, injection_params) - - return (model, latents) - - -class AnimateDiffUnload: - def __init__(self) -> None: - self.change = IsChangedHelper() - - @classmethod - def INPUT_TYPES(s): - return {"required": {"model": ("MODEL",)}} - - RETURN_TYPES = ("MODEL",) - CATEGORY = "Animate Diff 🎭🅐🅓" - FUNCTION = "unload_motion_modules" - - def unload_motion_modules(self, model: ModelPatcher): - # return model clone with ejected params - model = eject_params_from_model(model) - - return (model,) - - -class AnimateDiffCombine_Deprecated: - @classmethod - def INPUT_TYPES(s): - ffmpeg_path = shutil.which("ffmpeg") - #Hide ffmpeg formats if ffmpeg isn't available - if ffmpeg_path is not None: - ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")] - else: - ffmpeg_formats = [] - logger.warning("ffmpeg could not be found. Outputs that require it have been disabled") - return { - "required": { - "images": ("IMAGE",), - "frame_rate": ( - "INT", - {"default": 8, "min": 1, "max": 24, "step": 1}, - ), - "loop_count": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}), - "filename_prefix": ("STRING", {"default": "AnimateDiff"}), - "format": (["image/gif", "image/webp"] + ffmpeg_formats,), - "pingpong": ("BOOLEAN", {"default": False}), - "save_image": ("BOOLEAN", {"default": True}), - }, - "hidden": { - "prompt": "PROMPT", - "extra_pnginfo": "EXTRA_PNGINFO", - }, - } - - RETURN_TYPES = ("GIF",) - OUTPUT_NODE = True - CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" - FUNCTION = "generate_gif" - - def generate_gif( - self, - images, - frame_rate: int, - loop_count: int, - filename_prefix="AnimateDiff", - format="image/gif", - pingpong=False, - save_image=True, - prompt=None, - extra_pnginfo=None, - ): - # convert images to numpy - frames: List[Image.Image] = [] - for image in images: - img = 255.0 * image.cpu().numpy() - img = Image.fromarray(np.clip(img, 0, 255).astype(np.uint8)) - frames.append(img) - - # get output information - output_dir = ( - folder_paths.get_output_directory() - if save_image - else folder_paths.get_temp_directory() - ) - ( - full_output_folder, - filename, - counter, - subfolder, - _, - ) = folder_paths.get_save_image_path(filename_prefix, output_dir) - - metadata = PngInfo() - if prompt is not None: - metadata.add_text("prompt", json.dumps(prompt)) - if extra_pnginfo is not None: - for x in extra_pnginfo: - metadata.add_text(x, json.dumps(extra_pnginfo[x])) - - # save first frame as png to keep metadata - file = f"{filename}_{counter:05}_.png" - file_path = os.path.join(full_output_folder, file) - frames[0].save( - file_path, - pnginfo=metadata, - compress_level=4, - ) - if pingpong: - frames = frames + frames[-2:0:-1] - - format_type, format_ext = format.split("/") - file = f"{filename}_{counter:05}_.{format_ext}" - file_path = os.path.join(full_output_folder, file) - if format_type == "image": - # Use pillow directly to save an animated image - frames[0].save( - file_path, - format=format_ext.upper(), - save_all=True, - append_images=frames[1:], - duration=round(1000 / frame_rate), - loop=loop_count, - compress_level=4, - ) - else: - # Use ffmpeg to save a video - ffmpeg_path = shutil.which("ffmpeg") - if ffmpeg_path is None: - #Should never be reachable - raise ProcessLookupError("Could not find ffmpeg") - - video_format_path = folder_paths.get_full_path("video_formats", format_ext + ".json") - with open(video_format_path, 'r') as stream: - video_format = json.load(stream) - file = f"{filename}_{counter:05}_.{video_format['extension']}" - file_path = os.path.join(full_output_folder, file) - dimensions = f"{frames[0].width}x{frames[0].height}" - args = [ffmpeg_path, "-v", "error", "-f", "rawvideo", "-pix_fmt", "rgb24", - "-s", dimensions, "-r", str(frame_rate), "-i", "-"] \ - + video_format['main_pass'] + [file_path] - - env=os.environ.copy() - if "environment" in video_format: - env.update(video_format["environment"]) - with subprocess.Popen(args, stdin=subprocess.PIPE, env=env) as proc: - for frame in frames: - proc.stdin.write(frame.tobytes()) - - previews = [ - { - "filename": file, - "subfolder": subfolder, - "type": "output" if save_image else "temp", - "format": format, - } - ] - return {"ui": {"gifs": previews}} - -class CheckpointLoaderSimpleWithNoiseSelect: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), - "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.LINEAR), ) - }, - } - RETURN_TYPES = ("MODEL", "CLIP", "VAE") - FUNCTION = "load_checkpoint" - - CATEGORY = "Animate Diff 🎭🅐🅓/extras" - - def load_checkpoint(self, ckpt_name, beta_schedule, output_vae=True, output_clip=True): - ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) - out = load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) - # register chosen beta schedule on model - convert to beta_schedule name recognized by ComfyUI - out[0].model.model_sampling = BetaSchedules.to_model_sampling(beta_schedule, out[0]) - return out - - -class EmptyLatentImageLarge: - def __init__(self, device="cpu"): - self.device = device - - @classmethod - def INPUT_TYPES(s): - return {"required": { "width": ("INT", {"default": 512, "min": 64, "max": comfy_nodes.MAX_RESOLUTION, "step": 8}), - "height": ("INT", {"default": 512, "min": 64, "max": comfy_nodes.MAX_RESOLUTION, "step": 8}), - "batch_size": ("INT", {"default": 1, "min": 1, "max": 262144})}} - RETURN_TYPES = ("LATENT",) - FUNCTION = "generate" - - CATEGORY = "Animate Diff 🎭🅐🅓/extras" - - def generate(self, width, height, batch_size=1): - latent = torch.zeros([batch_size, 4, height // 8, width // 8]) - return ({"samples":latent}, ) - - NODE_CLASS_MAPPINGS = { "ADE_AnimateDiffUniformContextOptions": AnimateDiffUniformContextOptions, "ADE_AnimateDiffLoaderWithContext": AnimateDiffLoaderWithContext, "ADE_AnimateDiffLoRALoader": AnimateDiffLoRALoader, + "ADE_AnimateDiffModelSettings_Release": AnimateDiffModelSettings, + # Experimental Nodes "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths, + # Extras Nodes "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, "CheckpointLoaderSimpleWithNoiseSelect": CheckpointLoaderSimpleWithNoiseSelect, + # Deprecated Nodes "AnimateDiffLoaderV1": AnimateDiffLoader_Deprecated, "ADE_AnimateDiffLoaderV1Advanced": AnimateDiffLoaderAdvanced_Deprecated, "ADE_AnimateDiffCombine": AnimateDiffCombine_Deprecated, @@ -578,13 +195,17 @@ 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_AnimateDiffModelSettingsAdvancedAttnStrengths": "Motion Model Settings (Adv. Attn) 🎭🅐🅓", + "ADE_AnimateDiffModelSettings_Release": "Motion Model Settings 🎭🅐🅓", + # Experimental Nodes + "ADE_AnimateDiffModelSettingsSimple": "EXP Motion Model Settings (Simple) 🎭🅐🅓", + "ADE_AnimateDiffModelSettings": "EXP Motion Model Settings (Advanced) 🎭🅐🅓", + "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": "EXP Motion Model Settings (Adv. Attn) 🎭🅐🅓", + # Extras Nodes "ADE_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓", "CheckpointLoaderSimpleWithNoiseSelect": "Load Checkpoint w/ Noise Select 🎭🅐🅓", + # Deprecated Nodes "AnimateDiffLoaderV1": "AnimateDiff Loader [DEPRECATED] 🎭🅐🅓", "ADE_AnimateDiffLoaderV1Advanced": "AnimateDiff Loader (Advanced) [DEPRECATED] 🎭🅐🅓", - "ADE_AnimateDiffCombine": "AnimateDiff Combine [DEPRECATED] 🎭🅐🅓", + "ADE_AnimateDiffCombine": "DO NOT USE, USE VideoCombine from ComfyUI-VideoHelperSuite instead! AnimateDiff Combine [DEPRECATED, DO NOT USE] 🎭🅐🅓", } diff --git a/animatediff/nodes_deprecated.py b/animatediff/nodes_deprecated.py new file mode 100644 index 0000000..492bb40 --- /dev/null +++ b/animatediff/nodes_deprecated.py @@ -0,0 +1,254 @@ +import json +import os +import shutil +import subprocess +from typing import Dict, List + +import numpy as np +import torch +from PIL import Image +from PIL.PngImagePlugin import PngInfo + +import folder_paths +from comfy.model_patcher import ModelPatcher +from .context import ContextSchedules +from .logger import logger +from .model_utils import get_available_motion_models, BetaSchedules, \ + raise_if_not_checkpoint_sd1_5 +from .motion_module import InjectorVersion, InjectionParams +from .motion_module import inject_params_into_model, load_motion_module + + +class AnimateDiffLoader_Deprecated: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "latents": ("LATENT",), + "model_name": (get_available_motion_models(),), + "unlimited_area_hack": ("BOOLEAN", {"default": False},), + "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.SQRT_LINEAR),), + }, + } + + RETURN_TYPES = ("MODEL", "LATENT") + CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" + FUNCTION = "load_mm_and_inject_params" + + def load_mm_and_inject_params( + self, + model: ModelPatcher, + latents: Dict[str, torch.Tensor], + model_name: str, unlimited_area_hack: bool, beta_schedule: str, + ): + raise_if_not_checkpoint_sd1_5(model) + # load motion module + load_motion_module(model_name) + # get total frames + init_frames_len = len(latents["samples"]) + # set injection params + injection_params = InjectionParams( + video_length=init_frames_len, + unlimited_area_hack=unlimited_area_hack, + apply_mm_groupnorm_hack=True, + beta_schedule=beta_schedule, + injector=InjectorVersion.V1_V2, + model_name=model_name, + ) + # inject for use in sampling code + model = inject_params_into_model(model, injection_params) + + return (model, latents) + + +class AnimateDiffLoaderAdvanced_Deprecated: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "latents": ("LATENT",), + "model_name": (get_available_motion_models(),), + "unlimited_area_hack": ("BOOLEAN", {"default": False},), + "context_length": ("INT", {"default": 16, "min": 0, "max": 1000}), + "context_stride": ("INT", {"default": 1, "min": 1, "max": 1000}), + "context_overlap": ("INT", {"default": 4, "min": 0, "max": 1000}), + "context_schedule": (ContextSchedules.CONTEXT_SCHEDULE_LIST,), + "closed_loop": ("BOOLEAN", {"default": False},), + "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.SQRT_LINEAR),), + }, + } + + RETURN_TYPES = ("MODEL", "LATENT") + CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" + FUNCTION = "load_mm_and_inject_params" + + def load_mm_and_inject_params(self, + model: ModelPatcher, + latents: Dict[str, torch.Tensor], + model_name: str, unlimited_area_hack: bool, + context_length: int, context_stride: int, context_overlap: int, context_schedule: str, closed_loop: bool, + beta_schedule: str, + ): + raise_if_not_checkpoint_sd1_5(model) + # load motion module + load_motion_module(model_name) + # get total frames + init_frames_len = len(latents["samples"]) + # set injection params + injection_params = InjectionParams( + video_length=init_frames_len, + unlimited_area_hack=unlimited_area_hack, + apply_mm_groupnorm_hack=True, + beta_schedule=beta_schedule, + injector=InjectorVersion.V1_V2, + model_name=model_name, + ) + # set context settings + injection_params.set_context( + context_length=context_length, + context_stride=context_stride, + context_overlap=context_overlap, + context_schedule=context_schedule, + closed_loop=closed_loop + ) + # inject for use in sampling code + model = inject_params_into_model(model, injection_params) + + return (model, latents) + + +class AnimateDiffCombine_Deprecated: + @classmethod + def INPUT_TYPES(s): + ffmpeg_path = shutil.which("ffmpeg") + #Hide ffmpeg formats if ffmpeg isn't available + if ffmpeg_path is not None: + ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")] + else: + ffmpeg_formats = [] + logger.warning("This warning can be ignored, you should not be using the deprecated AnimateDiff Combine node anyway. If you are, use Video Combine from ComfyUI-VideoHelperSuite instead. ffmpeg could not be found. Outputs that require it have been disabled") + return { + "required": { + "images": ("IMAGE",), + "frame_rate": ( + "INT", + {"default": 8, "min": 1, "max": 24, "step": 1}, + ), + "loop_count": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}), + "filename_prefix": ("STRING", {"default": "AnimateDiff"}), + "format": (["image/gif", "image/webp"] + ffmpeg_formats,), + "pingpong": ("BOOLEAN", {"default": False}), + "save_image": ("BOOLEAN", {"default": True}), + }, + "hidden": { + "prompt": "PROMPT", + "extra_pnginfo": "EXTRA_PNGINFO", + }, + } + + RETURN_TYPES = ("GIF",) + OUTPUT_NODE = True + CATEGORY = "Animate Diff 🎭🅐🅓/deprecated (DO NOT USE)" + FUNCTION = "generate_gif" + + def generate_gif( + self, + images, + frame_rate: int, + loop_count: int, + filename_prefix="AnimateDiff", + format="image/gif", + pingpong=False, + save_image=True, + prompt=None, + extra_pnginfo=None, + ): + logger.warning("Do not use AnimateDiff Combine node, it is deprecated. Use Video Combine node from ComfyUI-VideoHelperSuite instead. Video nodes from VideoHelperSuite are actively maintained, more feature-rich, and also automatically attempts to get ffmpeg.") + # convert images to numpy + frames: List[Image.Image] = [] + for image in images: + img = 255.0 * image.cpu().numpy() + img = Image.fromarray(np.clip(img, 0, 255).astype(np.uint8)) + frames.append(img) + + # get output information + output_dir = ( + folder_paths.get_output_directory() + if save_image + else folder_paths.get_temp_directory() + ) + ( + full_output_folder, + filename, + counter, + subfolder, + _, + ) = folder_paths.get_save_image_path(filename_prefix, output_dir) + + metadata = PngInfo() + if prompt is not None: + metadata.add_text("prompt", json.dumps(prompt)) + if extra_pnginfo is not None: + for x in extra_pnginfo: + metadata.add_text(x, json.dumps(extra_pnginfo[x])) + + # save first frame as png to keep metadata + file = f"{filename}_{counter:05}_.png" + file_path = os.path.join(full_output_folder, file) + frames[0].save( + file_path, + pnginfo=metadata, + compress_level=4, + ) + if pingpong: + frames = frames + frames[-2:0:-1] + + format_type, format_ext = format.split("/") + file = f"{filename}_{counter:05}_.{format_ext}" + file_path = os.path.join(full_output_folder, file) + if format_type == "image": + # Use pillow directly to save an animated image + frames[0].save( + file_path, + format=format_ext.upper(), + save_all=True, + append_images=frames[1:], + duration=round(1000 / frame_rate), + loop=loop_count, + compress_level=4, + ) + else: + # Use ffmpeg to save a video + ffmpeg_path = shutil.which("ffmpeg") + if ffmpeg_path is None: + #Should never be reachable + raise ProcessLookupError("Could not find ffmpeg") + + video_format_path = folder_paths.get_full_path("video_formats", format_ext + ".json") + with open(video_format_path, 'r') as stream: + video_format = json.load(stream) + file = f"{filename}_{counter:05}_.{video_format['extension']}" + file_path = os.path.join(full_output_folder, file) + dimensions = f"{frames[0].width}x{frames[0].height}" + args = [ffmpeg_path, "-v", "error", "-f", "rawvideo", "-pix_fmt", "rgb24", + "-s", dimensions, "-r", str(frame_rate), "-i", "-"] \ + + video_format['main_pass'] + [file_path] + + env=os.environ.copy() + if "environment" in video_format: + env.update(video_format["environment"]) + with subprocess.Popen(args, stdin=subprocess.PIPE, env=env) as proc: + for frame in frames: + proc.stdin.write(frame.tobytes()) + + previews = [ + { + "filename": file, + "subfolder": subfolder, + "type": "output" if save_image else "temp", + "format": format, + } + ] + return {"ui": {"gifs": previews}} diff --git a/animatediff/nodes_experimental.py b/animatediff/nodes_experimental.py new file mode 100644 index 0000000..0e31347 --- /dev/null +++ b/animatediff/nodes_experimental.py @@ -0,0 +1,122 @@ + +import torch + +from .motion_module import MotionModelSettings + + +class AnimateDiffModelSettingsSimple: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "motion_pe_stretch": ("INT", {"default": 0, "min": 0, "step": 1}), + }, + "optional": { + "mask": ("MASK",), + "min_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), + "max_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), + } + } + + RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings/experimental" + FUNCTION = "get_motion_model_settings" + + def get_motion_model_settings(self, motion_pe_stretch: int, mask: torch.Tensor=None, min_scale: float=1.0, max_scale: float=1.0): + motion_model_settings = MotionModelSettings( + motion_pe_stretch=motion_pe_stretch + ) + + return (motion_model_settings,) + + +class AnimateDiffModelSettingsAdvanced: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "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}), + "final_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}), + }, + } + + RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings/experimental" + 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( + pe_strength=pe_strength, + attn_strength=attn_strength, + other_strength=other_strength, + cap_initial_pe_length=cap_initial_pe_length, + 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,) + + +class AnimateDiffModelSettingsAdvancedAttnStrengths: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "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}), + "attn_q_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), + "attn_k_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), + "attn_v_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), + "attn_out_weight_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}), + "attn_out_bias_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}), + "final_pe_idx_offset": ("INT", {"default": 0, "min": 0, "step": 1}), + }, + } + + RETURN_TYPES = ("MOTION_MODEL_SETTINGS",) + CATEGORY = "Animate Diff 🎭🅐🅓/motion settings/experimental" + FUNCTION = "get_motion_model_settings" + + def get_motion_model_settings(self, pe_strength: float, attn_strength: float, + attn_q_strength: float, + attn_k_strength: float, + attn_v_strength: float, + attn_out_weight_strength: float, + attn_out_bias_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( + pe_strength=pe_strength, + attn_strength=attn_strength, + attn_q_strength=attn_q_strength, + attn_k_strength=attn_k_strength, + attn_v_strength=attn_v_strength, + attn_out_weight_strength=attn_out_weight_strength, + attn_out_bias_strength=attn_out_bias_strength, + other_strength=other_strength, + cap_initial_pe_length=cap_initial_pe_length, + 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,) diff --git a/animatediff/nodes_extras.py b/animatediff/nodes_extras.py new file mode 100644 index 0000000..60fba49 --- /dev/null +++ b/animatediff/nodes_extras.py @@ -0,0 +1,69 @@ +import torch + +import folder_paths +import nodes as comfy_nodes +from comfy.model_patcher import ModelPatcher +from comfy.sd import load_checkpoint_guess_config +from .logger import logger +from .model_utils import IsChangedHelper, BetaSchedules +from .motion_module import eject_params_from_model + + +class AnimateDiffUnload: + def __init__(self) -> None: + self.change = IsChangedHelper() + + @classmethod + def INPUT_TYPES(s): + return {"required": {"model": ("MODEL",)}} + + RETURN_TYPES = ("MODEL",) + CATEGORY = "Animate Diff 🎭🅐🅓/extras" + FUNCTION = "unload_motion_modules" + + def unload_motion_modules(self, model: ModelPatcher): + # return model clone with ejected params + model = eject_params_from_model(model) + + return (model,) + + +class CheckpointLoaderSimpleWithNoiseSelect: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), + "beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.LINEAR), ) + }, + } + RETURN_TYPES = ("MODEL", "CLIP", "VAE") + FUNCTION = "load_checkpoint" + + CATEGORY = "Animate Diff 🎭🅐🅓/extras" + + def load_checkpoint(self, ckpt_name, beta_schedule, output_vae=True, output_clip=True): + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + out = load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) + # register chosen beta schedule on model - convert to beta_schedule name recognized by ComfyUI + out[0].model.model_sampling = BetaSchedules.to_model_sampling(beta_schedule, out[0]) + return out + + +class EmptyLatentImageLarge: + def __init__(self, device="cpu"): + self.device = device + + @classmethod + def INPUT_TYPES(s): + return {"required": { "width": ("INT", {"default": 512, "min": 64, "max": comfy_nodes.MAX_RESOLUTION, "step": 8}), + "height": ("INT", {"default": 512, "min": 64, "max": comfy_nodes.MAX_RESOLUTION, "step": 8}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 262144})}} + RETURN_TYPES = ("LATENT",) + FUNCTION = "generate" + + CATEGORY = "Animate Diff 🎭🅐🅓/extras" + + def generate(self, width, height, batch_size=1): + latent = torch.zeros([batch_size, 4, height // 8, width // 8]) + return ({"samples":latent}, ) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index fa4b607..b3d6a59 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -131,6 +131,7 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: # get amount of latents passed in, and inject into model latents = args[-1] params.video_length = latents.size(0) + params.full_length = latents.size(0) model = inject_params_into_model(model, params) # reset global state ADGS.reset() @@ -172,6 +173,13 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: # apply scale multiplier, if needed motion_module.set_scale_multiplier(params.motion_model_settings.attn_scale) + # apply scale mask, if needed + motion_module.set_masks( + masks=params.motion_model_settings.mask_attn_scale, + min_val=params.motion_model_settings.mask_attn_scale_min, + max_val=params.motion_model_settings.mask_attn_scale_max + ) + # handle GLOBALSTATE vars and step tally ADGS.motion_module = motion_module ADGS.update_with_inject_params(params) @@ -192,10 +200,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: # attempt to eject motion module eject_motion_module(model=model) if motion_module is not None: - # reset motion module scale multiplier - motion_module.reset_scale_multiplier() - # reset motion module sub_idxs - motion_module.set_sub_idxs(None) + # reset motion module + motion_module.reset() # if loras are present, remove model so it can be re-loaded next time with fresh weights if motion_module.has_loras(): unload_motion_module(motion_module) @@ -516,10 +522,8 @@ def sliding_sampling_function(model_function, x, timestep, uncond, cond, cond_sc # perform calc_cond_uncond_batch per context window for ctx_idxs in context_scheduler(ADGS.current_step, ADGS.total_steps, ADGS.video_length, ADGS.context_frames, ADGS.context_stride, ADGS.context_overlap, ADGS.closed_loop): - # idxs of positional encoders in motion module to use, if needed (experimental, so disabled for now) - if ADGS.sync_context_to_pe: - ADGS.sub_idxs = ctx_idxs - ADGS.motion_module.set_sub_idxs(ADGS.sub_idxs) + ADGS.sub_idxs = ctx_idxs + ADGS.motion_module.set_sub_idxs(ADGS.sub_idxs) # account for all portions of input frames full_idxs = [] for n in range(axes_factor):