From 5540ee6314e40d5d34dc540f396fe16f798fc704 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 4 Feb 2024 17:35:42 -0600 Subject: [PATCH] Finished adding support for AnimateLCM model, added lcm beta_schedules --- animatediff/model_injection.py | 11 +++-- animatediff/motion_module_ad.py | 9 +++-- animatediff/utils_model.py | 72 +++++++++++++++++++++++++++++---- 3 files changed, 78 insertions(+), 14 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 7da63f8..3ad7a61 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -13,7 +13,7 @@ from comfy.model_base import BaseModel from .ad_settings import AnimateDiffSettings from .context import ContextOptions, ContextOptions, ContextOptionsGroup -from .motion_module_ad import AnimateDiffModel, has_mid_block, normalize_ad_state_dict +from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, has_mid_block, normalize_ad_state_dict from .logger import logger from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max from .motion_lora import MotionLoraInfo, MotionLoraList @@ -365,7 +365,8 @@ def load_motion_module_gen1(model_name: str, model: ModelPatcher, motion_lora: M ad_wrapper = AnimateDiffModel(mm_state_dict=mm_state_dict, mm_info=mm_info) ad_wrapper.to(model.model_dtype()) ad_wrapper.to(model.offload_device) - load_result = ad_wrapper.load_state_dict(mm_state_dict) + is_animatelcm = mm_info.mm_format==AnimateDiffFormat.ANIMATELCM + load_result = ad_wrapper.load_state_dict(mm_state_dict, strict=not is_animatelcm) # TODO: report load_result of motion_module loading? # wrap motion_module into a ModelPatcher, to allow motion lora patches motion_model = MotionModelPatcher(model=ad_wrapper, load_device=model.load_device, offload_device=model.offload_device) @@ -389,7 +390,11 @@ def load_motion_module_gen2(model_name: str, motion_model_settings: AnimateDiffS ad_wrapper = AnimateDiffModel(mm_state_dict=mm_state_dict, mm_info=mm_info) ad_wrapper.to(comfy.model_management.unet_dtype()) ad_wrapper.to(comfy.model_management.unet_offload_device()) - load_result = ad_wrapper.load_state_dict(mm_state_dict) + is_animatelcm = mm_info.mm_format==AnimateDiffFormat.ANIMATELCM + load_result = ad_wrapper.load_state_dict(mm_state_dict, strict=not is_animatelcm) + # TODO: manually check load_results for AnimateLCM models + if is_animatelcm: + pass # TODO: report load_result of motion_module loading? # wrap motion_module into a ModelPatcher, to allow motion lora patches motion_model = MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(), diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index a29fdc6..8d582c8 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -95,9 +95,9 @@ def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str for key in mm_state_dict.keys(): if key.endswith("pos_encoder.pe"): return mm_state_dict[key].size(1) # get middle dim - # AnimateLCM models should have no pos_encoder entries + # AnimateLCM models should have no pos_encoder entries, and assumed to be 64 if mm_format == AnimateDiffFormat.ANIMATELCM: - return None + return 64 raise MotionCompatibilityError(f"No pos_encoder.pe found in mm_state_dict - {mm_name} is not a valid AnimateDiff motion module!") @@ -211,7 +211,10 @@ class AnimateDiffModel(nn.Module): def get_best_beta_schedule(self, log=False) -> str: to_return = None if self.mm_info.sd_type == ModelTypeSD.SD1_5: - to_return = BetaSchedules.SQRT_LINEAR + if self.mm_info.mm_format == AnimateDiffFormat.ANIMATELCM: + to_return = BetaSchedules.LCM # while LCM_100 is the intended schedule, I find LCM to have much less flicker + else: + to_return = BetaSchedules.SQRT_LINEAR elif self.mm_info.sd_type == ModelTypeSD.SDXL: if self.mm_info.mm_format == AnimateDiffFormat.HOTSHOTXL: to_return = BetaSchedules.LINEAR diff --git a/animatediff/utils_model.py b/animatediff/utils_model.py index 8fa3aac..34fe71d 100644 --- a/animatediff/utils_model.py +++ b/animatediff/utils_model.py @@ -11,6 +11,9 @@ from comfy.model_base import SD21UNCLIP, SDXL, BaseModel, SDXLRefiner, SVD_img2v from comfy.model_management import xformers_enabled from comfy.model_patcher import ModelPatcher +import comfy.model_sampling +import comfy_extras.nodes_model_advanced + class IsChangedHelper: def __init__(self): @@ -24,49 +27,102 @@ class IsChangedHelper: class ModelSamplingConfig: - def __init__(self, beta_schedule: str): + def __init__(self, beta_schedule: str, linear_start: float=None, linear_end: float=None): self.sampling_settings = {"beta_schedule": beta_schedule} + if linear_start is not None: + self.sampling_settings["linear_start"] = linear_start + if linear_end is not None: + self.sampling_settings["linear_end"] = linear_end self.beta_schedule = beta_schedule # keeping this for backwards compatibility +class ModelSamplingType: + EPS = "eps" + V_PREDICTION = "v_prediction" + LCM = "lcm" + + +def factory_model_sampling_discrete_distilled(original_timesteps=50): + class ModelSamplingDiscreteDistilledEvolved(comfy_extras.nodes_model_advanced.ModelSamplingDiscreteDistilled): + def __init__(self, *args, **kwargs): + self.original_timesteps = original_timesteps # normal LCM has 50 + super().__init__(*args, **kwargs) + return ModelSamplingDiscreteDistilledEvolved + + +# based on code in comfy_extras/nodes_model_advanced.py +def evolved_model_sampling(model_config: ModelSamplingConfig, model_type: int, alias: str): + # if LCM, need to handle manually + if BetaSchedules.is_lcm(alias): + sampling_type = comfy_extras.nodes_model_advanced.LCM + if alias == BetaSchedules.LCM_100: + sampling_base = factory_model_sampling_discrete_distilled(original_timesteps=100) + elif alias == BetaSchedules.LCM_25: + sampling_base = factory_model_sampling_discrete_distilled(original_timesteps=25) + else: + sampling_base = comfy_extras.nodes_model_advanced.ModelSamplingDiscreteDistilled + class ModelSamplingAdvancedEvolved(sampling_base, sampling_type): + pass + # NOTE: if I want to support zsnr, this is where I would add that code + return ModelSamplingAdvancedEvolved(model_config) + # otherwise, use vanilla model_sampling function + return model_sampling(model_config, model_type) + + class BetaSchedules: AUTOSELECT = "autoselect" SQRT_LINEAR = "sqrt_linear (AnimateDiff)" LINEAR_ADXL = "linear (AnimateDiff-SDXL)" LINEAR = "linear (HotshotXL/default)" + LCM = "lcm" + LCM_100 = "lcm[100_ots]" + LCM_25 = "lcm[25_ots]" + LCM_SQRT_LINEAR = "lcm >> sqrt_linear" USE_EXISTING = "use existing" SQRT = "sqrt" COSINE = "cosine" SQUAREDCOS_CAP_V2 = "squaredcos_cap_v2" - ALIAS_LIST = [AUTOSELECT, SQRT_LINEAR, LINEAR_ADXL, LINEAR, - USE_EXISTING, SQRT, COSINE, SQUAREDCOS_CAP_V2] + LCM_LIST = [LCM, LCM_100, LCM_25, LCM_SQRT_LINEAR] + + ALIAS_LIST = [AUTOSELECT, USE_EXISTING, SQRT_LINEAR, LINEAR_ADXL, LINEAR, LCM, LCM_100, LCM_SQRT_LINEAR, # LCM_25 is purposely omitted + SQRT, COSINE, SQUAREDCOS_CAP_V2] ALIAS_MAP = { SQRT_LINEAR: "sqrt_linear", LINEAR_ADXL: "linear", # also linear, but has different linear_end (0.020) LINEAR: "linear", + LCM_100: "linear", # distilled, 100 original timesteps + LCM_25: "linear", # distilled, 25 original timesteps + LCM: "linear", # distilled + LCM_SQRT_LINEAR: "sqrt_linear", # distilled, sqrt_linear SQRT: "sqrt", COSINE: "cosine", SQUAREDCOS_CAP_V2: "squaredcos_cap_v2", } + @classmethod + def is_lcm(cls, alias: str): + return alias in cls.LCM_LIST + @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)) + linear_start = None + linear_end = None + if alias == cls.LINEAR_ADXL: + # uses linear_end=0.020 + linear_end = 0.020 + return ModelSamplingConfig(cls.to_name(alias), linear_start=linear_start, linear_end=linear_end) @classmethod def to_model_sampling(cls, alias: str, model: ModelPatcher): if alias == cls.USE_EXISTING: return None - ms_obj = model_sampling(cls.to_config(alias), model_type=model.model.model_type) - if alias == cls.LINEAR_ADXL: - # uses linear_end=0.020 - ms_obj._register_schedule(given_betas=None, beta_schedule=cls.to_name(alias), timesteps=1000, linear_start=0.00085, linear_end=0.020, cosine_s=8e-3) + ms_obj = evolved_model_sampling(cls.to_config(alias), model_type=model.model.model_type, alias=alias) return ms_obj @staticmethod