From b7104b0da2aae600182830ae70b216f0ce87fe7f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 17:32:44 -0500 Subject: [PATCH] Refactored custom_cfg code to no longer use sampler_cfg_function patch, allowing that patch to not be overridden when using Custom CFG (some stuff might fail if it expects cond_scale to be a float instead of a tensor) --- animatediff/sample_settings.py | 9 +++++++++ animatediff/sampling.py | 4 ++-- 2 files changed, 11 insertions(+), 2 deletions(-) diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index 71086f4..a2c0764 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -583,7 +583,16 @@ class CustomCFGKeyframeGroup: # update steps current context is used self._current_used_steps += 1 + def get_cfg_scale(self, cond: Tensor): + cond_scale = self.cfg_multival + if isinstance(cond_scale, Tensor): + cond_scale = prepare_mask_batch(cond_scale.to(cond.dtype).to(cond.device), cond.shape) + cond_scale = extend_to_batch_size(cond_scale, cond.shape[0]) + return cond_scale + def patch_model(self, model: ModelPatcher) -> ModelPatcher: + # NOTE: no longer used at the moment, as most sampler_cfg_function patches should work with tensor cfg_scales, + # meaning get_cfg_scale is a direct replacement def evolved_custom_cfg(args): cond: Tensor = args["cond"] uncond: Tensor = args["uncond"] diff --git a/animatediff/sampling.py b/animatediff/sampling.py index d6be33a..ee573d4 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -414,8 +414,6 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> cached_noise = None function_injections = FunctionInjectionHolder() try: - if model.sample_settings.custom_cfg is not None: - model = model.sample_settings.custom_cfg.patch_model(model) # clone params from model params = model.motion_injection_params.clone() # get amount of latents passed in, and store in params @@ -629,6 +627,8 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, cond_pred, uncond_pred = sliding_calc_conds_batch(model, [cond, uncond_], x, timestep, model_options) if hasattr(comfy.samplers, "cfg_function"): + if ADGS.sample_settings.custom_cfg is not None: + cond_scale = ADGS.sample_settings.custom_cfg.get_cfg_scale(cond_pred) try: cached_calc_cond_batch = comfy.samplers.calc_cond_batch # support hooks and sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch