diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index 1ca129d..163dc6b 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -10,7 +10,7 @@ import torch from torch import Tensor, nn import folder_paths -from comfy.model_base import SDXL, BaseModel +from comfy.model_base import SDXL, BaseModel, model_sampling from comfy.model_patcher import ModelPatcher from comfy.model_management import xformers_enabled @@ -26,6 +26,11 @@ class IsChangedHelper: self.val = (self.val + 1) % 100 +class ModelSamplingConfig: + def __init__(self, beta_schedule: str): + self.beta_schedule = beta_schedule + + class BetaSchedules: SQRT_LINEAR = "sqrt_linear (AnimateDiff)" LINEAR = "linear (HotshotXL/default)" @@ -46,6 +51,14 @@ class BetaSchedules: @classmethod def to_name(cls, alias: str): return cls.ALIAS_MAP[alias] + + @classmethod + def to_config(cls, alias: str) -> ModelSamplingConfig: + return ModelSamplingConfig(cls.to_name(alias)) + + @classmethod + def to_model_sampling(cls, alias: str, model: ModelPatcher): + return model_sampling(cls.to_config(alias), model_type=model.model.model_type) @staticmethod def get_alias_list_with_first_element(first_element: str): @@ -57,18 +70,14 @@ class BetaSchedules: class BetaScheduleCache: def __init__(self, model: ModelPatcher): - self.betas = model.model.betas.cpu().clone().detach() - self.linear_start = model.model.linear_start - self.linear_end = model.model.linear_end + self.model_sampling = model.model.model_sampling def use_cached_beta_schedule_and_clean(self, model: ModelPatcher): - model.model.register_schedule(given_betas=self.betas.clone().detach(), linear_start=self.linear_start, linear_end=self.linear_end) + model.model.model_sampling = self.model_sampling self.clean() def clean(self): - self.betas = None - self.linear_start = None - self.linear_end = None + self.model_sampling = None class Folders: diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index 2c080fa..f84b20e 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -98,6 +98,32 @@ def interpolate_pe_to_length(model_dict: dict[str, Tensor], key: str, new_length del temp_pe +def interpolate_pe_to_length_diffs(model_dict: dict[str, Tensor], key: str, new_length: int): + # TODO: fill out and try out + pe_shape = model_dict[key].shape + temp_pe = rearrange(model_dict[key], "(t b) f d -> t b f d", t=1) + temp_pe = F.interpolate(temp_pe, size=(new_length, pe_shape[-1]), mode="bilinear") + temp_pe = rearrange(temp_pe, "t b f d -> (t b) f d", t=1) + model_dict[key] = temp_pe + del temp_pe + + +def interpolate_pe_to_length_pingpong(model_dict: dict[str, Tensor], key: str, new_length: int): + if model_dict[key].shape[1] < new_length: + temp_pe = model_dict[key] + flipped_temp_pe = torch.flip(temp_pe[:, 1:-1, :], [1]) + use_flipped = True + preview_pe = None + while model_dict[key].shape[1] < new_length: + preview_pe = model_dict[key] + model_dict[key] = torch.cat([model_dict[key], flipped_temp_pe if use_flipped else temp_pe], dim=1) + use_flipped = not use_flipped + del temp_pe + del flipped_temp_pe + del preview_pe + model_dict[key] = model_dict[key][:, :new_length] + + 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 @@ -123,9 +149,23 @@ def apply_mm_settings(model_dict: dict[str, Tensor], mm_settings: 'MotionModelSe # apply final_pe_idx_offset, if needed if mm_settings.has_final_pe_idx_offset(): model_dict[key] = model_dict[key][:, mm_settings.final_pe_idx_offset:] - # apply attn_strenth, if needed - elif mm_settings.has_attn_strength(): - model_dict[key] *= mm_settings.attn_strength + else: + # apply attn_strenth, if needed + if mm_settings.has_attn_strength(): + model_dict[key] *= mm_settings.attn_strength + # apply specific attn_strengths, if needed + if mm_settings.has_any_attn_sub_strength(): + if "to_q" in key and mm_settings.has_attn_q_strength(): + model_dict[key] *= mm_settings.attn_q_strength + elif "to_k" in key and mm_settings.has_attn_k_strength(): + model_dict[key] *= mm_settings.attn_k_strength + elif "to_v" in key and mm_settings.has_attn_v_strength(): + model_dict[key] *= mm_settings.attn_v_strength + elif "to_out" in key: + if key.strip().endswith("weight") and mm_settings.has_attn_out_weight_strength(): + model_dict[key] *= mm_settings.attn_out_weight_strength + elif key.strip().endswith("bias") and mm_settings.has_attn_out_bias_strength(): + model_dict[key] *= mm_settings.attn_out_bias_strength # apply other strength, if needed elif mm_settings.has_other_strength(): model_dict[key] *= mm_settings.other_strength @@ -506,16 +546,30 @@ def del_injected_unet_version(model: ModelPatcher): class MotionModelSettings: def __init__(self, - pe_strength: float=1.0, attn_strength: float=1.0, other_strength: float=1.0, + pe_strength: float=1.0, + attn_strength: float=1.0, + attn_q_strength: float=1.0, + attn_k_strength: float=1.0, + attn_v_strength: float=1.0, + attn_out_weight_strength: float=1.0, + attn_out_bias_strength: float=1.0, + other_strength: float=1.0, cap_initial_pe_length: int=0, interpolate_pe_to_length: int=0, initial_pe_idx_offset: int=0, final_pe_idx_offset: int=0, motion_pe_stretch: int=0, attn_scale: float=1.0, ): - # PE-interpolation settings + # general strengths self.pe_strength = pe_strength self.attn_strength = attn_strength self.other_strength = other_strength + # specific attn strengths + self.attn_q_strength = attn_q_strength + self.attn_k_strength = attn_k_strength + self.attn_v_strength = attn_v_strength + self.attn_out_weight_strength = attn_out_weight_strength + self.attn_out_bias_strength = attn_out_bias_strength + # PE-interpolation settings self.cap_initial_pe_length = cap_initial_pe_length self.interpolate_pe_to_length = interpolate_pe_to_length self.initial_pe_idx_offset = initial_pe_idx_offset @@ -556,4 +610,27 @@ class MotionModelSettings: or self.has_interpolate_pe_to_length() \ or self.has_initial_pe_idx_offset() \ or self.has_final_pe_idx_offset() \ - or self.has_motion_pe_stretch() + or self.has_motion_pe_stretch() \ + or self.has_any_attn_sub_strength() + + def has_any_attn_sub_strength(self) -> bool: + return self.has_attn_q_strength() \ + or self.has_attn_k_strength() \ + or self.has_attn_v_strength() \ + or self.has_attn_out_weight_strength() \ + or self.has_attn_out_bias_strength() + + def has_attn_q_strength(self) -> bool: + return self.attn_q_strength != 1.0 + + def has_attn_k_strength(self) -> bool: + return self.attn_k_strength != 1.0 + + def has_attn_v_strength(self) -> bool: + return self.attn_v_strength != 1.0 + + def has_attn_out_weight_strength(self) -> bool: + return self.attn_out_weight_strength != 1.0 + + def has_attn_out_bias_strength(self) -> bool: + return self.attn_out_bias_strength != 1.0 diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 0a20fdf..3e70b15 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -116,6 +116,60 @@ class AnimateDiffModelSettingsAdvanced: 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): @@ -511,6 +565,7 @@ NODE_CLASS_MAPPINGS = { "ADE_AnimateDiffLoRALoader": AnimateDiffLoRALoader, "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, + "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths, "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, "CheckpointLoaderSimpleWithNoiseSelect": CheckpointLoaderSimpleWithNoiseSelect, @@ -524,6 +579,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "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_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓", "CheckpointLoaderSimpleWithNoiseSelect": "Load Checkpoint w/ Noise Select 🎭🅐🅓", diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 0761be0..7342676 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -115,6 +115,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: if not is_injected_mm_params(model): return orig_comfy_sample(model, *args, **kwargs) # otherwise, injection time + motion_module = None + orig_beta_cache = None try: # get params - clone to keep from resetting values on cached model params = get_injected_mm_params(model).clone() @@ -152,9 +154,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: # inject motion module into unet inject_motion_module(model=model, motion_module=motion_module, params=params) - # apply suggested beta schedule - beta_schedule = BetaSchedules.to_name(params.beta_schedule) - model.model.register_schedule(given_betas=None, beta_schedule=beta_schedule, timesteps=1000, linear_start=0.00085, linear_end=0.012, cosine_s=8e-3) + # apply suggested beta schedule (model_sampling) + model.model.model_sampling = BetaSchedules.to_model_sampling(params.beta_schedule, model) # apply scale multiplier, if needed motion_module.set_scale_multiplier(params.motion_model_settings.attn_scale) @@ -178,14 +179,15 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: finally: # attempt to eject motion module eject_motion_module(model=model) - # reset motion module scale multiplier - motion_module.reset_scale_multiplier() - # reset motion module sub_idxs - motion_module.set_sub_idxs(None) - # if loras are present, remove model so it can be re-loaded next time with fresh weights - if motion_module.has_loras(): - unload_motion_module(motion_module) - del motion_module + 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) + # 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) + del motion_module ############################################## # Restoration model_management.maximum_batch_area = orig_maximum_batch_area @@ -194,7 +196,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable: GroupNormAD.forward = orig_groupnormad_forward comfy_samplers.sampling_function = orig_sampling_function # reapply previous beta schedule - orig_beta_cache.use_cached_beta_schedule_and_clean(model) + if orig_beta_cache is not None: + orig_beta_cache.use_cached_beta_schedule_and_clean(model) # reset global state ADGS.reset() ##############################################