Implemented mask_motion_scale into Motion Model Settings. Refactored a lot of code, moved extra/experimental/deprecated nodes into separate files,

This commit is contained in:
Jedrzej Kosinski
2023-11-07 22:27:05 -06:00
parent 1fd6a45b8d
commit 0186edee1e
10 changed files with 776 additions and 484 deletions
View File
+32 -4
View File
@@ -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
+68 -16
View File
@@ -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)
+72 -21
View File
@@ -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)
+99 -8
View File
@@ -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
+48 -427
View File
@@ -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] 🎭🅐🅓",
}
+254
View File
@@ -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}}
+122
View File
@@ -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,)
+69
View File
@@ -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}, )
+12 -8
View File
@@ -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):