diff --git a/.gitignore b/.gitignore index d9005f2..57e650e 100644 --- a/.gitignore +++ b/.gitignore @@ -150,3 +150,6 @@ cython_debug/ # and can be added to the global gitignore or merged into this file. For a more nuclear # option (not recommended) you can uncomment the following to ignore the entire idea folder. #.idea/ + +# Misc +*.ipynb diff --git a/README.md b/README.md index f6097b3..bab0455 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,9 @@ # ComfyUI_restart_sampling Unofficial [ComfyUI](https://github.com/comfyanonymous/ComfyUI) nodes for restart sampling based on the paper "Restart Sampling for Improving Generative Processes" -[[paper]](https://arxiv.org/abs/2306.14878) [[repo]](https://github.com/Newbeeer/diffusion_restart_sampling) + +Paper: https://arxiv.org/abs/2306.14878 + +Repo: https://github.com/Newbeeer/diffusion_restart_sampling ## Installation @@ -16,6 +19,8 @@ Nodes can be found in the node menu under `sampling`: |Node|Image|Description| | --- | --- | --- | | KSampler With Restarts | ![image](https://github.com/ssitu/ComfyUI_restart_sampling/assets/57548627/7696da21-ea8c-4263-91a9-658d0f87dc47) | Has all the inputs of a KSampler, but with an added string widget for configuring the Restart segments and a widget for the scheduler for the Restart segments. Not all samplers and schedulers from KSampler are currently supported. Restart sampling is done with ODE samplers and are not supposed to be used with SDE samplers.
The format for `segments` is a sequence of comma separated arrays of ${[N_{\textrm{Restart}}, K, t_{\textrm{min}}, t_{\textrm{max}}]}$. For example, [4, 1, 19.35, 40.79], [4, 1, 1.09, 1.92], [4, 5, 0.59, 1.09], [4, 5, 0.30, 0.59], [6, 6, 0.06, 0.30] would be a valid sequence. Segments may overwrite each other if their $t_{\textrm{min}}$ parameters are too close to each other. Each segment will add $(N_{\textrm{Restart}} - 1) \cdot K$ steps to the sampling process. For more information on the Restart parameters, refer to the paper.
The `restart_scheduler` is used as the scheduler for the backwards process during restart intervals. The researchers used the Karras scheduler in their experiments, but use the same scheduler as the sampler schedule in their implementation. | +| KSampler With Restarts (Simple) | | Instead of having a restart segment scheduler, segments will use the same scheduler as the KSampler scheduler. | +| KSampler With Restarts (Advanced) | | Has all the inputs for an Advanced KSampler with all the inputs for restart sampling. It should be noted that there is a possibility for invalid segments when using it to end the denoising process early or starting it late (e.g. 20 steps, start at step 0, end at step 10) and invalid segments will be ignored. An invalid segment means that the closest $t_{\textrm{min}}$ in the noise schedule is higher than the segment's $t_{\textrm{max}}$, so the segment would have restarted the denoising process at $t_{\textrm{max}}$ then try to go to a higher noise level (when it should've gone to a lower noise level near $t_{\textrm{min}}$) which will destroy the sample. | --- diff --git a/nodes.py b/nodes.py index 1f75b69..8a9c269 100644 --- a/nodes.py +++ b/nodes.py @@ -4,20 +4,24 @@ from .restart_sampling import restart_sampling, SCHEDULER_MAPPING def get_supported_samplers(): samplers = comfy.samplers.KSampler.SAMPLERS.copy() - samplers.remove("uni_pc") - samplers.remove("uni_pc_bh2") # SDE samplers cannot be used with restarts + samplers.remove("uni_pc") + samplers.remove("uni_pc_bh2") samplers.remove("dpmpp_sde") samplers.remove("dpmpp_sde_gpu") samplers.remove("dpmpp_2m_sde") samplers.remove("dpmpp_2m_sde_gpu") + samplers.remove("dpmpp_3m_sde") + samplers.remove("dpmpp_3m_sde_gpu") return samplers def get_supported_restart_schedulers(): return list(SCHEDULER_MAPPING.keys()) +DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]" + class KRestartSamplerSimple: @classmethod @@ -34,7 +38,7 @@ class KRestartSamplerSimple: "negative": ("CONDITIONING", ), "latent_image": ("LATENT", ), "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "segments": ("STRING", {"default": "[3,2,0.06,0.30],[3,1,0.30,0.59]", "multiline": False}), + "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), } } @@ -43,7 +47,7 @@ class KRestartSamplerSimple: CATEGORY = "sampling" def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments): - return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, scheduler) + return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise) class KRestartSampler: @@ -61,7 +65,7 @@ class KRestartSampler: "negative": ("CONDITIONING", ), "latent_image": ("LATENT", ), "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "segments": ("STRING", {"default": "[3,2,0.06,0.30],[3,1,0.30,0.59]", "multiline": False}), + "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), } } @@ -71,15 +75,55 @@ class KRestartSampler: CATEGORY = "sampling" def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler): - return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler) + return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise) + + +class KRestartSamplerAdv: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "add_noise": (["enable", "disable"], ), + "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), + "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), + "sampler_name": (get_supported_samplers(), ), + "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "latent_image": ("LATENT", ), + "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), + "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), + "return_with_leftover_noise": (["disable", "enable"], ), + "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), + "restart_scheduler": (get_supported_restart_schedulers(), ), + } + } + + RETURN_TYPES = ("LATENT",) + FUNCTION = "sample" + CATEGORY = "sampling" + + def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler): + force_full_denoise = True + if return_with_leftover_noise == "enable": + force_full_denoise = False + disable_noise = False + if add_noise == "disable": + disable_noise = True + return restart_sampling(model, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, start_step=start_at_step, last_step=end_at_step, force_full_denoise=force_full_denoise) NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, "KRestartSampler": KRestartSampler, + "KRestartSamplerAdv": KRestartSamplerAdv, } NODE_DISPLAY_NAME_MAPPINGS = { "KRestartSamplerSimple": "KSampler With Restarts (Simple)", "KRestartSampler": "KSampler With Restarts", + "KRestartSamplerAdv": "KSampler With Restarts (Advanced)", } diff --git a/restart_sampling.py b/restart_sampling.py index 264a3c5..a36065d 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,4 +1,5 @@ import ast +import warnings import torch from tqdm.auto import trange from nodes import common_ksampler @@ -31,14 +32,6 @@ def prepare_restart_segments(restart_info): return restart_segments -def round_restart_segments(sigmas, restart_segments): - t_min_mapping = {} - for segment in reversed(restart_segments): # Reversed to prioritize segments to the front - t_min_neighbor = min(sigmas, key=lambda s: abs(s - segment['t_min'])).item() - t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']} - return t_min_mapping - - def segments_to_timesteps(restart_segments, model): timesteps = [] for segment in restart_segments: @@ -49,10 +42,23 @@ def segments_to_timesteps(restart_segments, model): return timesteps -def round_restart_segments_timesteps(timesteps, restart_segments): +def round_restart_segments(ts, restart_segments): + """ + Map nearest timestep/sigma min to the nearest timestep/sigma to segments. + :param ts: Timesteps or sigmas of the original denoising schedule + :param restart_segments: Restart segments dict of the form {'t_min': t_min, 'n': n, 'k': k, 't_max': t_max} + :return: dict of the form {nearest_t_min: {'n': n, 'k': k, 't_max': t_max}} + """ t_min_mapping = {} for segment in reversed(restart_segments): # Reversed to prioritize segments to the front - t_min_neighbor = min(timesteps, key=lambda ts: abs(ts - segment['t_min'])).item() + t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment['t_min'])).item() + if t_min_neighbor > segment['t_max']: + warnings.warn( + f"\n[Restart Sampling] t_min neighbor {t_min_neighbor} is greater than t_max {segment['t_max']}, ignoring segment...", stacklevel=2) + continue + if t_min_neighbor in t_min_mapping: + warnings.warn( + f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']} is {t_min_neighbor}", stacklevel=2) t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']} return t_min_mapping @@ -73,7 +79,7 @@ _restart_segments = None _restart_scheduler = None -def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, restart_info, restart_scheduler): +def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False): global _total_steps, _restart_segments, _restart_scheduler _restart_scheduler = restart_scheduler _restart_segments = prepare_restart_segments(restart_info) @@ -93,14 +99,56 @@ def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, ProgressBar.update_absolute = pbar_update_absolute_wrapper try: - samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, - positive, negative, latent_image, denoise=denoise) + samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=denoise, + disable_noise=disable_noise, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise) finally: sampler_wrapper.cleanup() ProgressBar.update_absolute = pbar_update_absolute return samples +class OneStepSampler: + + def __init__(self, model, steps, cfg, sampler, scheduler, positive, negative, latent_image, denoise): + # Keep parameters for sampler + self.model = model + self.steps = steps + self.cfg = cfg + self.sampler = sampler + self.scheduler = scheduler + self.positive = positive + self.negative = negative + self.latent_image = latent_image + self.denoise = denoise + + # Get the sampler function + match sampler: + case "ddim": + sampler = DDIMSampler(self.model, device=self.device) + sampler.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) + z_enc = sampler.stochastic_encode(latent_image, torch.tensor( + [len(timesteps) - 1] * noise.shape[0]).to(self.device), noise=noise, max_denoise=max_denoise) + samples, _ = sampler.sample_custom(ddim_timesteps=timesteps, + conditioning=positive, + batch_size=noise.shape[0], + shape=noise.shape[1:], + verbose=False, + unconditional_guidance_scale=cfg, + unconditional_conditioning=negative, + eta=0.0, + x_T=z_enc, + x0=latent_image, + img_callback=ddim_callback, + denoise_function=sampling_function, + extra_args=extra_args, + mask=noise_mask, + to_zero=sigmas[-1] == 0, + end_step=sigmas.shape[0] - 1, + disable_pbar=disable_pbar) + case _: + sample = getattr(k_diffusion_sampling, f"sample_{sampler}") + + class RestartWrapper: def cleanup(self): @@ -113,7 +161,7 @@ class KSamplerRestartWrapper(RestartWrapper): def __init__(self, sampler_name): self.sample_func_name = "sample_{}".format(sampler_name) - self.__class__.ksampler = getattr(k_diffusion_sampling, self.sample_func_name) + KSamplerRestartWrapper.ksampler = getattr(k_diffusion_sampling, self.sample_func_name) setattr(k_diffusion_sampling, self.sample_func_name, self.ksampler_restart_wrapper) def cleanup(self): @@ -123,7 +171,7 @@ class KSamplerRestartWrapper(RestartWrapper): @torch.no_grad() def ksampler_restart_wrapper(model, x, sigmas, extra_args=None, callback=None, disable=None): global _total_steps, _restart_segments, _restart_scheduler - ksampler = __class__.ksampler + ksampler = KSamplerRestartWrapper.ksampler segments = round_restart_segments(sigmas, _restart_segments) _total_steps = len(sigmas) - 1 + calc_restart_steps(segments) step = 0 @@ -174,18 +222,17 @@ class DDIMWrapper(RestartWrapper): ddim_sampler = __class__.sample_custom model_denoise = CompVisVDenoiser(self.model) segments = segments_to_timesteps(_restart_segments, model_denoise) - segments = round_restart_segments_timesteps(ddim_timesteps, segments) + segments = round_restart_segments(ddim_timesteps, segments) _total_steps = len(ddim_timesteps) - 1 + calc_restart_steps(segments) step = 0 def callback_wrapper(pred_x0, i): img_callback(pred_x0, step) - def ddim_simplified(x, timesteps, x_T=None, disable_pbar=False): - if x_T is None: - self.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) - x_T = self.stochastic_encode(x, torch.tensor( - [len(timesteps) - 1] * x.shape[0]).to(self.device), noise=torch.zeros_like(x), max_denoise=False) + def ddim_simplified(x, timesteps, disable_pbar=False): + self.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) + x_T = self.stochastic_encode(x, torch.tensor( + [len(timesteps) - 1] * x.shape[0]).to(self.device), noise=torch.zeros_like(x), max_denoise=False) x, intermediates = ddim_sampler( self, timesteps, conditioning, callback=callback, img_callback=callback_wrapper, quantize_x0=quantize_x0, eta=eta, mask=mask, x0=x, temperature=temperature, noise_dropout=noise_dropout, score_corrector=score_corrector, @@ -198,8 +245,7 @@ class DDIMWrapper(RestartWrapper): with trange(_total_steps, disable=disable_pbar) as pbar: for i in reversed(range(len(ddim_timesteps) - 1)): - x0, intermediates = ddim_simplified(x0, ddim_timesteps[i:i + 2], x_T=x_T, disable_pbar=True) - x_T = None + x0, intermediates = ddim_simplified(x0, ddim_timesteps[i:i + 2], disable_pbar=True) pbar.update(1) step += 1 if ddim_timesteps[i].item() in segments: