Merge pull request #18 from blepping/feat_restart_sampler

Make restart a sampler, add node to generate sigmas
This commit is contained in:
ssitu
2024-05-07 18:02:07 -04:00
committed by GitHub
4 changed files with 569 additions and 233 deletions
+6 -4
View File
@@ -5,7 +5,7 @@ Paper: https://arxiv.org/abs/2306.14878
Repo: https://github.com/Newbeeer/diffusion_restart_sampling
This has been tested for ComfyUI for the following commit: [d14bdb1](https://github.com/comfyanonymous/ComfyUI/commit/d14bdb18967f7413852a364747c49599de537eec)
This has been tested for ComfyUI for the following commit: [72508a8](https://github.com/comfyanonymous/ComfyUI/commit/72508a8d19121e2814ea4dfbce8a5311f37dcd61)
## Installation
@@ -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. Setting `denoise` also can't adjust the steps: it can only shorten the sigmas you pass to the node. |
| `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
+131 -6
View File
@@ -1,6 +1,20 @@
import os
import comfy
from .restart_sampling import DEFAULT_SEGMENTS, SCHEDULER_MAPPING, restart_sampling
from .restart_sampling import (
DEFAULT_SEGMENTS,
NORMAL_SCHEDULER_MAPPING,
RESTART_SCHEDULER_MAPPING,
VERBOSE,
RestartPlan,
RestartSampler,
restart_sampling,
)
INCLUDE_SELFTEST = (
os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1"
)
def get_supported_samplers():
@@ -23,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:
@@ -36,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",),
@@ -92,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",),
@@ -160,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",),
@@ -234,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",),
@@ -296,11 +314,93 @@ class KRestartSamplerCustom:
)
class RestartSchedulerNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"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",),
},
}
RETURN_TYPES = ("SIGMAS",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/schedulers"
def go(
self,
model,
steps,
scheduler,
segments,
restart_scheduler,
denoise,
start_at_step=0,
end_at_step=10000,
sigmas_opt=None,
):
plan = RestartPlan(
model,
steps,
scheduler,
segments,
restart_scheduler,
denoise=denoise,
step_range=(start_at_step, end_at_step),
sigmas=sigmas_opt,
)
if VERBOSE:
plan.explain(chunked=True)
return (plan.sigmas(),)
class RestartSamplerNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sampler": ("SAMPLER",),
"chunked_mode": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("SAMPLER",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/samplers"
def go(self, sampler, chunked_mode):
restart_options = {
"restart_chunked": chunked_mode,
"restart_wrapped_sampler": sampler,
}
restart_sampler = comfy.samplers.KSAMPLER(
RestartSampler.sampler_function,
extra_options=sampler.extra_options | restart_options,
inpaint_options=sampler.inpaint_options,
)
return (restart_sampler,)
NODE_CLASS_MAPPINGS = {
"KRestartSamplerSimple": KRestartSamplerSimple,
"KRestartSampler": KRestartSampler,
"KRestartSamplerAdv": KRestartSamplerAdv,
"KRestartSamplerCustom": KRestartSamplerCustom,
"RestartScheduler": RestartSchedulerNode,
"RestartSampler": RestartSamplerNode,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -309,3 +409,28 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"KRestartSamplerAdv": "KSampler With Restarts (Advanced)",
"KRestartSamplerCustom": "KSampler With Restarts (Custom)",
}
if INCLUDE_SELFTEST:
class RestartSelfTestNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"enabled": ("BOOLEAN", {"default": True}),
"min_steps": ("INT", {"default": 2, "min": 0}),
"max_steps": ("INT", {"default": 100, "min": 2}),
},
}
RETURN_TYPES = ("BOOLEAN",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/samplers"
def go(self, model, enabled=True, min_steps=2, max_steps=100):
if enabled:
RestartPlan.self_test(model, min_steps=min_steps, max_steps=max_steps)
return (True,)
NODE_CLASS_MAPPINGS["RestartSelfTest"] = RestartSelfTestNode
+404 -221
View File
@@ -1,3 +1,5 @@
from __future__ import annotations
import ast
import os
import warnings
@@ -7,11 +9,11 @@ import comfy
import latent_preview
import torch
from comfy.sample import prepare_noise, sample_custom
from comfy.samplers import KSAMPLER, sampler_object
from comfy.samplers import KSAMPLER, 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"
@@ -42,14 +44,7 @@ def resolve_t_value(val, ms):
def prepare_restart_segments(restart_info, ms, sigmas):
restart_info = restart_info.strip().lower()
if restart_info == "":
# No restarts.
return []
restart_arrays = None
if restart_info == "default":
restart_info = DEFAULT_SEGMENTS
elif restart_info == "a1111":
def get_a1111_segment():
# Emulate A1111 WebUI's restart sampler behavior.
steps = len(sigmas) - 1
if steps < 20:
@@ -58,20 +53,48 @@ def prepare_restart_segments(restart_info, ms, sigmas):
a1111_t_max = sigmas[int(torch.argmin(abs(sigmas - 2.0), dim=0))].item()
if steps < 36:
# Less than 36 steps - one restart with 9 steps.
restart_arrays = [[10, 1, 0.1, a1111_t_max]]
else:
# Otherwise two restarts with steps // 4 steps.
restart_arrays = [[(steps // 4) + 1, 2, 0.1, a1111_t_max]]
return [10, 1, 0.1, a1111_t_max]
# Otherwise two restarts with steps // 4 steps.
return [(steps // 4) + 1, 2, 0.1, a1111_t_max]
restart_info = restart_info.strip().lower()
if restart_info == "":
# No restarts.
return []
restart_arrays = None
if restart_info == "default":
restart_info = DEFAULT_SEGMENTS
elif restart_info == "a1111":
restart_arrays = [get_a1111_segment()]
if restart_arrays == [[]]:
return []
if restart_arrays is None:
try:
restart_arrays = ast.literal_eval(f"[{restart_info}]")
except SyntaxError:
print("Ill-formed restart segments")
raise
temp = []
default_segments = ast.literal_eval(DEFAULT_SEGMENTS)
# This phase expands any preset strings into actual 4-item restart segments.
for idx in range(len(restart_arrays)):
item = restart_arrays[idx]
if not isinstance(item, str):
temp.append(item)
continue
preset = item.strip().lower()
if preset == "default":
temp += default_segments
elif preset == "a1111":
temp.append(get_a1111_segment())
else:
raise ValueError("Ill-formed restart segment")
restart_arrays = temp
restart_segments = []
# Now we build the actual restart segments.
for arr in restart_arrays:
if len(arr) != 4:
raise ValueError("Restart segment must have 4 values")
if not isinstance(arr, (list, tuple)) or len(arr) != 4:
raise ValueError("Restart segment must be a list with 4 values")
n_restart, k, val_min, val_max = arr
n_restart, k = int(n_restart), int(k)
t_min = resolve_t_value(val_min, ms)
@@ -123,15 +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_restart_steps(restart_segments):
restart_steps = 0
for segment in restart_segments.values():
restart_steps += (segment["n"] - 1) * segment["k"]
return restart_steps
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(
@@ -156,58 +181,39 @@ def restart_sampling(
sigmas=None,
):
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
effective_steps = (
steps if step_range is not None or denoise > 0.9999 else int(steps / denoise)
)
if sigmas is None:
sigmas = calc_sigmas(
scheduler,
effective_steps,
float(real_model.model_sampling.sigma_min),
float(real_model.model_sampling.sigma_max),
real_model,
model.load_device,
# Only possible to determine this when the sampler is passed by name. When using
# a custom sampler, the user will need to slice sigmas when desirable.
discard_penultimate_sigma = sampler in getattr(
KSampler,
"DISCARD_PENULTIMATE_SIGMA_SAMPLERS",
set(),
)
sampler = sampler_object(sampler)
else:
sigmas = sigmas.detach().clone().to(model.load_device)
if step_range is not None:
start_step, last_step = step_range
discard_penultimate_sigma = False
if last_step < (len(sigmas) - 1):
sigmas = sigmas[: last_step + 1]
if force_full_denoise:
sigmas[-1] = 0
if start_step < (len(sigmas) - 1):
sigmas = sigmas[start_step:]
elif effective_steps != steps:
sigmas = sigmas[-(steps + 1) :]
restart_segments = prepare_restart_segments(
plan = RestartPlan(
model,
steps,
scheduler,
restart_info,
real_model.model_sampling,
sigmas,
restart_scheduler,
denoise=denoise,
step_range=step_range,
force_full_denoise=force_full_denoise,
sigmas=sigmas,
discard_penultimate_sigma=discard_penultimate_sigma,
)
sampler_wrapper = KSamplerRestartWrapper(
sampler,
real_model,
restart_scheduler,
restart_segments,
seed,
custom_noise,
chunked=chunked_mode,
)
if VERBOSE:
plan.explain(chunked_mode)
total_steps = plan.total_steps
sigmas = plan.sigmas().to(model.load_device)
latent = latent_image
latent_image = latent["samples"]
if disable_noise:
torch.manual_seed(
seed,
@@ -227,21 +233,27 @@ def restart_sampling(
noise_mask = latent["noise_mask"]
x0_output = {}
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
callback = latent_preview.prepare_callback(model, plan.total_steps, x0_output)
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
sampler = KSAMPLER(
sampler_wrapper.ksampler_restart_wrapper,
extra_options=sampler.extra_options | {},
restart_options = {
"restart_chunked": chunked_mode,
"restart_wrapped_sampler": sampler,
"restart_custom_noise": custom_noise,
}
ksampler = KSAMPLER(
RestartSampler.sampler_function,
extra_options=sampler.extra_options | restart_options,
inpaint_options=sampler.inpaint_options | {},
)
# Add the additional steps to the progress bar
pbar_update_absolute = ProgressBar.update_absolute
def pbar_update_absolute_wrapper(self, value, total=None, preview=None):
pbar_update_absolute(self, value, sampler_wrapper.total_steps, preview)
def pbar_update_absolute_wrapper(self, value, total=None, preview=None): # noqa: ARG001
pbar_update_absolute(self, value, total_steps, preview)
ProgressBar.update_absolute = pbar_update_absolute_wrapper
@@ -250,7 +262,7 @@ def restart_sampling(
model,
noise,
cfg,
sampler,
ksampler,
sigmas,
positive,
negative,
@@ -278,81 +290,165 @@ def restart_sampling(
return (out, out_denoised)
# PlanItem:
# 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: 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.
# Convenience properties:
# total_steps
# s_min, s_max: None if no restart sigmas.
# n_restart: 0 if no restart sigmas.
class PlanItem(
namedtuple(
"PlanItem",
["sigmas", "k", "s_min", "s_max", "restart_sigmas"],
defaults=[None, 0, 0.0, 0.0, None],
["sigmas", "k", "restart_sigmas"],
defaults=[None, 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)
if self.k < 1 or self.restart_sigmas is None:
return x
noise_sampler = get_noise_sampler(x, self.s_min, self.s_max)
for kidx in range(self.k):
x += (
noise_sampler(self.restart_sigmas[0], self.restart_sigmas[-1])
* (self.s_max**2 - self.s_min**2) ** 0.5
def __new__(cls, *args: list, **kwargs: dict):
obj = super().__new__(cls, *args, **kwargs)
obj.validate()
return obj
def validate(self, threshold=1e-06):
if len(self.sigmas) < 2:
raise ValueError("PlanItem: invalid normal sigmas: too short")
t = self.sigmas.sort(descending=True, stable=True)[0].unique_consecutive()
if not torch.equal(self.sigmas, t):
errstr = (
f"PlanItem: invalid normal sigmas: out of order or contains duplicates: {self}",
)
x = sample(x, self.restart_sigmas, kidx)
return x
raise ValueError(errstr)
if self.k == 0:
return
if self.k < 0:
raise ValueError("PlanItem: invalid negative k value")
if len(self.restart_sigmas) < 2:
raise ValueError("PlanItem: invalid restart sigmas: too short")
if self.s_min >= self.s_max:
raise ValueError("PlanItem: invalid min/max: min >= max")
if self.sigmas[-1] - self.restart_sigmas[0] > threshold:
raise ValueError(
"PlanItem: invalid sigmas: last normal sigma >= first restart sigma",
)
if self.sigmas[-1] - self.restart_sigmas[-1] > threshold:
errstr = (
f"PlanItem: invalid sigmas: last restart sigma {self.restart_sigmas[-1]} < last normal sigma {self.sigmas[-1]}",
)
raise ValueError(errstr)
t = self.restart_sigmas.sort(descending=True, stable=True)[
0
].unique_consecutive()
if not torch.equal(self.restart_sigmas, t):
errstr = (
f"PlanItem: invalid restart sigmas: out of order or contains duplicates: {self}",
)
raise ValueError(errstr)
@property
def total_steps(self):
if self.k < 1:
return len(self.sigmas) - 1
return (len(self.sigmas) - 1) + (len(self.restart_sigmas) - 1) * self.k
@property
def s_min(self):
return None if self.k < 1 else self.restart_sigmas[-1].item()
@property
def s_max(self):
return None if self.k < 1 else self.restart_sigmas[0].item()
@property
def n_restart(self):
return 0 if self.k < 1 else len(self.restart_sigmas) - 1
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.
class RestartPlan:
def __init__(
self,
sampler,
real_model,
model,
steps,
scheduler,
restart_info,
restart_scheduler,
restart_segments,
seed,
make_noise_sampler=None,
chunked=True,
denoise=1.0,
step_range=None,
force_full_denoise=False,
sigmas=None,
discard_penultimate_sigma=False,
):
self.ksampler = sampler
self.real_model = real_model
self.restart_scheduler = restart_scheduler
self.restart_segments = restart_segments
self.total_steps = 0
self.seed = seed
self.make_noise_sampler = make_noise_sampler
self.chunked = chunked
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:
effective_steps = steps if denoise > 0.9999 else int(steps / denoise)
sigmas = calc_sigmas(
scheduler,
effective_steps + int(discard_penultimate_sigma), # True evaluates to 1
float(ms.sigma_min),
float(ms.sigma_max),
model.model,
"cpu",
restart_segment=False,
)
if discard_penultimate_sigma:
sigmas = torch.cat((sigmas[:-2], sigmas[-1:]))
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
if last_step < len(sigmas) - 1:
sigmas = sigmas[: last_step + 1]
if force_full_denoise:
sigmas[-1] = 0
if start_step < len(sigmas) - 1:
sigmas = sigmas[start_step:]
restart_segments = prepare_restart_segments(restart_info, ms, sigmas)
self.plan, self.total_steps = self.build_plan_items(
model.model,
restart_segments,
restart_scheduler,
sigmas,
"cpu",
)
def __repr__(self) -> str:
return f"<RestartPlan: steps={self.total_steps}, plan={self.plan}>"
# 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.
@staticmethod
@torch.no_grad()
def build_plan(self, sigmas, device):
segments = round_restart_segments(sigmas, self.restart_segments)
total_steps = len(sigmas) - 1 + calc_restart_steps(segments)
def build_plan_items(
model,
restart_segments,
restart_scheduler,
sigmas,
device,
) -> tuple[list, int]:
model_sigma_min = float(model.model_sampling.sigma_min)
segments = round_restart_segments(sigmas, restart_segments)
plan = []
range_start = -1
for i in range(len(sigmas) - 1):
@@ -364,102 +460,194 @@ class KSamplerRestartWrapper:
if seg is None:
continue
s_max, k, n_restart = seg["t_max"], seg["k"], seg["n"]
seg_sigmas = calc_sigmas(
self.restart_scheduler,
if k < 1 or n_restart < 2:
continue
if s_max <= model_sigma_min:
errstr = f"Restart: Invalid restart segment t_max {s_max:.05} <= model minimum sigma {model_sigma_min:.05}"
raise ValueError(errstr)
normal_sigmas = sigmas[range_start : i + 2]
restart_sigmas = calc_sigmas(
restart_scheduler,
n_restart,
s_min,
max(model_sigma_min, s_min),
s_max,
self.real_model,
model,
device=device,
)
plan.append(
PlanItem(sigmas[range_start : i + 2], k, s_min, s_max, seg_sigmas[:-1]),
)
if normal_sigmas[-1] != 0:
restart_sigmas = restart_sigmas[:-1]
restart_sigmas[-1] = s_min # Force the restart segment to end at s_min.
plan.append(PlanItem(normal_sigmas, k, restart_sigmas))
range_start = -1
if range_start != -1:
# Include sigmas after the last restart segments in the plan.
plan.append(PlanItem(sigmas[range_start:]))
return plan, total_steps
return plan, sum(pi.total_steps for pi in plan)
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:
yield pi.sigmas if prev_last != pi.sigmas[0] else pi.sigmas[1:]
prev_last = pi.restart_sigmas[-1] if pi.k > 0 else pi.sigmas[-1]
for _ in range(pi.k):
yield pi.restart_sigmas
return torch.flatten(torch.cat(tuple(sigmas_generator())))
# Dumps information about the plan to the console. It uses the normal plan execute
# logic.
def explain_plan(self, plan, total_steps):
def explain(self, chunked=True):
def pretty_sigmas(sigmas):
return ", ".join(f"{sig:.4}" for sig in sigmas.tolist())
print(f"** Dumping restart sampling plan (total steps {total_steps}):")
def dump_steps(step, sigmas, restart=0):
rlabel = f"R{restart:>3}" if restart > 0 else " "
if chunked:
chunk_size = len(sigmas) - 2
step += 1
print(
f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigmas)}",
)
step += chunk_size
return step
for i in range(len(sigmas) - 1):
step += 1
print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigmas[i:i+2])}")
return step
print(f"** Dumping restart sampling plan (total steps {self.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 " "
last_kidx = kidx
if not self.chunked:
for i in range(len(sigs) - 1):
step += 1
print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigs[i:i+2])}")
return x
chunk_size = len(sigs) - 2
step += 1
print(
f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigs)}",
)
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)
for pi in self.plan:
step = dump_steps(step, pi.sigmas)
for kidx in range(pi.k):
step = dump_steps(step, pi.restart_sigmas, kidx + 1)
print(
"** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.",
)
@staticmethod
def self_test(
model,
schedules=None,
restart_schedules=None,
segments=None,
min_steps=2,
max_steps=100,
) -> None:
if schedules is None:
schedules = NORMAL_SCHEDULER_MAPPING.keys()
if restart_schedules is None:
restart_schedules = RESTART_SCHEDULER_MAPPING.keys()
if segments is None:
segments = ("default", "a1111")
for schname in schedules:
for rschname in restart_schedules:
for tsegs in segments:
print(
f"--- Test: {min_steps}..{max_steps} steps, schedules {schname}/{rschname}, segments {tsegs}",
)
for tsteps in range(min_steps, max_steps + 1):
label = f"** {tsteps:03}: {schname}, {rschname}, {tsegs}:"
try:
_plan = RestartPlan(
model,
tsteps,
schname,
tsegs,
rschname,
1.0,
)
except ValueError as err:
print(f"{label}\n\t!! FAIL: {err}")
raise
continue
print("\n|| Done test")
class RestartSampler:
@staticmethod
def get_segment(sigmas: torch.Tensor) -> torch.Tensor:
# A normal segment ends when we either reach the end of the list or
# encounter a sigma higher than the previous.
last_sigma = sigmas[0]
for idx in range(1, len(sigmas)):
sigma = sigmas[idx]
if sigma > last_sigma:
return sigmas[:idx]
last_sigma = sigma
return sigmas
@classmethod
def split_sigmas(cls, sigmas):
# This function just splits the sigmas into chunks that are sorted descending.
# If the first sigma of a chunk is > the last sigma of the previous chunk then this
# is a restart segment: noising the restart uses s_min=prev_chunk[-1], s_max=chunk[0].
# It's a generator that yields tuples of (noise_scale, chunk_sigmas).
prev_seg = None
while len(sigmas) > 1:
seg = cls.get_segment(sigmas)
sigmas = sigmas[len(seg) :]
if prev_seg is not None and seg[0] > prev_seg[-1]:
s_min, s_max = prev_seg[-1], seg[0]
noise_scale = ((s_max**2 - s_min**2) ** 0.5).item()
else:
noise_scale = 0.0
prev_seg = seg
yield (noise_scale, seg)
# Some extra explanation for a couple of these arguments:
#
# restart_chunked:
# When False, the sampling function is called step-by-step with only two sigmas at a time.
# When 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.
#
# restart_custom_noise:
# 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.
@classmethod
@torch.no_grad()
def ksampler_restart_wrapper(
self,
def sampler_function(
cls,
model,
x,
sigmas,
*args,
extra_args=None,
*args: list,
restart_wrapped_sampler=None,
restart_chunked=True,
restart_custom_noise=None,
callback=None,
disable=None,
**kwargs,
):
ksampler = self.ksampler
**kwargs: dict,
) -> torch.Tensor:
if not restart_wrapped_sampler:
raise ValueError("RestartSampler: missing restart_sampler option!")
def restart_noise(x, _s_min, _s_max, _seed):
return lambda _s, _sn: torch.randn_like(x)
seed = (kwargs.get("extra_args", {}) or {}).get("seed", 42)
if restart_custom_noise is not None:
restart_noise = restart_custom_noise
sampler = restart_wrapped_sampler.sampler_function
chunks = tuple(cls.split_sigmas(sigmas))
total_steps = sum(len(chunk) - 1 for _noise, chunk in chunks)
step = 0
seed = self.seed
plan, self.total_steps = self.build_plan(sigmas, x.device)
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:
return noise_sampler
result = self.make_noise_sampler(x, s_min, s_max, seed)
seed += 1
return result
with trange(self.total_steps, disable=disable) as pbar:
noise_count = 0
with trange(total_steps, disable=disable) as pbar:
last_cb_sigma = None
def callback_wrapper(cb_state):
def cb_wrapper(cb_state):
nonlocal step, last_cb_sigma
curr_sigma = cb_state.get("sigma")
curr_sigma = (
@@ -477,33 +665,28 @@ class KSamplerRestartWrapper:
if callback is not None:
callback(cb_state)
# Convenience function for code reuse.
def sampler_function(x, sigs):
return ksampler.sampler_function(
def do_sample(x, sigmas):
return sampler(
model,
x,
sigs,
sigmas,
*args,
extra_args=extra_args,
callback=callback_wrapper,
callback=cb_wrapper,
disable=True,
**kwargs,
)
def do_sample(x, sigs, _kidx=-1):
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)
for noise_scale, chunk_sigmas in chunks:
if noise_scale != 0:
s_min, s_max = chunk_sigmas[-1], chunk_sigmas[0]
x += (
restart_noise(x, s_min, s_max, seed + noise_count)(s_max, s_min)
* noise_scale
)
noise_count += 1
if restart_chunked:
x = do_sample(x, chunk_sigmas)
continue
for i in range(len(chunk_sigmas) - 1):
x = do_sample(x, chunk_sigmas[i : i + 2])
return x
+28 -2
View File
@@ -1,3 +1,4 @@
import comfy
import torch
from comfy.k_diffusion import sampling as k_diffusion_sampling
@@ -7,7 +8,7 @@ from comfy.k_diffusion import sampling as k_diffusion_sampling
# These two may be wrong for v-pred... but it seems to work?
# Copied from k_diffusion
def sigma_to_t(ms, sigma, quantize=True):
log_sigmas = ms.log_sigmas
log_sigmas = ms.log_sigmas.cpu()
log_sigma = sigma.log()
dists = log_sigma - log_sigmas[:, None]
if quantize:
@@ -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,
}