diff --git a/comfy/sd.py b/comfy/sd.py index 95fc6d271..43f07e1d3 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -682,6 +682,12 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c if output_model: model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=model_management.unet_offload_device()) + if model_config.ztsnr: + class ModelSamplingAdvanced(comfy.model_sampling.ModelSamplingDiscrete, comfy.model_sampling.V_PREDICTION): + pass + model_sampling = ModelSamplingAdvanced(model.model_config) + model_sampling.set_sigmas(comfy.utils.rescale_zero_terminal_snr_sigmas(model_sampling.sigmas)) + model_patcher.add_object_patch("model_sampling", model_sampling) if inital_load_device != torch.device("cpu"): logging.info("loaded straight to GPU") model_management.load_models_gpu([model_patcher], force_full_load=True) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index 9931f4c5d..3a9674e63 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -197,6 +197,8 @@ class SDXL(supported_models_base.BASE): self.sampling_settings["sigma_min"] = float(state_dict["edm_vpred.sigma_min"].item()) return model_base.ModelType.V_PREDICTION_EDM elif "v_pred" in state_dict: + if "ztsnr" in state_dict: + self.ztsnr = True return model_base.ModelType.V_PREDICTION else: return model_base.ModelType.EPS diff --git a/comfy/supported_models_base.py b/comfy/supported_models_base.py index 54573abb1..5b1ac5d12 100644 --- a/comfy/supported_models_base.py +++ b/comfy/supported_models_base.py @@ -40,6 +40,7 @@ class BASE: clip_vision_prefix = None noise_aug_config = None sampling_settings = {} + ztsnr = False latent_format = latent_formats.LatentFormat vae_key_prefix = ["first_stage_model."] text_encoder_key_prefix = ["cond_stage_model."] diff --git a/comfy/utils.py b/comfy/utils.py index 3c5d06a4f..ae48940a2 100644 --- a/comfy/utils.py +++ b/comfy/utils.py @@ -869,3 +869,22 @@ def reshape_mask(input_mask, output_shape): mask = mask.repeat((1, output_shape[1]) + (1,) * dims)[:,:output_shape[1]] mask = comfy.utils.repeat_to_batch_size(mask, output_shape[0]) return mask + +def rescale_zero_terminal_snr_sigmas(sigmas): + alphas_cumprod = 1 / ((sigmas * sigmas) + 1) + alphas_bar_sqrt = alphas_cumprod.sqrt() + + # Store old values. + alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone() + alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone() + + # Shift so the last timestep is zero. + alphas_bar_sqrt -= (alphas_bar_sqrt_T) + + # Scale so the first timestep is back to the old value. + alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T) + + # Convert alphas_bar_sqrt to betas + alphas_bar = alphas_bar_sqrt**2 # Revert sqrt + alphas_bar[-1] = 4.8973451890853435e-08 + return ((1 - alphas_bar) / alphas_bar) ** 0.5 diff --git a/comfy_extras/nodes_model_advanced.py b/comfy_extras/nodes_model_advanced.py index 918e6085a..c7c864708 100644 --- a/comfy_extras/nodes_model_advanced.py +++ b/comfy_extras/nodes_model_advanced.py @@ -5,6 +5,8 @@ import comfy.latent_formats import nodes import torch +from comfy.utils import rescale_zero_terminal_snr_sigmas + class LCM(comfy.model_sampling.EPS): def calculate_denoised(self, sigma, model_output, model_input): timestep = self.timestep(sigma).view(sigma.shape[:1] + (1,) * (model_output.ndim - 1)) @@ -51,25 +53,6 @@ class ModelSamplingDiscreteDistilled(comfy.model_sampling.ModelSamplingDiscrete) return log_sigma.exp().to(timestep.device) -def rescale_zero_terminal_snr_sigmas(sigmas): - alphas_cumprod = 1 / ((sigmas * sigmas) + 1) - alphas_bar_sqrt = alphas_cumprod.sqrt() - - # Store old values. - alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone() - alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone() - - # Shift so the last timestep is zero. - alphas_bar_sqrt -= (alphas_bar_sqrt_T) - - # Scale so the first timestep is back to the old value. - alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T) - - # Convert alphas_bar_sqrt to betas - alphas_bar = alphas_bar_sqrt**2 # Revert sqrt - alphas_bar[-1] = 4.8973451890853435e-08 - return ((1 - alphas_bar) / alphas_bar) ** 0.5 - class ModelSamplingDiscrete: @classmethod def INPUT_TYPES(s):