Refactor to match recent ComfyUI updates

This commit is contained in:
Jedrzej Kosinski
2023-11-01 01:07:27 -05:00
parent a65083727c
commit 0fe243dfd4
2 changed files with 32 additions and 20 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, 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
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()
##############################################