Add advanced node, remove DPM++ 3m SDE

This commit is contained in:
ssit
2023-09-10 11:05:46 -04:00
parent 8c65bc0090
commit c3fec18633
4 changed files with 128 additions and 30 deletions
+3
View File
@@ -150,3 +150,6 @@ cython_debug/
# and can be added to the global gitignore or merged into this file. For a more nuclear # 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. # option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/ #.idea/
# Misc
*.ipynb
+6 -1
View File
@@ -1,6 +1,9 @@
# ComfyUI_restart_sampling # ComfyUI_restart_sampling
Unofficial [ComfyUI](https://github.com/comfyanonymous/ComfyUI) nodes for restart sampling based on the paper "Restart Sampling for Improving Generative Processes" 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 ## Installation
@@ -16,6 +19,8 @@ Nodes can be found in the node menu under `sampling`:
|Node|Image|Description| |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. <br>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. <br>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 | ![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. <br>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. <br>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. |
--- ---
+50 -6
View File
@@ -4,20 +4,24 @@ from .restart_sampling import restart_sampling, SCHEDULER_MAPPING
def get_supported_samplers(): def get_supported_samplers():
samplers = comfy.samplers.KSampler.SAMPLERS.copy() samplers = comfy.samplers.KSampler.SAMPLERS.copy()
samplers.remove("uni_pc")
samplers.remove("uni_pc_bh2")
# SDE samplers cannot be used with restarts # SDE samplers cannot be used with restarts
samplers.remove("uni_pc")
samplers.remove("uni_pc_bh2")
samplers.remove("dpmpp_sde") samplers.remove("dpmpp_sde")
samplers.remove("dpmpp_sde_gpu") samplers.remove("dpmpp_sde_gpu")
samplers.remove("dpmpp_2m_sde") samplers.remove("dpmpp_2m_sde")
samplers.remove("dpmpp_2m_sde_gpu") samplers.remove("dpmpp_2m_sde_gpu")
samplers.remove("dpmpp_3m_sde")
samplers.remove("dpmpp_3m_sde_gpu")
return samplers return samplers
def get_supported_restart_schedulers(): def get_supported_restart_schedulers():
return list(SCHEDULER_MAPPING.keys()) return list(SCHEDULER_MAPPING.keys())
DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]"
class KRestartSamplerSimple: class KRestartSamplerSimple:
@classmethod @classmethod
@@ -34,7 +38,7 @@ class KRestartSamplerSimple:
"negative": ("CONDITIONING", ), "negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ), "latent_image": ("LATENT", ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "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" CATEGORY = "sampling"
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments): 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: class KRestartSampler:
@@ -61,7 +65,7 @@ class KRestartSampler:
"negative": ("CONDITIONING", ), "negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ), "latent_image": ("LATENT", ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "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(), ), "restart_scheduler": (get_supported_restart_schedulers(), ),
} }
} }
@@ -71,15 +75,55 @@ class KRestartSampler:
CATEGORY = "sampling" CATEGORY = "sampling"
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler): 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 = { NODE_CLASS_MAPPINGS = {
"KRestartSamplerSimple": KRestartSamplerSimple, "KRestartSamplerSimple": KRestartSamplerSimple,
"KRestartSampler": KRestartSampler, "KRestartSampler": KRestartSampler,
"KRestartSamplerAdv": KRestartSamplerAdv,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"KRestartSamplerSimple": "KSampler With Restarts (Simple)", "KRestartSamplerSimple": "KSampler With Restarts (Simple)",
"KRestartSampler": "KSampler With Restarts", "KRestartSampler": "KSampler With Restarts",
"KRestartSamplerAdv": "KSampler With Restarts (Advanced)",
} }
+69 -23
View File
@@ -1,4 +1,5 @@
import ast import ast
import warnings
import torch import torch
from tqdm.auto import trange from tqdm.auto import trange
from nodes import common_ksampler from nodes import common_ksampler
@@ -31,14 +32,6 @@ def prepare_restart_segments(restart_info):
return restart_segments 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): def segments_to_timesteps(restart_segments, model):
timesteps = [] timesteps = []
for segment in restart_segments: for segment in restart_segments:
@@ -49,10 +42,23 @@ def segments_to_timesteps(restart_segments, model):
return timesteps 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 = {} t_min_mapping = {}
for segment in reversed(restart_segments): # Reversed to prioritize segments to the front 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']} t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']}
return t_min_mapping return t_min_mapping
@@ -73,7 +79,7 @@ _restart_segments = None
_restart_scheduler = 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 global _total_steps, _restart_segments, _restart_scheduler
_restart_scheduler = restart_scheduler _restart_scheduler = restart_scheduler
_restart_segments = prepare_restart_segments(restart_info) _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 ProgressBar.update_absolute = pbar_update_absolute_wrapper
try: try:
samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=denoise,
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: finally:
sampler_wrapper.cleanup() sampler_wrapper.cleanup()
ProgressBar.update_absolute = pbar_update_absolute ProgressBar.update_absolute = pbar_update_absolute
return samples 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: class RestartWrapper:
def cleanup(self): def cleanup(self):
@@ -113,7 +161,7 @@ class KSamplerRestartWrapper(RestartWrapper):
def __init__(self, sampler_name): def __init__(self, sampler_name):
self.sample_func_name = "sample_{}".format(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) setattr(k_diffusion_sampling, self.sample_func_name, self.ksampler_restart_wrapper)
def cleanup(self): def cleanup(self):
@@ -123,7 +171,7 @@ class KSamplerRestartWrapper(RestartWrapper):
@torch.no_grad() @torch.no_grad()
def ksampler_restart_wrapper(model, x, sigmas, extra_args=None, callback=None, disable=None): def ksampler_restart_wrapper(model, x, sigmas, extra_args=None, callback=None, disable=None):
global _total_steps, _restart_segments, _restart_scheduler global _total_steps, _restart_segments, _restart_scheduler
ksampler = __class__.ksampler ksampler = KSamplerRestartWrapper.ksampler
segments = round_restart_segments(sigmas, _restart_segments) segments = round_restart_segments(sigmas, _restart_segments)
_total_steps = len(sigmas) - 1 + calc_restart_steps(segments) _total_steps = len(sigmas) - 1 + calc_restart_steps(segments)
step = 0 step = 0
@@ -174,18 +222,17 @@ class DDIMWrapper(RestartWrapper):
ddim_sampler = __class__.sample_custom ddim_sampler = __class__.sample_custom
model_denoise = CompVisVDenoiser(self.model) model_denoise = CompVisVDenoiser(self.model)
segments = segments_to_timesteps(_restart_segments, model_denoise) 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) _total_steps = len(ddim_timesteps) - 1 + calc_restart_steps(segments)
step = 0 step = 0
def callback_wrapper(pred_x0, i): def callback_wrapper(pred_x0, i):
img_callback(pred_x0, step) img_callback(pred_x0, step)
def ddim_simplified(x, timesteps, x_T=None, disable_pbar=False): def ddim_simplified(x, timesteps, disable_pbar=False):
if x_T is None: self.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False)
self.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) x_T = self.stochastic_encode(x, torch.tensor(
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)
[len(timesteps) - 1] * x.shape[0]).to(self.device), noise=torch.zeros_like(x), max_denoise=False)
x, intermediates = ddim_sampler( x, intermediates = ddim_sampler(
self, timesteps, conditioning, callback=callback, img_callback=callback_wrapper, quantize_x0=quantize_x0, 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, 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: with trange(_total_steps, disable=disable_pbar) as pbar:
for i in reversed(range(len(ddim_timesteps) - 1)): 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) x0, intermediates = ddim_simplified(x0, ddim_timesteps[i:i + 2], disable_pbar=True)
x_T = None
pbar.update(1) pbar.update(1)
step += 1 step += 1
if ddim_timesteps[i].item() in segments: if ddim_timesteps[i].item() in segments: