Finished adding support for AnimateLCM model, added lcm beta_schedules

This commit is contained in:
Jedrzej Kosinski
2024-02-04 17:35:42 -06:00
parent 3f8202a382
commit 5540ee6314
3 changed files with 78 additions and 14 deletions
+8 -3
View File
@@ -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(),
+6 -3
View File
@@ -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
+64 -8
View File
@@ -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