diff --git a/README.md b/README.md index 98bd053..31052c1 100644 --- a/README.md +++ b/README.md @@ -29,6 +29,8 @@ information about the steps it's going to run to the console. | 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. | | KSampler With Restarts (Custom) | | Essentially the same as `KSampler With Restarts (Advanced)` but it takes a `SAMPLER` input like the built in `SamplerCustom` node. Note that it is possible to input samplers that don't work properly or are incompatible with Restart sampling like SDE and UniPC samplers.| +| `RestartScheduler` | | For use with custom sampling: This node will output sigmas like other scheduler nodes with restart segments inserted. Must be used with `RestartSampler`. Like stand alone samplers, the node takes parameters for restart segments and schedules. You may also optionally connect sigmas to it, in which case it will use the supplied sigmas for the main schedule. **Note**: When sigmas are connected, the `steps` and `scheduler` parameters have no effect. | +| `RestartSampler` | | For use with custom sampling: Should be used in conjunction with `RestartScheduler` and takes a `SAMPLER` input. This node arranges for the restart noise to be injected at the appropriate points and delegates to the supplied sampler for actual sampling. | ### Segments @@ -44,9 +46,9 @@ You may freely mix the different formats. For example, `[2, 2, -500, "10%"], [3, **Special segment values**: -* Enter `default` to use the default segment list. -* Enter `a1111` to emulate A1111 WebUI's segment calculation behavior. -For full emulation, enabled chunked mode, set both schedulers to `karras` and the sampler to `heun`. +* Enter `default` by itself to use the default segment list. +* Enter `a1111` by itself to emulate A1111 WebUI's segment calculation behavior. For full emulation, enabled chunked mode, set both schedulers to `karras` and the sampler to `heun`. +* You may also enter `"default"` or `"a1111"` in place of a segment definition (note the quotes). This will insert preset segments at the point the quoted preset name appears. For example `[1,2,3,4], "default"` is the same as `[1,2,3,4], [3,2,0.06,0.30], [3,1,0.30,0.59]`. ### Chunked Mode diff --git a/nodes.py b/nodes.py index a46eee2..fa957ad 100644 --- a/nodes.py +++ b/nodes.py @@ -4,7 +4,8 @@ import comfy from .restart_sampling import ( DEFAULT_SEGMENTS, - SCHEDULER_MAPPING, + NORMAL_SCHEDULER_MAPPING, + RESTART_SCHEDULER_MAPPING, VERBOSE, RestartPlan, RestartSampler, @@ -36,7 +37,11 @@ def get_supported_samplers(): def get_supported_restart_schedulers(): - return list(SCHEDULER_MAPPING.keys()) + return tuple(RESTART_SCHEDULER_MAPPING.keys()) + + +def get_supported_normal_schedulers(): + return tuple(NORMAL_SCHEDULER_MAPPING.keys()) class KRestartSamplerSimple: @@ -49,7 +54,7 @@ class KRestartSamplerSimple: "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,), + "scheduler": (get_supported_restart_schedulers(),), "positive": ("CONDITIONING",), "negative": ("CONDITIONING",), "latent_image": ("LATENT",), @@ -105,7 +110,7 @@ class KRestartSampler: "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": (tuple(SCHEDULER_MAPPING.keys()),), + "scheduler": (get_supported_normal_schedulers(),), "positive": ("CONDITIONING",), "negative": ("CONDITIONING",), "latent_image": ("LATENT",), @@ -173,7 +178,7 @@ class KRestartSamplerAdv: "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": (tuple(SCHEDULER_MAPPING.keys()),), + "scheduler": (get_supported_normal_schedulers(),), "positive": ("CONDITIONING",), "negative": ("CONDITIONING",), "latent_image": ("LATENT",), @@ -247,7 +252,7 @@ class KRestartSamplerCustom: "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), "sampler": ("SAMPLER",), - "scheduler": (tuple(SCHEDULER_MAPPING.keys()),), + "scheduler": (get_supported_normal_schedulers(),), "positive": ("CONDITIONING",), "negative": ("CONDITIONING",), "latent_image": ("LATENT",), @@ -316,13 +321,15 @@ class RestartSchedulerNode: "required": { "model": ("MODEL",), "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), - "scheduler": (tuple(SCHEDULER_MAPPING.keys()),), + "scheduler": (get_supported_normal_schedulers(),), "segments": ( "STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}, ), "restart_scheduler": (get_supported_restart_schedulers(),), "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), + "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), + "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), }, "optional": { "sigmas_opt": ("SIGMAS",), @@ -341,6 +348,8 @@ class RestartSchedulerNode: segments, restart_scheduler, denoise, + start_at_step=0, + end_at_step=10000, sigmas_opt=None, ): plan = RestartPlan( @@ -349,7 +358,8 @@ class RestartSchedulerNode: scheduler, segments, restart_scheduler, - denoise, + denoise=denoise, + step_range=(start_at_step, end_at_step), sigmas=sigmas_opt, ) if VERBOSE: diff --git a/restart_sampling.py b/restart_sampling.py index 7c6a1f7..c7f0851 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -13,7 +13,7 @@ from comfy.samplers import KSAMPLER, sampler_object from comfy.utils import ProgressBar from tqdm.auto import trange -from .restart_schedulers import SCHEDULER_MAPPING +from .restart_schedulers import NORMAL_SCHEDULER_MAPPING, RESTART_SCHEDULER_MAPPING VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1" @@ -146,8 +146,17 @@ def round_restart_segments(ts, restart_segments): return t_min_mapping -def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device): - return SCHEDULER_MAPPING[scheduler](model, n, sigma_min, sigma_max, device) +def calc_sigmas( + scheduler, + n, + sigma_min, + sigma_max, + model, + device, + restart_segment=True, +): + mapping = RESTART_SCHEDULER_MAPPING if restart_segment else NORMAL_SCHEDULER_MAPPING + return mapping[scheduler](model, n, sigma_min, sigma_max, device) def restart_sampling( @@ -354,6 +363,14 @@ class RestartPlan: force_full_denoise=False, sigmas=None, ): + if ( + denoise <= 0 + or (sigmas is None and steps < 1) + or (sigmas is not None and len(sigmas) < 2) + ): + self.plan = [] + self.total_steps = 0 + return ms = model.get_model_object("model_sampling") if sigmas is None: @@ -365,11 +382,14 @@ class RestartPlan: float(ms.sigma_max), model.model, "cpu", + restart_segment=False, ) else: steps = effective_steps = len(sigmas) - 1 steps = steps if denoise > 0.9999 else int(effective_steps * denoise) sigmas = sigmas.clone().detach().cpu() + if effective_steps != steps: + sigmas = sigmas[-(steps + 1) :] if step_range is not None: start_step, last_step = step_range @@ -380,8 +400,6 @@ class RestartPlan: if start_step < len(sigmas) - 1: sigmas = sigmas[start_step:] - elif effective_steps != steps: - sigmas = sigmas[-(steps + 1) :] restart_segments = prepare_restart_segments(restart_info, ms, sigmas) self.plan, self.total_steps = self.build_plan_items( @@ -408,7 +426,6 @@ class RestartPlan: device, ) -> tuple[list, int]: model_sigma_min = float(model.model_sampling.sigma_min) - model_sigma_max = float(model.model_sampling.sigma_max) segments = round_restart_segments(sigmas, restart_segments) plan = [] range_start = -1 @@ -431,7 +448,7 @@ class RestartPlan: restart_scheduler, n_restart, max(model_sigma_min, sigmas[i + 1]), - min(model_sigma_max, s_max), + s_max, model, device=device, ) @@ -448,6 +465,9 @@ class RestartPlan: def sigmas(self) -> torch.Tensor: # Flattens a plan into sigmas. When the first normal sigma matches the last item's # final sigma, we strip the first normal sigma to avoid creating duplicates. + if not self.plan or self.total_steps < 1: + return torch.FloatTensor([]) + def sigmas_generator(): prev_last = None for pi in self.plan: @@ -499,9 +519,9 @@ class RestartPlan: max_steps=100, ) -> None: if schedules is None: - schedules = SCHEDULER_MAPPING.keys() + schedules = NORMAL_SCHEDULER_MAPPING.keys() if restart_schedules is None: - restart_schedules = SCHEDULER_MAPPING.keys() + restart_schedules = RESTART_SCHEDULER_MAPPING.keys() if segments is None: segments = ("default", "a1111") for schname in schedules: diff --git a/restart_schedulers.py b/restart_schedulers.py index e8e78c5..78cab30 100644 --- a/restart_schedulers.py +++ b/restart_schedulers.py @@ -1,3 +1,4 @@ +import comfy import torch from comfy.k_diffusion import sampling as k_diffusion_sampling @@ -81,6 +82,10 @@ def get_sigmas_ddim_uniform(model, n, s_min, s_max, device): return torch.tensor(sigs, device=device) +def get_sigmas_sgm_uniform(model, n, s_min, s_max, device): + return normal_scheduler(model, n, s_min, s_max, sgm=True).to(device) + + def get_sigmas_simple_test(model, n, s_min, s_max, device): ms = model.model_sampling min_idx = torch.argmin(torch.abs(ms.sigmas - s_min)) @@ -91,11 +96,32 @@ def get_sigmas_simple_test(model, n, s_min, s_max, device): return torch.tensor(sigs, device=device) -SCHEDULER_MAPPING = { +def get_comfy_scheduler_fn(name): + return ( + lambda model, + steps, + _smin, + _smax, + device="cpu": comfy.samplers.calculate_sigmas( + model.model_sampling, + name, + steps, + ).to(device) + ) + + +RESTART_SCHEDULER_MAPPING = { "normal": get_sigmas_normal, "karras": get_sigmas_karras, "exponential": get_sigmas_exponential, "simple": get_sigmas_simple, "ddim_uniform": get_sigmas_ddim_uniform, + "sgm_uniform": get_sigmas_sgm_uniform, + "simple_test": get_sigmas_simple_test, +} + +NORMAL_SCHEDULER_MAPPING = { + k: get_comfy_scheduler_fn(k) for k in comfy.samplers.SCHEDULER_NAMES +} | { "simple_test": get_sigmas_simple_test, }