Add some comments describing what's going on in plan generation and execution

Allow specifying segments "default" to use the default segments

Allow specifying segments "a1111" to calculate segments like A1111

Allow setting environment variable COMFYUI_VERBOSE_RESTART_SAMPLING=1 to get some debug info

Add chunked_mode to samplers (except for simple which will use the default of True)

Documentation updates
This commit is contained in:
blepping
2024-03-27 05:45:06 -06:00
parent bbbddbd7cb
commit e2e25e7a28
3 changed files with 103 additions and 57 deletions
+18
View File
@@ -16,6 +16,9 @@ git clone https://github.com/ssitu/ComfyUI_restart_sampling
The Restart sampler nodes can be found in the node menu under `sampling`.
If you set the environment variable `COMFYUI_VERBOSE_RESTART_SAMPLING` to `1`, restart sampling will dump
information about the steps it's going to run to the console.
### Nodes
|Node|Image|Description|
@@ -37,6 +40,21 @@ Both $t_{\textrm{min}}$ and $t_{\textrm{max}}$ within a segment definition may b
You may freely mix the different formats. For example, `[2, 2, -500, "10%"], [3, 2, 5.3, -3]` would be a valid sequence. Note: Random numbers used for example only, not recommended.
**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`.
### Chunked Mode
When chunked mode is enabled, the sampler is called with as many steps as possible up to the next segment. When disabled, the sampler
is only called with a single step at a time. Some samplers such as SDE samplers, momentum samplers, second order samplers
like dpmpp_2m use state from previous steps - when called step-by-step, this state is lost. Using chunked mode may make those
samplers more accurate.
*Note*: Using SDE or momentum samplers with restart is likely not an improvement over normal sampling.
## Visual Example
Consider the default segments of `[3,2,0.06,0.30],[3,1,0.30,0.59]`.
+12 -51
View File
@@ -1,5 +1,5 @@
import comfy
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING, DEFAULT_SEGMENTS
def get_supported_samplers():
@@ -24,8 +24,6 @@ def get_supported_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
@@ -42,7 +40,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": DEFAULT_SEGMENTS, "multiline": False}),
"segments": ("STRING", {"default": "default", "multiline": False}),
}
}
@@ -50,7 +48,7 @@ class KRestartSamplerSimple:
FUNCTION = "sample"
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, chunked_mode=False):
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise)
@@ -71,6 +69,7 @@ class KRestartSampler:
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
"restart_scheduler": (get_supported_restart_schedulers(), ),
"chunked_mode": ("BOOLEAN", {"default": True}),
}
}
@@ -78,8 +77,8 @@ class KRestartSampler:
FUNCTION = "sample"
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, segments, restart_scheduler, denoise=denoise)
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=False):
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise, chunked_mode=chunked_mode)
class KRestartSamplerAdv:
@@ -103,6 +102,7 @@ class KRestartSamplerAdv:
"return_with_leftover_noise": (["disable", "enable"], ),
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
"restart_scheduler": (get_supported_restart_schedulers(), ),
"chunked_mode": ("BOOLEAN", {"default": True}),
}
}
@@ -110,10 +110,10 @@ class KRestartSamplerAdv:
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):
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, chunked_mode=False):
force_full_denoise = return_with_leftover_noise != "enable"
disable_noise = add_noise == "disable"
return restart_sampling(model, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise)
return restart_sampling(model, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, chunked_mode=chunked_mode)
class KRestartSamplerCustom:
@@ -137,6 +137,7 @@ class KRestartSamplerCustom:
"return_with_leftover_noise": (["disable", "enable"], ),
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
"restart_scheduler": (get_supported_restart_schedulers(), ),
"chunked_mode": ("BOOLEAN", {"default": True}),
}
}
@@ -145,56 +146,16 @@ class KRestartSamplerCustom:
FUNCTION = "sample"
CATEGORY = "sampling"
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler):
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, chunked_mode=False):
force_full_denoise = return_with_leftover_noise != "enable"
disable_noise = add_noise == "disable"
return restart_sampling(model, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, output_only=False)
class KRestartSamplerCustomNoise:
@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": ("SAMPLER", ),
"scheduler": (tuple(SCHEDULER_MAPPING.keys()), ),
"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(),),
"chunked_mode": ("BOOLEAN", {"default": False}),
},
"optional": {
"custom_noise_opt": ("SONAR_CUSTOM_NOISE",),
},
}
RETURN_TYPES = ("LATENT","LATENT")
RETURN_NAMES = ("output", "denoised_output")
FUNCTION = "sample"
CATEGORY = "sampling"
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, custom_noise_opt=None, chunked_mode="disable"):
force_full_denoise = return_with_leftover_noise != "enable"
disable_noise = add_noise == "disable"
return restart_sampling(model, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, output_only=False, custom_noise=custom_noise_opt.make_noise_sampler if custom_noise_opt else None, chunked_mode=chunked_mode)
return restart_sampling(model, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, output_only=False, chunked_mode=chunked_mode)
NODE_CLASS_MAPPINGS = {
"KRestartSamplerSimple": KRestartSamplerSimple,
"KRestartSampler": KRestartSampler,
"KRestartSamplerAdv": KRestartSamplerAdv,
"KRestartSamplerCustom": KRestartSamplerCustom,
"KRestartSamplerCustomNoise": KRestartSamplerCustomNoise,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+73 -6
View File
@@ -1,5 +1,6 @@
import ast
from collections import namedtuple
import os
import warnings
import torch
from tqdm.auto import trange
@@ -10,6 +11,10 @@ from comfy.samplers import KSAMPLER, sampler_object
from comfy.utils import ProgressBar
from .restart_schedulers import SCHEDULER_MAPPING
VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1"
DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]"
def add_restart_segment(restart_segments, n_restart, k, t_min, t_max):
if restart_segments is None:
@@ -34,7 +39,25 @@ def resolve_t_value(val, ms):
raise ValueError("bad t_min or t_max value")
def prepare_restart_segments(restart_info, ms):
def prepare_restart_segments(restart_info, ms, sigmas):
restart_info = restart_info.strip().lower()
if restart_info == "":
# No restarts.
return []
if restart_info == "default":
restart_info = DEFAULT_SEGMENTS
elif restart_info == "a1111":
# Emulate A1111 WebUI's restart sampler behavior.
steps = len(sigmas) - 1
if steps < 20:
# Less than 20 steps - no restarts.
return []
if steps < 36:
# Less than 36 steps - one restart with 9 steps.
restart_info = "[10,1,0.1,0.2]"
else:
# Otherwise two restarts with steps // 4 steps.
restart_info = f"[{(steps // 4) + 1}, 2, 0.1, 0.2]"
try:
restart_arrays = ast.literal_eval(f"[{restart_info}]")
except SyntaxError as e:
@@ -88,18 +111,15 @@ def calc_restart_steps(restart_segments):
return restart_steps
def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=False):
def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=True):
if isinstance(sampler, str):
sampler = sampler_object(sampler)
comfy.model_management.load_models_gpu([model])
real_model = model
while hasattr(real_model, "model"):
real_model = real_model.model
restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling)
effective_steps = steps if step_range is not None or denoise > 0.9999 else int(steps / denoise)
sigmas = calc_sigmas(scheduler, effective_steps,
float(real_model.model_sampling.sigma_min), float(real_model.model_sampling.sigma_max),
@@ -118,6 +138,8 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
elif effective_steps != steps:
sigmas = sigmas[-(steps + 1):]
restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling, sigmas)
sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, seed, custom_noise, chunked=chunked_mode)
latent = latent_image
@@ -175,8 +197,20 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], defaults=[None, 0, 0., 0., None])):
# sigmas: Sigmas for normal (outside of a restart segment) sampling. They start from after the previous PlanItem's steps
# if there is one or simply the beginning of sampling.
# k, s_min, s_max: This is the same as the restart segment definition. Set to 0 if there is no restart segment.
# restart_sigmas: Sigmas for the restart segment if it exists, otherwise None.
# Note: n_restart is not included as it can be calculated from the length of restart_sigmas.
__slots__ = ()
# Execute a plan item: runs sampling on the main sigmas, handles injecting noise for restarts
# as well as sampling the restart steps.
# sample: Function used sample sigmas. It takes x, a tensor with the sigmas to sample and
# the restart index (k) or -1 for sampling that isn't within a restart segment.
# get_noise_sampler: Return the noise sampler for restart segment noise injection.
# It takes x, and sigma_min, sigma_max (basically the same arguments as ComfyUI's
# BrownianTreeNoiseSampler class init function).
@torch.no_grad()
def execute(self, x, sample, get_noise_sampler):
x = sample(x, self.sigmas, -1)
@@ -190,6 +224,18 @@ class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart
class KSamplerRestartWrapper:
# Some extra explanation for a couple of these arguments:
#
# chunked:
# When chunked is False, the sampling function is called step-by-step with only two sigmas at a time.
# When chunked is is True, the sampling function will be called with sigmas for multiple steps at a time.
# this means either the steps up to the next restart segment (or the end of sampling) or the steps within
# a restart segment.
#
# make_noise_sampler:
# If set to None, restart noise will just use torch.randn_like (gaussian) for noise generation. Otherwise
# this should contain a function that takes x, sigma_min, sigma_max, seed and returns a noise sampler
# function (which takes sigma, sigma_next) and returns a noisy tensor.
def __init__(self, sampler, real_model, restart_scheduler, restart_segments, seed, make_noise_sampler=None, chunked=True):
self.ksampler = sampler
self.real_model = real_model
@@ -200,6 +246,9 @@ class KSamplerRestartWrapper:
self.make_noise_sampler = make_noise_sampler
self.chunked = chunked
# Builds a list of PlanItems and calculates the total number of steps. See the comments for PlanItem
# for more information about plans.
# Returns two values: the plan and the total steps.
@torch.no_grad()
def build_plan(self, sigmas, device):
segments = round_restart_segments(sigmas, self.restart_segments)
@@ -222,9 +271,15 @@ class KSamplerRestartWrapper:
plan.append(PlanItem(sigmas[range_start:]))
return plan, total_steps
# Dumps information about the plan to the console. It uses the normal plan execute
# logic.
def explain_plan(self, plan, total_steps):
print(f"** Dumping restart sampling plan (total steps {total_steps}):")
step = 0
last_kidx = -1
# Instead of actually sampling, we just dump information about the steps.
# When kidx==-1 this is a normal step, otherwise kidx==0 is the first restart,
# kidx==1 is the second, etc.
def do_sample(x, sigs, kidx=-1):
nonlocal step, last_kidx
rlabel = f"R{kidx+1:>3}" if kidx > last_kidx else " "
@@ -240,11 +295,13 @@ class KSamplerRestartWrapper:
step += chunk_size
return x
# Stub function to satisfy PlanItem.execute
def get_noise_sampler(*_args):
return lambda *_args: 0.0
for pi in plan:
pi.execute(0.0, do_sample, get_noise_sampler)
print("** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.")
@torch.no_grad()
@@ -254,11 +311,16 @@ class KSamplerRestartWrapper:
seed = self.seed
plan, self.total_steps = self.build_plan(sigmas, x.device)
self.explain_plan(plan, self.total_steps)
if VERBOSE:
self.explain_plan(plan, self.total_steps)
def noise_sampler(*_args):
return torch.randn_like(x)
# Passed to the PlanItem .execute method. Most of the time, self.make_noise_sampler
# is going to be None so this is just a wrapper for torch.randn_like.
# Otherwise we call the noise sampler factory and increment seed to ensure that restarts
# don't all use the same noise.
def get_noise_sampler(x, s_min, s_max):
nonlocal seed
if not self.make_noise_sampler:
@@ -276,6 +338,7 @@ class KSamplerRestartWrapper:
if callback is not None:
callback(x)
# Convenience function for code reuse.
def sampler_function(x, sigs):
return ksampler.sampler_function(
model, x, sigs, *args, extra_args=extra_args, callback=callback_wrapper, disable=True,
@@ -285,11 +348,15 @@ class KSamplerRestartWrapper:
if isinstance(sigs, (list, tuple)):
sigs = torch.tensor(sigs, device=x.device)
if self.chunked or len(sigs) < 3:
# If running un chunked mode or there are already 2 or less sigmas, we can just
# pass the sigmas to the sampling function.
return sampler_function(x, sigs)
# Otherwise we call the sampling function step by step on slices of 2 sigmas.
for i in range(len(sigs)-1):
x = sampler_function(x, sigs[i:i+2])
return x
# Execute the plan items in sequence.
for pi in plan:
x = pi.execute(x, do_sample, get_noise_sampler)