Merge fix for recent ComfyUI updates

Merge fix for recent ComfyUI updates
This commit is contained in:
Jedrzej Kosinski
2023-11-01 01:14:33 -05:00
committed by GitHub
4 changed files with 171 additions and 26 deletions
+17 -8
View File
@@ -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:
+83 -6
View File
@@ -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
+56
View File
@@ -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
View File
@@ -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()
##############################################