Refactor to match 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, 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:
|
||||
|
||||
+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