diff --git a/nodes.py b/nodes.py index 9059fde..c71e61d 100644 --- a/nodes.py +++ b/nodes.py @@ -5,11 +5,13 @@ import math original_sampling_function = deepcopy(comfy.samplers.sampling_function) minimum_sigma_to_disable_uncond = 1 +no_uncond_at_all = False def sampling_function_patched(model, x, timestep, uncond, cond, cond_scale, model_options={}, seed=None): - if math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False or timestep[0] <= minimum_sigma_to_disable_uncond: + if math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False or timestep[0] <= minimum_sigma_to_disable_uncond or no_uncond_at_all: uncond_ = None - cond_scale = 1 + if not no_uncond_at_all: + cond_scale = 1 else: uncond_ = uncond @@ -54,13 +56,16 @@ class advancedDynamicCFG: return {"required": { "model": ("MODEL",), "center_mean_post_cfg" : ("BOOLEAN", {"default": True}), - "center_mean_to_sigma" : ("BOOLEAN", {"default": True}), - "automatic_cfg" : (["None","soft","hard","include_boost"], {"default": "hard"},), + "center_mean_to_sigma" : ("BOOLEAN", {"default": False}), + "automatic_cfg" : (["None","soft","hard","progressive","include_boost"], {"default": "hard"},), "sigma_boost" : ("BOOLEAN", {"default": True}), "sigma_boost_percentage": ("FLOAT", {"default": 6.86, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01}), "lerp_uncond" : ("BOOLEAN", {"default": False}), - "lerp_uncond_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 10.0, "step": 0.1, "round": 0.1}), - # "debug_print" : ("BOOLEAN", {"default": False}), + "lerp_uncond_strength": ("FLOAT", {"default": 1, "min": 0.0, "max": 10.0, "step": 0.1, "round": 0.1}), + "post_cfg_scale" : ("BOOLEAN", {"default": False}), + "post_cfg_scale_value": ("FLOAT", {"default": 0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.1}), + "no_uncond_mode" : ("BOOLEAN", {"default": False}), + "debug_print" : ("BOOLEAN", {"default": False}), }} RETURN_TYPES = ("MODEL",) FUNCTION = "patch" @@ -68,15 +73,17 @@ class advancedDynamicCFG: CATEGORY = "model_patches" def patch(self, model, center_mean_post_cfg, center_mean_to_sigma, - automatic_cfg, sigma_boost, sigma_boost_percentage, lerp_uncond=False, lerp_uncond_strength=1, debug_print=False): + automatic_cfg, sigma_boost, sigma_boost_percentage, lerp_uncond=False, lerp_uncond_strength=1, + post_cfg_scale=False, post_cfg_scale_value=8, no_uncond_mode=False, debug_print=False): - global minimum_sigma_to_disable_uncond + global minimum_sigma_to_disable_uncond, no_uncond_at_all + no_uncond_at_all = no_uncond_mode model_sampling = model.model.model_sampling sigmin = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_min)) sigmax = model_sampling.sigma(model_sampling.timestep(model_sampling.sigma_max)) - + low_sigma_threshold = (sigmax - sigmin) / 100 * sigma_boost_percentage if sigma_boost_percentage > 0 and sigma_boost: - minimum_sigma_to_disable_uncond = (sigmax - sigmin) / 100 * sigma_boost_percentage + minimum_sigma_to_disable_uncond = low_sigma_threshold comfy.samplers.sampling_function = sampling_function_patched print(f"Sampling function patched. Trigger when sigmas are at: {round(minimum_sigma_to_disable_uncond.item(),4)}") else: @@ -90,13 +97,19 @@ class advancedDynamicCFG: input_x = args["input"] cond_pred = args["cond_denoised"] uncond_pred = args["uncond_denoised"] + sigma = args["sigma"][0] + if lerp_uncond: - uncond_pred = torch.lerp(cond_pred,uncond_pred,lerp_uncond_strength) - # uncond_pred = uncond_pred * cond_pred.norm() / uncond_pred.norm() + lerp_weight = lerp_uncond_strength if lerp_uncond_strength > 0 else max(sigma.item(), 1) + if lerp_weight != 1: + uncond_pred = torch.lerp(cond_pred, uncond_pred, lerp_weight) cond = input_x - cond_pred uncond = input_x - uncond_pred - sigma = args["sigma"][0] + if no_uncond_mode: + self.last_cfg_ht_one = cond_scale + return cond + if sigma == sigmax or cond_scale > 1: self.last_cfg_ht_one = cond_scale @@ -126,13 +139,18 @@ class advancedDynamicCFG: min_val = abs(torch.mean(min_values).item()) elif automatic_cfg == "hard" or automatic_cfg == "include_boost": min_val = torch.mean(torch.abs(min_values)).item() + elif automatic_cfg == "progressive": + min_val = torch.mean(torch.abs(min_values)).item() + s_progression = map_sigma(sigma, sigmax, sigmin) + target_intensity = 1.1 * target_intensity * s_progression + 0.9 * target_intensity * (1 - s_progression) denoised_range = (max_val + min_val) / 2 - scale_correction = target_intensity / denoised_range + scale_correction = target_intensity / denoised_range tmp_scale = reference_cfg * scale_correction if debug_print: print(f"c{c}: {tmp_scale} / {scale_correction}") + print(f"denoised_range: {denoised_range}") if cond_scale > 1: denoised_tmp[b][c] = uncond[b][c] + tmp_scale * (cond[b][c] - uncond[b][c]) @@ -152,11 +170,33 @@ class advancedDynamicCFG: denoised = center_latent_mean_values(denoised, False, mult) return denoised + def rescale_post_cfg(args): + denoised = args["denoised"] + sigma = args["sigma"][0] + if sigma <= minimum_sigma_to_disable_uncond: + return denoised + for b in range(len(denoised)): + for c in range(len(denoised[b])): #TODO make a function for the scaling + channel = denoised[b][c] + max_values = torch.topk(channel, k=int(len(channel)*top_k), largest=True ).values + min_values = torch.topk(channel, k=int(len(channel)*top_k), largest=False).values + max_val = torch.mean(max_values).item() + min_val = torch.mean(torch.abs(min_values)).item() + denoised_range = (max_val + min_val) / 2 + if no_uncond_mode or post_cfg_scale_value == 0: + target_intensity = self.last_cfg_ht_one / 10 + else: + target_intensity = post_cfg_scale_value / 10 + scale_correction = target_intensity / denoised_range + denoised[b][c] = channel * scale_correction + return denoised + m = model.clone() m.set_model_sampler_cfg_function(linear_cfg, disable_cfg1_optimization=False) - if center_mean_post_cfg: + if center_mean_post_cfg or no_uncond_mode: m.set_model_sampler_post_cfg_function(center_mean_latent_post_cfg) - + if post_cfg_scale or no_uncond_mode: + m.set_model_sampler_post_cfg_function(rescale_post_cfg) return (m, ) class simpleDynamicCFG: @@ -182,7 +222,7 @@ class simpleDynamicCFGlerpUncond: return {"required": { "model": ("MODEL",), "boost" : ("BOOLEAN", {"default": True}), - "negative_strength": ("FLOAT", {"default": 1, "min": 0.0, "max": 10.0, "step": 0.1, "round": 0.1}), + "negative_strength": ("FLOAT", {"default": 1, "min": 0.0, "max": 5.0, "step": 0.1, "round": 0.1}), }} RETURN_TYPES = ("MODEL",) FUNCTION = "patch" @@ -191,5 +231,26 @@ class simpleDynamicCFGlerpUncond: def patch(self, model, boost, negative_strength): advcfg = advancedDynamicCFG() - m = advcfg.patch(model, False, False, "hard", boost, 6.86, negative_strength != 1, negative_strength / 2)[0] + # automatic_cfg="progressive" if negative_strength == 1 else "hard" + m = advcfg.patch(model=model, center_mean_post_cfg=False, center_mean_to_sigma=False, + automatic_cfg="hard", sigma_boost=boost, sigma_boost_percentage=6.86, + lerp_uncond=negative_strength != 1, lerp_uncond_strength=negative_strength)[0] return (m, ) + +class simpleDynamicCFGNoUncond: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ("MODEL",), + }} + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "model_patches" + + def patch(self, model): + advcfg = advancedDynamicCFG() + m = advcfg.patch(model=model, center_mean_post_cfg=True, center_mean_to_sigma=False, + automatic_cfg="None", sigma_boost="None", sigma_boost_percentage=6.86, + no_uncond_mode=True)[0] + return (m, ) \ No newline at end of file