diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index 1ca129d..e500be9 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, ModelSamplingDiscrete, ModelType, 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/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() ##############################################