Merge pull request #18 from blepping/feat_restart_sampler
Make restart a sampler, add node to generate sigmas
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user