From bb0aefd6936a134d86a6ffd92424134c549c6b75 Mon Sep 17 00:00:00 2001 From: Haoming Date: Wed, 8 Oct 2025 13:58:31 +0800 Subject: [PATCH] uni_pc --- modules/sd_samplers_extra.py | 54 ++++ modules/sd_samplers_kdiffusion.py | 1 + modules/sd_samplers_timesteps.py | 1 - modules/sd_samplers_timesteps_impl.py | 38 --- modules/shared_options.py | 4 - modules/uni_pc/__init__.py | 1 - modules/uni_pc/sampler.py | 101 ------- modules/uni_pc/uni_pc.py | 409 +++++--------------------- 8 files changed, 134 insertions(+), 475 deletions(-) delete mode 100644 modules/uni_pc/__init__.py delete mode 100644 modules/uni_pc/sampler.py diff --git a/modules/sd_samplers_extra.py b/modules/sd_samplers_extra.py index ce9bbc19..50e4a212 100644 --- a/modules/sd_samplers_extra.py +++ b/modules/sd_samplers_extra.py @@ -2,6 +2,8 @@ import k_diffusion.sampling import torch import tqdm +from modules.uni_pc import uni_pc as unipc_impl + @torch.no_grad() def restart_sampler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1.0, restart_list=None): @@ -73,3 +75,55 @@ def restart_sampler(model, x, sigmas, extra_args=None, callback=None, disable=No last_sigma = new_sigma return x + + +class SigmaConvert: + schedule = "" + + def marginal_log_mean_coeff(self, sigma): + return 0.5 * torch.log(1 / ((sigma * sigma) + 1)) + + def marginal_alpha(self, t): + return torch.exp(self.marginal_log_mean_coeff(t)) + + def marginal_std(self, t): + return torch.sqrt(1.0 - torch.exp(2.0 * self.marginal_log_mean_coeff(t))) + + def marginal_lambda(self, t): + """ + Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T] + """ + log_mean_coeff = self.marginal_log_mean_coeff(t) + log_std = 0.5 * torch.log(1.0 - torch.exp(2.0 * log_mean_coeff)) + return log_mean_coeff - log_std + + +def predict_eps_sigma(model, input, sigma_in, **kwargs): + sigma = sigma_in.view(sigma_in.shape[:1] + (1,) * (input.ndim - 1)) + input = input * ((sigma**2 + 1.0) ** 0.5) + return (input - model(input, sigma_in, **kwargs)) / sigma + + +def sample_unipc(model, x, sigmas, extra_args=None, callback=None, disable=False, variant="bh1"): + timesteps = sigmas.clone() + if sigmas[-1] < 0.001: + timesteps[-1] = 0.001 + + ns = SigmaConvert() + + x = x / torch.sqrt(1.0 + timesteps[0] ** 2.0) + model_type = "noise" + + model_fn = unipc_impl.model_wrapper( + lambda input, sigma, **kwargs: predict_eps_sigma(model, input, sigma, **kwargs), + ns, + model_type=model_type, + guidance_type="uncond", + model_kwargs=extra_args, + ) + + order = min(3, len(timesteps) - 2) + uni_pc = unipc_impl.UniPC(model_fn, ns, predict_x0=True, thresholding=False, variant=variant) + x = uni_pc.sample(x, timesteps=timesteps, skip_type="time_uniform", method="multistep", order=order, lower_order_final=True, callback=callback, disable_pbar=disable) + x /= ns.marginal_alpha(timesteps[-1]) + return x diff --git a/modules/sd_samplers_kdiffusion.py b/modules/sd_samplers_kdiffusion.py index fd4b6657..70faaf31 100644 --- a/modules/sd_samplers_kdiffusion.py +++ b/modules/sd_samplers_kdiffusion.py @@ -23,6 +23,7 @@ samplers_k_diffusion = [ ("Heun", "sample_heun", ["k_heun"], {"second_order": True}), ("DPM2", "sample_dpm_2", ["k_dpm_2"], {"scheduler": "karras", "discard_next_to_last_sigma": True, "second_order": True}), ("Restart", sd_samplers_extra.restart_sampler, ["restart"], {"scheduler": "karras", "second_order": True}), + ("UniPC", sd_samplers_extra.sample_unipc, ["unipc"], {"discard_next_to_last_sigma": True}), ("DDPM", "sample_ddpm", ["ddpm"], {}), ] diff --git a/modules/sd_samplers_timesteps.py b/modules/sd_samplers_timesteps.py index b826d294..c5f89371 100644 --- a/modules/sd_samplers_timesteps.py +++ b/modules/sd_samplers_timesteps.py @@ -13,7 +13,6 @@ from modules.shared import opts samplers_timesteps = [ ("DDIM", sd_samplers_timesteps_impl.ddim, ["ddim"], {}), ("PLMS", sd_samplers_timesteps_impl.plms, ["plms"], {}), - ("UniPC", sd_samplers_timesteps_impl.unipc, ["unipc"], {}), ] diff --git a/modules/sd_samplers_timesteps_impl.py b/modules/sd_samplers_timesteps_impl.py index 5e2c5347..edee7ce2 100644 --- a/modules/sd_samplers_timesteps_impl.py +++ b/modules/sd_samplers_timesteps_impl.py @@ -3,9 +3,7 @@ import numpy as np import torch import tqdm -from modules import shared from modules.torch_utils import float64 -from modules.uni_pc import uni_pc @torch.no_grad() @@ -105,39 +103,3 @@ def plms(model, x, timesteps, extra_args=None, callback=None, disable=None): callback({"x": x, "i": i, "sigma": 0, "sigma_hat": 0, "denoised": pred_x0}) return x - - -class UniPCCFG(uni_pc.UniPC): - def __init__(self, cfg_model, extra_args, callback, *args, **kwargs): - super().__init__(None, *args, **kwargs) - - def after_update(x, model_x): - callback({"x": x, "i": self.index, "sigma": 0, "sigma_hat": 0, "denoised": model_x}) - self.index += 1 - - self.cfg_model = cfg_model - self.extra_args = extra_args - self.callback = callback - self.index = 0 - self.after_update = after_update - - def get_model_input_time(self, t_continuous): - return (t_continuous - 1.0 / self.noise_schedule.total_N) * 1000.0 - - def model(self, x, t): - t_input = self.get_model_input_time(t) - - res = self.cfg_model(x, t_input, **self.extra_args) - - return res - - -def unipc(model, x, timesteps, extra_args=None, callback=None, disable=None, is_img2img=False): - alphas_cumprod = model.inner_model.inner_model.alphas_cumprod - - ns = uni_pc.NoiseScheduleVP("discrete", alphas_cumprod=alphas_cumprod) - t_start = timesteps[-1] / 1000 + 1 / 1000 if is_img2img else None # this is likely off by a bit - if someone wants to fix it please by all means - unipc_sampler = UniPCCFG(model, extra_args, callback, ns, predict_x0=True, thresholding=False, variant=shared.opts.uni_pc_variant) - x = unipc_sampler.sample(x, steps=len(timesteps), t_start=t_start, skip_type=shared.opts.uni_pc_skip_type, method="multistep", order=shared.opts.uni_pc_order, lower_order_final=shared.opts.uni_pc_lower_order_final) - - return x diff --git a/modules/shared_options.py b/modules/shared_options.py index 4039e74d..568b0aa2 100644 --- a/modules/shared_options.py +++ b/modules/shared_options.py @@ -507,10 +507,6 @@ options_templates.update( "eta_noise_seed_delta": OptionInfo(0, "Eta noise seed delta", gr.Number, {"precision": 0}, infotext="ENSD"), "always_discard_next_to_last_sigma": OptionInfo(False, "Always discard next-to-last sigma", infotext="Discard penultimate sigma"), "sgm_noise_multiplier": OptionInfo(False, "SGM noise multiplier", infotext="SGM noise multiplier"), - "uni_pc_variant": OptionInfo("bh1", "UniPC variant", gr.Radio, {"choices": ("bh1", "bh2", "vary_coeff")}, infotext="UniPC variant"), - "uni_pc_skip_type": OptionInfo("time_uniform", "UniPC skip type", gr.Radio, {"choices": ("time_uniform", "time_quadratic", "logSNR")}, infotext="UniPC skip type"), - "uni_pc_order": OptionInfo(3, "UniPC order", gr.Slider, {"minimum": 1, "maximum": 50, "step": 1}, infotext="UniPC order"), - "uni_pc_lower_order_final": OptionInfo(True, "UniPC lower order final", infotext="UniPC lower order final"), "sd_noise_schedule": OptionInfo("Default", "Noise schedule for sampling", gr.Radio, {"choices": ("Default", "Zero Terminal SNR")}, infotext="Noise Schedule"), "beta_dist_alpha": OptionInfo(0.6, "Beta scheduler - alpha", gr.Slider, {"minimum": 0.01, "maximum": 2.0, "step": 0.01}, infotext="Beta scheduler alpha"), "beta_dist_beta": OptionInfo(0.6, "Beta scheduler - beta", gr.Slider, {"minimum": 0.01, "maximum": 2.0, "step": 0.01}, infotext="Beta scheduler beta"), diff --git a/modules/uni_pc/__init__.py b/modules/uni_pc/__init__.py deleted file mode 100644 index dbb35964..00000000 --- a/modules/uni_pc/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .sampler import UniPCSampler # noqa: F401 diff --git a/modules/uni_pc/sampler.py b/modules/uni_pc/sampler.py deleted file mode 100644 index 168700f7..00000000 --- a/modules/uni_pc/sampler.py +++ /dev/null @@ -1,101 +0,0 @@ -import torch - -from modules import devices, shared - -from .uni_pc import NoiseScheduleVP, UniPC, model_wrapper - - -class UniPCSampler(object): - def __init__(self, model, **kwargs): - super().__init__() - self.model = model - to_torch = lambda x: x.clone().detach().to(torch.float32).to(model.device) - self.before_sample = None - self.after_sample = None - self.register_buffer("alphas_cumprod", to_torch(model.alphas_cumprod)) - - def register_buffer(self, name, attr): - if type(attr) == torch.Tensor: - if attr.device != devices.device: - attr = attr.to(devices.device) - setattr(self, name, attr) - - def set_hooks(self, before_sample, after_sample, after_update): - self.before_sample = before_sample - self.after_sample = after_sample - self.after_update = after_update - - @torch.no_grad() - def sample( - self, - S, - batch_size, - shape, - conditioning=None, - callback=None, - normals_sequence=None, - img_callback=None, - quantize_x0=False, - eta=0.0, - mask=None, - x0=None, - temperature=1.0, - noise_dropout=0.0, - score_corrector=None, - corrector_kwargs=None, - verbose=True, - x_T=None, - log_every_t=100, - unconditional_guidance_scale=1.0, - unconditional_conditioning=None, - # this has to come in the same format as the conditioning, # e.g. as encoded tokens, ... - **kwargs, - ): - if conditioning is not None: - if isinstance(conditioning, dict): - ctmp = conditioning[list(conditioning.keys())[0]] - while isinstance(ctmp, list): - ctmp = ctmp[0] - cbs = ctmp.shape[0] - if cbs != batch_size: - print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}") - - elif isinstance(conditioning, list): - for ctmp in conditioning: - if ctmp.shape[0] != batch_size: - print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}") - - else: - if conditioning.shape[0] != batch_size: - print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}") - - # sampling - C, H, W = shape - size = (batch_size, C, H, W) - # print(f'Data shape for UniPC sampling is {size}') - - device = self.model.betas.device - if x_T is None: - img = torch.randn(size, device=device) - else: - img = x_T - - ns = NoiseScheduleVP("discrete", alphas_cumprod=self.alphas_cumprod) - - # SD 1.X is "noise", SD 2.X is "v" - model_type = "v" if self.model.parameterization == "v" else "noise" - - model_fn = model_wrapper( - lambda x, t, c: self.model.apply_model(x, t, c), - ns, - model_type=model_type, - guidance_type="classifier-free", - # condition=conditioning, - # unconditional_condition=unconditional_conditioning, - guidance_scale=unconditional_guidance_scale, - ) - - uni_pc = UniPC(model_fn, ns, predict_x0=True, thresholding=False, variant=shared.opts.uni_pc_variant, condition=conditioning, unconditional_condition=unconditional_conditioning, before_sample=self.before_sample, after_sample=self.after_sample, after_update=self.after_update) - x = uni_pc.sample(img, steps=S, skip_type=shared.opts.uni_pc_skip_type, method="multistep", order=shared.opts.uni_pc_order, lower_order_final=shared.opts.uni_pc_lower_order_final) - - return x.to(device), None diff --git a/modules/uni_pc/uni_pc.py b/modules/uni_pc/uni_pc.py index 6a46cb5d..16baf6be 100644 --- a/modules/uni_pc/uni_pc.py +++ b/modules/uni_pc/uni_pc.py @@ -1,7 +1,9 @@ +# reference: https://github.com/comfyanonymous/ComfyUI/blob/v0.3.63/comfy/extra_samplers/uni_pc.py + import math import torch -import tqdm +from tqdm.auto import trange class NoiseScheduleVP: @@ -13,85 +15,6 @@ class NoiseScheduleVP: continuous_beta_0=0.1, continuous_beta_1=20.0, ): - r"""Create a wrapper class for the forward SDE (VP type). - - *** - Update: We support discrete-time diffusion models by implementing a picewise linear interpolation for log_alpha_t. - We recommend to use schedule='discrete' for the discrete-time diffusion models, especially for high-resolution images. - *** - - The forward SDE ensures that the condition distribution q_{t|0}(x_t | x_0) = N ( alpha_t * x_0, sigma_t^2 * I ). - We further define lambda_t = log(alpha_t) - log(sigma_t), which is the half-logSNR (described in the DPM-Solver paper). - Therefore, we implement the functions for computing alpha_t, sigma_t and lambda_t. For t in [0, T], we have: - - log_alpha_t = self.marginal_log_mean_coeff(t) - sigma_t = self.marginal_std(t) - lambda_t = self.marginal_lambda(t) - - Moreover, as lambda(t) is an invertible function, we also support its inverse function: - - t = self.inverse_lambda(lambda_t) - - =============================================================== - - We support both discrete-time DPMs (trained on n = 0, 1, ..., N-1) and continuous-time DPMs (trained on t in [t_0, T]). - - 1. For discrete-time DPMs: - - For discrete-time DPMs trained on n = 0, 1, ..., N-1, we convert the discrete steps to continuous time steps by: - t_i = (i + 1) / N - e.g. for N = 1000, we have t_0 = 1e-3 and T = t_{N-1} = 1. - We solve the corresponding diffusion ODE from time T = 1 to time t_0 = 1e-3. - - Args: - betas: A `torch.Tensor`. The beta array for the discrete-time DPM. (See the original DDPM paper for details) - alphas_cumprod: A `torch.Tensor`. The cumprod alphas for the discrete-time DPM. (See the original DDPM paper for details) - - Note that we always have alphas_cumprod = cumprod(betas). Therefore, we only need to set one of `betas` and `alphas_cumprod`. - - **Important**: Please pay special attention for the args for `alphas_cumprod`: - The `alphas_cumprod` is the \hat{alpha_n} arrays in the notations of DDPM. Specifically, DDPMs assume that - q_{t_n | 0}(x_{t_n} | x_0) = N ( \sqrt{\hat{alpha_n}} * x_0, (1 - \hat{alpha_n}) * I ). - Therefore, the notation \hat{alpha_n} is different from the notation alpha_t in DPM-Solver. In fact, we have - alpha_{t_n} = \sqrt{\hat{alpha_n}}, - and - log(alpha_{t_n}) = 0.5 * log(\hat{alpha_n}). - - - 2. For continuous-time DPMs: - - We support two types of VPSDEs: linear (DDPM) and cosine (improved-DDPM). The hyperparameters for the noise - schedule are the default settings in DDPM and improved-DDPM: - - Args: - beta_min: A `float` number. The smallest beta for the linear schedule. - beta_max: A `float` number. The largest beta for the linear schedule. - cosine_s: A `float` number. The hyperparameter in the cosine schedule. - cosine_beta_max: A `float` number. The hyperparameter in the cosine schedule. - T: A `float` number. The ending time of the forward process. - - =============================================================== - - Args: - schedule: A `str`. The noise schedule of the forward SDE. 'discrete' for discrete-time DPMs, - 'linear' or 'cosine' for continuous-time DPMs. - Returns: - A wrapper object of the forward SDE (VP type). - - =============================================================== - - Example: - - # For discrete-time DPMs, given betas (the beta array for n = 0, 1, ..., N - 1): - >>> ns = NoiseScheduleVP('discrete', betas=betas) - - # For discrete-time DPMs, given alphas_cumprod (the \hat{alpha_n} array for n = 0, 1, ..., N - 1): - >>> ns = NoiseScheduleVP('discrete', alphas_cumprod=alphas_cumprod) - - # For continuous-time DPMs (VPSDE), linear schedule: - >>> ns = NoiseScheduleVP('linear', continuous_beta_0=0.1, continuous_beta_1=20.) - - """ if schedule not in ["discrete", "linear", "cosine"]: raise ValueError(f"Unsupported noise schedule {schedule}. The schedule needs to be 'discrete' or 'linear' or 'cosine'") @@ -117,15 +40,13 @@ class NoiseScheduleVP: self.cosine_log_alpha_0 = math.log(math.cos(self.cosine_s / (1.0 + self.cosine_s) * math.pi / 2.0)) self.schedule = schedule if schedule == "cosine": - # For the cosine schedule, T = 1 will have numerical issues. So we manually set the ending time T. - # Note that T = 0.9946 may be not the optimal setting. However, we find it works well. self.T = 0.9946 else: self.T = 1.0 def marginal_log_mean_coeff(self, t): """ - Compute log(alpha_t) of a given continuous-time label t in [0, T]. + Compute log(alpha_t) of a given continuous-time label t in [0, T] """ if self.schedule == "discrete": return interpolate_fn(t.reshape((-1, 1)), self.t_array.to(t.device), self.log_alpha_array.to(t.device)).reshape((-1)) @@ -138,19 +59,19 @@ class NoiseScheduleVP: def marginal_alpha(self, t): """ - Compute alpha_t of a given continuous-time label t in [0, T]. + Compute alpha_t of a given continuous-time label t in [0, T] """ return torch.exp(self.marginal_log_mean_coeff(t)) def marginal_std(self, t): """ - Compute sigma_t of a given continuous-time label t in [0, T]. + Compute sigma_t of a given continuous-time label t in [0, T] """ return torch.sqrt(1.0 - torch.exp(2.0 * self.marginal_log_mean_coeff(t))) def marginal_lambda(self, t): """ - Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T]. + Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T] """ log_mean_coeff = self.marginal_log_mean_coeff(t) log_std = 0.5 * torch.log(1.0 - torch.exp(2.0 * log_mean_coeff)) @@ -158,7 +79,7 @@ class NoiseScheduleVP: def inverse_lambda(self, lamb): """ - Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t. + Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t """ if self.schedule == "linear": tmp = 2.0 * (self.beta_1 - self.beta_0) * torch.logaddexp(-2.0 * lamb, torch.zeros((1,)).to(lamb)) @@ -179,105 +100,14 @@ def model_wrapper( model, noise_schedule, model_type="noise", - model_kwargs=None, + model_kwargs={}, guidance_type="uncond", - # condition=None, - # unconditional_condition=None, + condition=None, + unconditional_condition=None, guidance_scale=1.0, classifier_fn=None, - classifier_kwargs=None, + classifier_kwargs={}, ): - """Create a wrapper function for the noise prediction model. - - DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to - firstly wrap the model function to a noise prediction model that accepts the continuous time as the input. - - We support four types of the diffusion model by setting `model_type`: - - 1. "noise": noise prediction model. (Trained by predicting noise). - - 2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0). - - 3. "v": velocity prediction model. (Trained by predicting the velocity). - The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2]. - - [1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models." - arXiv preprint arXiv:2202.00512 (2022). - [2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models." - arXiv preprint arXiv:2210.02303 (2022). - - 4. "score": marginal score function. (Trained by denoising score matching). - Note that the score function and the noise prediction model follows a simple relationship: - ``` - noise(x_t, t) = -sigma_t * score(x_t, t) - ``` - - We support three types of guided sampling by DPMs by setting `guidance_type`: - 1. "uncond": unconditional sampling by DPMs. - The input `model` has the following format: - `` - model(x, t_input, **model_kwargs) -> noise | x_start | v | score - `` - - 2. "classifier": classifier guidance sampling [3] by DPMs and another classifier. - The input `model` has the following format: - `` - model(x, t_input, **model_kwargs) -> noise | x_start | v | score - `` - - The input `classifier_fn` has the following format: - `` - classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond) - `` - - [3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis," - in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794. - - 3. "classifier-free": classifier-free guidance sampling by conditional DPMs. - The input `model` has the following format: - `` - model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score - `` - And if cond == `unconditional_condition`, the model output is the unconditional DPM output. - - [4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance." - arXiv preprint arXiv:2207.12598 (2022). - - - The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999) - or continuous-time labels (i.e. epsilon to T). - - We wrap the model function to accept only `x` and `t_continuous` as inputs, and outputs the predicted noise: - `` - def model_fn(x, t_continuous) -> noise: - t_input = get_model_input_time(t_continuous) - return noise_pred(model, x, t_input, **model_kwargs) - `` - where `t_continuous` is the continuous time labels (i.e. epsilon to T). And we use `model_fn` for DPM-Solver. - - =============================================================== - - Args: - model: A diffusion model with the corresponding format described above. - noise_schedule: A noise schedule object, such as NoiseScheduleVP. - model_type: A `str`. The parameterization type of the diffusion model. - "noise" or "x_start" or "v" or "score". - model_kwargs: A `dict`. A dict for the other inputs of the model function. - guidance_type: A `str`. The type of the guidance for sampling. - "uncond" or "classifier" or "classifier-free". - condition: A pytorch tensor. The condition for the guided sampling. - Only used for "classifier" or "classifier-free" guidance type. - unconditional_condition: A pytorch tensor. The condition for the unconditional sampling. - Only used for "classifier-free" guidance type. - guidance_scale: A `float`. The scale for the guided sampling. - classifier_fn: A classifier function. Only used for the classifier guidance. - classifier_kwargs: A `dict`. A dict for the other inputs of the classifier function. - Returns: - A noise prediction model that accepts the noised data and the continuous time as the inputs. - """ - - model_kwargs = model_kwargs or {} - classifier_kwargs = classifier_kwargs or {} def get_model_input_time(t_continuous): """ @@ -294,10 +124,7 @@ def model_wrapper( if t_continuous.reshape((-1,)).shape[0] == 1: t_continuous = t_continuous.expand((x.shape[0])) t_input = get_model_input_time(t_continuous) - if cond is None: - output = model(x, t_input, None, **model_kwargs) - else: - output = model(x, t_input, cond, **model_kwargs) + output = model(x, t_input, **model_kwargs) if model_type == "noise": return output elif model_type == "x_start": @@ -313,18 +140,18 @@ def model_wrapper( dims = x.dim() return -expand_dims(sigma_t, dims) * output - def cond_grad_fn(x, t_input, condition): + def cond_grad_fn(x, t_input): """ - Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t). + Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t) """ with torch.enable_grad(): x_in = x.detach().requires_grad_(True) log_prob = classifier_fn(x_in, t_input, condition, **classifier_kwargs) return torch.autograd.grad(log_prob.sum(), x_in)[0] - def model_fn(x, t_continuous, condition, unconditional_condition): + def model_fn(x, t_continuous): """ - The noise prediction model function that is used for DPM-Solver. + The noise prediction model function that is used for DPM-Solver """ if t_continuous.reshape((-1,)).shape[0] == 1: t_continuous = t_continuous.expand((x.shape[0])) @@ -333,7 +160,7 @@ def model_wrapper( elif guidance_type == "classifier": assert classifier_fn is not None t_input = get_model_input_time(t_continuous) - cond_grad = cond_grad_fn(x, t_input, condition) + cond_grad = cond_grad_fn(x, t_input) sigma_t = noise_schedule.marginal_std(t_continuous) noise = noise_pred_fn(x, t_continuous) return noise - guidance_scale * expand_dims(sigma_t, dims=cond_grad.dim()) * cond_grad @@ -343,21 +170,7 @@ def model_wrapper( else: x_in = torch.cat([x] * 2) t_in = torch.cat([t_continuous] * 2) - if isinstance(condition, dict): - assert isinstance(unconditional_condition, dict) - c_in = {} - for k in condition: - if isinstance(condition[k], list): - c_in[k] = [torch.cat([unconditional_condition[k][i], condition[k][i]]) for i in range(len(condition[k]))] - else: - c_in[k] = torch.cat([unconditional_condition[k], condition[k]]) - elif isinstance(condition, list): - c_in = [] - assert isinstance(unconditional_condition, list) - for i in range(len(condition)): - c_in.append(torch.cat([unconditional_condition[i], condition[i]])) - else: - c_in = torch.cat([unconditional_condition, condition]) + c_in = torch.cat([unconditional_condition, condition]) noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2) return noise_uncond + guidance_scale * (noise - noise_uncond) @@ -367,39 +180,21 @@ def model_wrapper( class UniPC: - def __init__( - self, - model_fn, - noise_schedule, - predict_x0=True, - thresholding=False, - max_val=1.0, - variant="bh1", - condition=None, - unconditional_condition=None, - before_sample=None, - after_sample=None, - after_update=None, - ): + def __init__(self, model_fn, noise_schedule, predict_x0=True, thresholding=False, max_val=1.0, variant="bh1"): """ - Construct a UniPC. - We support both data_prediction and noise_prediction. + Construct a UniPC + We support both data_prediction and noise_prediction """ - self.model_fn_ = model_fn + self.model = model_fn self.noise_schedule = noise_schedule self.variant = variant self.predict_x0 = predict_x0 self.thresholding = thresholding self.max_val = max_val - self.condition = condition - self.unconditional_condition = unconditional_condition - self.before_sample = before_sample - self.after_sample = after_sample - self.after_update = after_update def dynamic_thresholding_fn(self, x0, t=None): """ - The dynamic thresholding method. + The dynamic thresholding method """ dims = x0.dim() p = self.dynamic_thresholding_ratio @@ -408,45 +203,30 @@ class UniPC: x0 = torch.clamp(x0, -s, s) / s return x0 - def model(self, x, t): - cond = self.condition - uncond = self.unconditional_condition - if self.before_sample is not None: - x, t, cond, uncond = self.before_sample(x, t, cond, uncond) - res = self.model_fn_(x, t, cond, uncond) - if self.after_sample is not None: - x, t, cond, uncond, res = self.after_sample(x, t, cond, uncond, res) - - if isinstance(res, tuple): - # (None, pred_x0) - res = res[1] - - return res - def noise_prediction_fn(self, x, t): """ - Return the noise prediction model. + Return the noise prediction model """ return self.model(x, t) def data_prediction_fn(self, x, t): """ - Return the data prediction model (with thresholding). + Return the data prediction model (with thresholding) """ noise = self.noise_prediction_fn(x, t) dims = x.dim() alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t) x0 = (x - expand_dims(sigma_t, dims) * noise) / expand_dims(alpha_t, dims) if self.thresholding: - p = 0.995 # A hyperparameter in the paper of "Imagen" [1]. + p = 0.995 s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1) s = expand_dims(torch.maximum(s, self.max_val * torch.ones_like(s).to(s.device)), dims) x0 = torch.clamp(x0, -s, s) / s - return x0.to(x) + return x0 def model_fn(self, x, t): """ - Convert the model to the noise prediction model or the data prediction model. + Convert the model to the noise prediction model or the data prediction model """ if self.predict_x0: return self.data_prediction_fn(x, t) @@ -454,7 +234,7 @@ class UniPC: return self.noise_prediction_fn(x, t) def get_time_steps(self, skip_type, t_T, t_0, N, device): - """Compute the intermediate time steps for sampling.""" + """Compute the intermediate time steps for sampling""" if skip_type == "logSNR": lambda_T = self.noise_schedule.marginal_lambda(torch.tensor(t_T).to(device)) lambda_0 = self.noise_schedule.marginal_lambda(torch.tensor(t_0).to(device)) @@ -471,7 +251,7 @@ class UniPC: def get_orders_and_timesteps_for_singlestep_solver(self, steps, order, skip_type, t_T, t_0, device): """ - Get the order of each step for sampling by the singlestep DPM-Solver. + Get the order of each step for sampling by the singlestep DPM-Solver """ if order == 3: K = steps // 3 + 1 @@ -494,7 +274,6 @@ class UniPC: else: raise ValueError("'order' must be '1' or '2' or '3'.") if skip_type == "logSNR": - # To reproduce the results in DPM-Solver paper timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, K, device) else: timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, steps, device)[torch.cumsum(torch.tensor([0] + orders), 0).to(device)] @@ -502,7 +281,7 @@ class UniPC: def denoise_to_zero_fn(self, x, s): """ - Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization. + Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization """ return self.data_prediction_fn(x, s) @@ -516,11 +295,9 @@ class UniPC: return self.multistep_uni_pc_vary_update(x, model_prev_list, t_prev_list, t, order, **kwargs) def multistep_uni_pc_vary_update(self, x, model_prev_list, t_prev_list, t, order, use_corrector=True): - # print(f'using unified predictor-corrector with order {order} (solver type: vary coeff)') ns = self.noise_schedule assert order <= len(model_prev_list) - # first compute rks t_prev_0 = t_prev_list[-1] lambda_prev_0 = ns.marginal_lambda(t_prev_0) lambda_t = ns.marginal_lambda(t) @@ -545,7 +322,6 @@ class UniPC: rks = torch.tensor(rks, device=x.device) K = len(rks) - # build C matrix C = [] col = torch.ones_like(rks) @@ -555,12 +331,11 @@ class UniPC: C = torch.stack(C, dim=1) if len(D1s) > 0: - D1s = torch.stack(D1s, dim=1) # (B, K) + D1s = torch.stack(D1s, dim=1) C_inv_p = torch.linalg.inv(C[:-1, :-1]) A_p = C_inv_p if use_corrector: - # print('using corrector') C_inv = torch.linalg.inv(C) A_c = C_inv @@ -577,13 +352,10 @@ class UniPC: model_t = None if self.predict_x0: x_t_ = sigma_t / sigma_prev_0 * x - alpha_t * h_phi_1 * model_prev_0 - # now predictor x_t = x_t_ if len(D1s) > 0: - # compute the residuals for predictor for k in range(K - 1): x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum("bkchw,k->bchw", D1s, A_p[k]) - # now corrector if use_corrector: model_t = self.model_fn(x_t, t) D1_t = model_t - model_prev_0 @@ -595,13 +367,10 @@ class UniPC: else: log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t) x_t_ = (torch.exp(log_alpha_t - log_alpha_prev_0)) * x - (sigma_t * h_phi_1) * model_prev_0 - # now predictor x_t = x_t_ if len(D1s) > 0: - # compute the residuals for predictor for k in range(K - 1): x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum("bkchw,k->bchw", D1s, A_p[k]) - # now corrector if use_corrector: model_t = self.model_fn(x_t, t) D1_t = model_t - model_prev_0 @@ -613,12 +382,10 @@ class UniPC: return x_t, model_t def multistep_uni_pc_bh_update(self, x, model_prev_list, t_prev_list, t, order, x_t=None, use_corrector=True): - # print(f'using unified predictor-corrector with order {order} (solver type: B(h))') ns = self.noise_schedule assert order <= len(model_prev_list) dims = x.dim() - # first compute rks t_prev_0 = t_prev_list[-1] lambda_prev_0 = ns.marginal_lambda(t_prev_0) lambda_t = ns.marginal_lambda(t) @@ -646,7 +413,7 @@ class UniPC: b = [] hh = -h[0] if self.predict_x0 else h[0] - h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1 + h_phi_1 = torch.expm1(hh) h_phi_k = h_phi_1 / hh - 1 factorial_i = 1 @@ -667,12 +434,10 @@ class UniPC: R = torch.stack(R) b = torch.tensor(b, device=x.device) - # now predictor use_predictor = len(D1s) > 0 and x_t is None if len(D1s) > 0: - D1s = torch.stack(D1s, dim=1) # (B, K) + D1s = torch.stack(D1s, dim=1) if x_t is None: - # for order 2, we use a simplified version if order == 2: rhos_p = torch.tensor([0.5], device=b.device) else: @@ -681,8 +446,6 @@ class UniPC: D1s = None if use_corrector: - # print('using corrector') - # for order 1, we use a simplified version if order == 1: rhos_c = torch.tensor([0.5], device=b.device) else: @@ -694,7 +457,7 @@ class UniPC: if x_t is None: if use_predictor: - pred_res = torch.einsum("k,bkchw->bchw", rhos_p, D1s) + pred_res = torch.tensordot(D1s, rhos_p, dims=([1], [0])) else: pred_res = 0 x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * pred_res @@ -702,7 +465,7 @@ class UniPC: if use_corrector: model_t = self.model_fn(x_t, t) if D1s is not None: - corr_res = torch.einsum("k,bkchw->bchw", rhos_c[:-1], D1s) + corr_res = torch.tensordot(D1s, rhos_c[:-1], dims=([1], [0])) else: corr_res = 0 D1_t = model_t - model_prev_0 @@ -711,7 +474,7 @@ class UniPC: x_t_ = expand_dims(torch.exp(log_alpha_t - log_alpha_prev_0), dims) * x - expand_dims(sigma_t * h_phi_1, dims) * model_prev_0 if x_t is None: if use_predictor: - pred_res = torch.einsum("k,bkchw->bchw", rhos_p, D1s) + pred_res = torch.tensordot(D1s, rhos_p, dims=([1], [0])) else: pred_res = 0 x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * pred_res @@ -719,68 +482,75 @@ class UniPC: if use_corrector: model_t = self.model_fn(x_t, t) if D1s is not None: - corr_res = torch.einsum("k,bkchw->bchw", rhos_c[:-1], D1s) + corr_res = torch.tensordot(D1s, rhos_c[:-1], dims=([1], [0])) else: corr_res = 0 D1_t = model_t - model_prev_0 x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t) return x_t, model_t - def sample(self, x, steps=20, t_start=None, t_end=None, order=3, skip_type="time_uniform", method="singlestep", lower_order_final=True, denoise_to_zero=False, solver_type="dpm_solver", atol=0.0078, rtol=0.05, corrector=False): - t_0 = 1.0 / self.noise_schedule.total_N if t_end is None else t_end - t_T = self.noise_schedule.T if t_start is None else t_start - device = x.device + def sample( + self, + x, + timesteps, + t_start=None, + t_end=None, + order=3, + skip_type="time_uniform", + method="singlestep", + lower_order_final=True, + denoise_to_zero=False, + solver_type="dpm_solver", + atol=0.0078, + rtol=0.05, + corrector=False, + callback=None, + disable_pbar=False, + ): + steps = len(timesteps) - 1 if method == "multistep": - assert steps >= order, "UniPC order must be < sampling steps" - timesteps = self.get_time_steps(skip_type=skip_type, t_T=t_T, t_0=t_0, N=steps, device=device) - # print(f"Running UniPC Sampling with {timesteps.shape[0]} timesteps, order {order}") + assert steps >= order assert timesteps.shape[0] - 1 == steps - with torch.no_grad(): - vec_t = timesteps[0].expand((x.shape[0])) - model_prev_list = [self.model_fn(x, vec_t)] - t_prev_list = [vec_t] - with tqdm.tqdm(total=steps) as pbar: - # Init the first `order` values by lower order multistep DPM-Solver. - for init_order in range(1, order): - vec_t = timesteps[init_order].expand(x.shape[0]) - x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, init_order, use_corrector=True) - if model_x is None: - model_x = self.model_fn(x, vec_t) - if self.after_update is not None: - self.after_update(x, model_x) - model_prev_list.append(model_x) - t_prev_list.append(vec_t) - pbar.update() - - for step in range(order, steps + 1): + for step_index in trange(steps, disable=disable_pbar): + if step_index == 0: + vec_t = timesteps[0].expand((x.shape[0])) + model_prev_list = [self.model_fn(x, vec_t)] + t_prev_list = [vec_t] + elif step_index < order: + init_order = step_index + vec_t = timesteps[init_order].expand(x.shape[0]) + x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, init_order, use_corrector=True) + if model_x is None: + model_x = self.model_fn(x, vec_t) + model_prev_list.append(model_x) + t_prev_list.append(vec_t) + else: + extra_final_step = 0 + if step_index == (steps - 1): + extra_final_step = 1 + for step in range(step_index, step_index + 1 + extra_final_step): vec_t = timesteps[step].expand(x.shape[0]) if lower_order_final: step_order = min(order, steps + 1 - step) else: step_order = order - # print('this step order:', step_order) if step == steps: - # print('do not run corrector at the last step') use_corrector = False else: use_corrector = True x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, step_order, use_corrector=use_corrector) - if self.after_update is not None: - self.after_update(x, model_x) for i in range(order - 1): t_prev_list[i] = t_prev_list[i + 1] model_prev_list[i] = model_prev_list[i + 1] t_prev_list[-1] = vec_t - # We do not need to evaluate the final model value. if step < steps: if model_x is None: model_x = self.model_fn(x, vec_t) model_prev_list[-1] = model_x - pbar.update() + if callback is not None: + callback({"x": x, "i": step_index, "denoised": model_prev_list[-1]}) else: raise NotImplementedError() - if denoise_to_zero: - x = self.denoise_to_zero_fn(x, torch.ones((x.shape[0],)).to(device) * t_0) return x @@ -790,18 +560,6 @@ class UniPC: def interpolate_fn(x, xp, yp): - """ - A piecewise linear function y = f(x), using xp and yp as keypoints. - We implement f(x) in a differentiable way (i.e. applicable for autograd). - The function f(x) is well-defined for all x-axis. (For x beyond the bounds of xp, we use the outmost points of xp to define the linear function.) - - Args: - x: PyTorch tensor with shape [N, C], where N is the batch size, C is the number of channels (we use C = 1 for DPM-Solver). - xp: PyTorch tensor with shape [C, K], where K is the number of keypoints. - yp: PyTorch tensor with shape [C, K]. - Returns: - The function values f(x), with shape [N, C]. - """ N, K = x.shape[0], xp.shape[1] all_x = torch.cat([x.unsqueeze(2), xp.unsqueeze(0).repeat((N, 1, 1))], dim=2) sorted_all_x, x_indices = torch.sort(all_x, dim=2) @@ -836,13 +594,4 @@ def interpolate_fn(x, xp, yp): def expand_dims(v, dims): - """ - Expand the tensor `v` to the dim `dims`. - - Args: - `v`: a PyTorch tensor with shape [N]. - `dim`: a `int`. - Returns: - a PyTorch tensor with shape [N, 1, 1, ..., 1] and the total dimension is `dims`. - """ return v[(...,) + (None,) * (dims - 1)]