From ae7ae379cdbb4de4c9dab45c5d657a7d0554f3ff Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 23 Aug 2024 04:43:47 -0500 Subject: [PATCH] In Create Raw Sigma Schedule, renamed lcm_zsnr to zsnr and changed logic to allow applying zsnr for non-lcm --- animatediff/nodes_sigma_schedule.py | 9 +++++---- animatediff/utils_model.py | 2 +- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/animatediff/nodes_sigma_schedule.py b/animatediff/nodes_sigma_schedule.py index 828e2ca..a27b36c 100644 --- a/animatediff/nodes_sigma_schedule.py +++ b/animatediff/nodes_sigma_schedule.py @@ -44,7 +44,7 @@ class RawSigmaScheduleNode: #"cosine_s": ("FLOAT", {"default": 8e-3, "min": 0.0, "max": 1.0, "step": 0.000001}), "sampling": (ModelSamplingType._FULL_LIST,), "lcm_original_timesteps": ("INT", {"default": 50, "min": 1, "max": 1000}), - "lcm_zsnr": ("BOOLEAN", {"default": False}), + "zsnr": ("BOOLEAN", {"default": False}), } } @@ -53,14 +53,15 @@ class RawSigmaScheduleNode: FUNCTION = "get_sigma_schedule" def get_sigma_schedule(self, raw_beta_schedule: str, linear_start: float, linear_end: float,# cosine_s: float, - sampling: str, lcm_original_timesteps: int, lcm_zsnr: bool): + sampling: str, lcm_original_timesteps: int, zsnr: bool, lcm_zsnr: bool=None): + if lcm_zsnr is not None: + zsnr = lcm_zsnr new_config = ModelSamplingConfig(beta_schedule=raw_beta_schedule, linear_start=linear_start, linear_end=linear_end) if sampling != ModelSamplingType.LCM: lcm_original_timesteps=None - lcm_zsnr=False model_type = ModelSamplingType.from_alias(sampling) new_model_sampling = BetaSchedules._to_model_sampling(alias=BetaSchedules.AUTOSELECT, model_type=model_type, config_override=new_config, original_timesteps=lcm_original_timesteps) - if lcm_zsnr: + if zsnr: SigmaSchedule.apply_zsnr(new_model_sampling=new_model_sampling) return (SigmaSchedule(model_sampling=new_model_sampling, model_type=model_type),) diff --git a/animatediff/utils_model.py b/animatediff/utils_model.py index 48c432b..930cdc9 100644 --- a/animatediff/utils_model.py +++ b/animatediff/utils_model.py @@ -194,7 +194,7 @@ class BetaSchedules: return ModelSamplingConfig(cls.to_name(alias), linear_start=linear_start, linear_end=linear_end) @classmethod - def _to_model_sampling(cls, alias: str, model_type: ModelType, config_override: ModelSamplingConfig=None, original_timesteps: int=None): + def _to_model_sampling(cls, alias: str, model_type: ModelType, config_override: Union[ModelSamplingConfig,None]=None, original_timesteps: Union[int,None]=None): if alias == cls.USE_EXISTING: return None elif config_override != None: