Merge fix for recent ComfyUI updates
Merge fix for recent ComfyUI updates
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 🎭🅐🅓",
|
||||
|
||||
+15
-12
@@ -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()
|
||||
##############################################
|
||||
|
||||
Reference in New Issue
Block a user