Code formatting + clean up some lints
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import comfy
|
||||
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING, DEFAULT_SEGMENTS
|
||||
|
||||
from .restart_sampling import DEFAULT_SEGMENTS, SCHEDULER_MAPPING, restart_sampling
|
||||
|
||||
|
||||
def get_supported_samplers():
|
||||
@@ -27,129 +28,273 @@ def get_supported_restart_schedulers():
|
||||
|
||||
class KRestartSamplerSimple:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"model": ("MODEL",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS,),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"latent_image": ("LATENT",),
|
||||
"denoise": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"segments": ("STRING", {"default": "default", "multiline": False}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments):
|
||||
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise)
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
denoise,
|
||||
segments,
|
||||
):
|
||||
return restart_sampling(
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
scheduler,
|
||||
denoise=denoise,
|
||||
)
|
||||
|
||||
|
||||
class KRestartSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"model": ("MODEL",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (tuple(SCHEDULER_MAPPING.keys()), ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"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}),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(), ),
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"scheduler": (tuple(SCHEDULER_MAPPING.keys()),),
|
||||
"positive": ("CONDITIONING",),
|
||||
"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},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=True):
|
||||
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise, chunked_mode=chunked_mode)
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
denoise,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
chunked_mode=True,
|
||||
):
|
||||
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:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"add_noise": (["enable", "disable"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"add_noise": (["enable", "disable"],),
|
||||
"noise_seed": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF},
|
||||
),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (tuple(SCHEDULER_MAPPING.keys()), ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"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(), ),
|
||||
"return_with_leftover_noise": (["disable", "enable"],),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, chunked_mode=True):
|
||||
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=True,
|
||||
):
|
||||
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, chunked_mode=chunked_mode)
|
||||
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:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"add_noise": (["enable", "disable"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"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", ),
|
||||
"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(), ),
|
||||
"return_with_leftover_noise": (["disable", "enable"],),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT","LATENT")
|
||||
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, chunked_mode=True):
|
||||
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=True,
|
||||
):
|
||||
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, 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,
|
||||
|
||||
+175
-50
@@ -1,14 +1,16 @@
|
||||
import ast
|
||||
from collections import namedtuple
|
||||
import os
|
||||
import warnings
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
import latent_preview
|
||||
from collections import namedtuple
|
||||
|
||||
import comfy
|
||||
from comfy.sample import sample_custom, prepare_noise
|
||||
import latent_preview
|
||||
import torch
|
||||
from comfy.sample import prepare_noise, sample_custom
|
||||
from comfy.samplers import KSAMPLER, sampler_object
|
||||
from comfy.utils import ProgressBar
|
||||
from tqdm.auto import trange
|
||||
|
||||
from .restart_schedulers import SCHEDULER_MAPPING
|
||||
|
||||
VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1"
|
||||
@@ -19,7 +21,7 @@ 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:
|
||||
restart_segments = []
|
||||
restart_segments.append({'n': n_restart, 'k': k, 't_min': t_min, 't_max': t_max})
|
||||
restart_segments.append({"n": n_restart, "k": k, "t_min": t_min, "t_max": t_max})
|
||||
return restart_segments
|
||||
|
||||
|
||||
@@ -63,9 +65,9 @@ def prepare_restart_segments(restart_info, ms, sigmas):
|
||||
if restart_arrays is None:
|
||||
try:
|
||||
restart_arrays = ast.literal_eval(f"[{restart_info}]")
|
||||
except SyntaxError as e:
|
||||
except SyntaxError:
|
||||
print("Ill-formed restart segments")
|
||||
raise e
|
||||
raise
|
||||
restart_segments = []
|
||||
for arr in restart_arrays:
|
||||
if len(arr) != 4:
|
||||
@@ -74,7 +76,13 @@ def prepare_restart_segments(restart_info, ms, sigmas):
|
||||
n_restart, k = int(n_restart), int(k)
|
||||
t_min = resolve_t_value(val_min, ms)
|
||||
t_max = resolve_t_value(val_max, ms)
|
||||
restart_segments = add_restart_segment(restart_segments, n_restart, k, t_min, t_max)
|
||||
restart_segments = add_restart_segment(
|
||||
restart_segments,
|
||||
n_restart,
|
||||
k,
|
||||
t_min,
|
||||
t_max,
|
||||
)
|
||||
return restart_segments
|
||||
|
||||
|
||||
@@ -86,20 +94,32 @@ def round_restart_segments(ts, restart_segments):
|
||||
:return: dict of the form {nearest_t_min: {'n': n, 'k': k, 't_max': t_max}}
|
||||
"""
|
||||
t_min_mapping = {}
|
||||
for segment in reversed(restart_segments): # Reversed to prioritize segments to the front
|
||||
t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment['t_min'])).item()
|
||||
for segment in reversed(
|
||||
restart_segments,
|
||||
): # Reversed to prioritize segments to the front
|
||||
t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment["t_min"])).item()
|
||||
if t_min_neighbor == ts[0]:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...", stacklevel=2)
|
||||
f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...",
|
||||
stacklevel=2,
|
||||
)
|
||||
continue
|
||||
if t_min_neighbor > segment['t_max']:
|
||||
if t_min_neighbor > segment["t_max"]:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...", stacklevel=2)
|
||||
f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...",
|
||||
stacklevel=2,
|
||||
)
|
||||
continue
|
||||
if t_min_neighbor in t_min_mapping:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}", stacklevel=2)
|
||||
t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']}
|
||||
f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}",
|
||||
stacklevel=2,
|
||||
)
|
||||
t_min_mapping[t_min_neighbor] = {
|
||||
"n": segment["n"],
|
||||
"k": segment["k"],
|
||||
"t_max": segment["t_max"],
|
||||
}
|
||||
return t_min_mapping
|
||||
|
||||
|
||||
@@ -110,11 +130,31 @@ def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device):
|
||||
def calc_restart_steps(restart_segments):
|
||||
restart_steps = 0
|
||||
for segment in restart_segments.values():
|
||||
restart_steps += (segment['n'] - 1) * segment['k']
|
||||
restart_steps += (segment["n"] - 1) * segment["k"]
|
||||
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=True, sigmas=None):
|
||||
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,
|
||||
sigmas=None,
|
||||
):
|
||||
if isinstance(sampler, str):
|
||||
sampler = sampler_object(sampler)
|
||||
|
||||
@@ -123,11 +163,17 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
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)
|
||||
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,
|
||||
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,
|
||||
)
|
||||
else:
|
||||
sigmas = sigmas.detach().clone().to(model.load_device)
|
||||
@@ -135,26 +181,45 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
start_step, last_step = step_range
|
||||
|
||||
if last_step < (len(sigmas) - 1):
|
||||
sigmas = sigmas[:last_step + 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):]
|
||||
sigmas = sigmas[-(steps + 1) :]
|
||||
|
||||
restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling, sigmas)
|
||||
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)
|
||||
sampler_wrapper = KSamplerRestartWrapper(
|
||||
sampler,
|
||||
real_model,
|
||||
restart_scheduler,
|
||||
restart_segments,
|
||||
seed,
|
||||
custom_noise,
|
||||
chunked=chunked_mode,
|
||||
)
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
if disable_noise:
|
||||
torch.manual_seed(seed) # workaround for https://github.com/comfyanonymous/ComfyUI/issues/2833
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
torch.manual_seed(
|
||||
seed,
|
||||
) # workaround for https://github.com/comfyanonymous/ComfyUI/issues/2833
|
||||
noise = torch.zeros(
|
||||
latent_image.size(),
|
||||
dtype=latent_image.dtype,
|
||||
layout=latent_image.layout,
|
||||
device="cpu",
|
||||
)
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
batch_inds = latent.get("batch_index", None)
|
||||
noise = prepare_noise(latent_image, seed, batch_inds)
|
||||
|
||||
noise_mask = None
|
||||
@@ -167,11 +232,11 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
|
||||
sampler = KSAMPLER(
|
||||
sampler_wrapper.ksampler_restart_wrapper, extra_options=sampler.extra_options | {},
|
||||
sampler_wrapper.ksampler_restart_wrapper,
|
||||
extra_options=sampler.extra_options | {},
|
||||
inpaint_options=sampler.inpaint_options | {},
|
||||
)
|
||||
|
||||
|
||||
# Add the additional steps to the progress bar
|
||||
pbar_update_absolute = ProgressBar.update_absolute
|
||||
|
||||
@@ -182,8 +247,19 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
|
||||
try:
|
||||
samples = sample_custom(
|
||||
model, noise, cfg, sampler, sigmas, positive, negative, latent_image,
|
||||
noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed)
|
||||
model,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
noise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
finally:
|
||||
ProgressBar.update_absolute = pbar_update_absolute
|
||||
|
||||
@@ -202,7 +278,13 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], defaults=[None, 0, 0., 0., None])):
|
||||
class PlanItem(
|
||||
namedtuple(
|
||||
"PlanItem",
|
||||
["sigmas", "k", "s_min", "s_max", "restart_sigmas"],
|
||||
defaults=[None, 0, 0.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.
|
||||
@@ -224,7 +306,10 @@ class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart
|
||||
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
|
||||
x += (
|
||||
noise_sampler(self.restart_sigmas[0], self.restart_sigmas[-1])
|
||||
* (self.s_max**2 - self.s_min**2) ** 0.5
|
||||
)
|
||||
x = sample(x, self.restart_sigmas, kidx)
|
||||
return x
|
||||
|
||||
@@ -242,7 +327,16 @@ class KSamplerRestartWrapper:
|
||||
# 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):
|
||||
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
|
||||
self.restart_scheduler = restart_scheduler
|
||||
@@ -269,10 +363,18 @@ class KSamplerRestartWrapper:
|
||||
seg = segments.get(s_min)
|
||||
if seg is None:
|
||||
continue
|
||||
s_max, k, n_restart = seg['t_max'], seg['k'], seg['n']
|
||||
seg_sigmas = calc_sigmas(self.restart_scheduler, n_restart, s_min,
|
||||
s_max, self.real_model, device=device)
|
||||
plan.append(PlanItem(sigmas[range_start:i+2], k, s_min, s_max, seg_sigmas[:-1]))
|
||||
s_max, k, n_restart = seg["t_max"], seg["k"], seg["n"]
|
||||
seg_sigmas = calc_sigmas(
|
||||
self.restart_scheduler,
|
||||
n_restart,
|
||||
s_min,
|
||||
s_max,
|
||||
self.real_model,
|
||||
device=device,
|
||||
)
|
||||
plan.append(
|
||||
PlanItem(sigmas[range_start : i + 2], k, s_min, s_max, seg_sigmas[:-1]),
|
||||
)
|
||||
range_start = -1
|
||||
if range_start != -1:
|
||||
# Include sigmas after the last restart segments in the plan.
|
||||
@@ -284,9 +386,11 @@ class KSamplerRestartWrapper:
|
||||
def explain_plan(self, plan, total_steps):
|
||||
def pretty_sigmas(sigmas):
|
||||
return ", ".join(f"{sig:.4}" for sig in sigmas.tolist())
|
||||
|
||||
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.
|
||||
@@ -295,13 +399,15 @@ class KSamplerRestartWrapper:
|
||||
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):
|
||||
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)}")
|
||||
print(
|
||||
f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigs)}",
|
||||
)
|
||||
step += chunk_size
|
||||
return x
|
||||
|
||||
@@ -311,11 +417,22 @@ class KSamplerRestartWrapper:
|
||||
|
||||
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.")
|
||||
|
||||
print(
|
||||
"** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.",
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs):
|
||||
def ksampler_restart_wrapper(
|
||||
self,
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args,
|
||||
extra_args=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
**kwargs,
|
||||
):
|
||||
ksampler = self.ksampler
|
||||
step = 0
|
||||
seed = self.seed
|
||||
@@ -340,6 +457,7 @@ class KSamplerRestartWrapper:
|
||||
return result
|
||||
|
||||
with trange(self.total_steps, disable=disable) as pbar:
|
||||
|
||||
def callback_wrapper(x):
|
||||
nonlocal step
|
||||
step += 1
|
||||
@@ -351,10 +469,17 @@ class KSamplerRestartWrapper:
|
||||
# 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,
|
||||
**kwargs)
|
||||
model,
|
||||
x,
|
||||
sigs,
|
||||
*args,
|
||||
extra_args=extra_args,
|
||||
callback=callback_wrapper,
|
||||
disable=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def do_sample(x, sigs, kidx=-1):
|
||||
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:
|
||||
@@ -362,8 +487,8 @@ class KSamplerRestartWrapper:
|
||||
# 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])
|
||||
for i in range(len(sigs) - 1):
|
||||
x = sampler_function(x, sigs[i : i + 2])
|
||||
return x
|
||||
|
||||
# Execute the plan items in sequence.
|
||||
|
||||
+15
-14
@@ -1,7 +1,9 @@
|
||||
import torch
|
||||
from comfy.k_diffusion import sampling as k_diffusion_sampling
|
||||
|
||||
# from comfy.samplers import normal_scheduler
|
||||
|
||||
|
||||
# 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):
|
||||
@@ -18,6 +20,7 @@ def sigma_to_t(ms, sigma, quantize=True):
|
||||
t = (1 - w) * low_idx + w * high_idx
|
||||
return t.view(sigma.shape)
|
||||
|
||||
|
||||
# Copied from k_diffusion
|
||||
def t_to_sigma(ms, t):
|
||||
t = t.float()
|
||||
@@ -25,15 +28,16 @@ def t_to_sigma(ms, t):
|
||||
log_sigma = (1 - w) * ms.log_sigmas[low_idx] + w * ms.log_sigmas[high_idx]
|
||||
return log_sigma.exp()
|
||||
|
||||
def get_sigmas_karras(model, n, s_min, s_max, device):
|
||||
|
||||
def get_sigmas_karras(_model, n, s_min, s_max, device):
|
||||
return k_diffusion_sampling.get_sigmas_karras(n, s_min, s_max, device=device)
|
||||
|
||||
|
||||
def get_sigmas_exponential(model, n, s_min, s_max, device):
|
||||
def get_sigmas_exponential(_model, n, s_min, s_max, device):
|
||||
return k_diffusion_sampling.get_sigmas_exponential(n, s_min, s_max, device=device)
|
||||
|
||||
|
||||
def normal_scheduler(model, steps, s_min, s_max, sgm=False, floor=False):
|
||||
def normal_scheduler(model, steps, s_min, s_max, sgm=False):
|
||||
"""
|
||||
Pulled from comfy.samplers.normal_scheduler
|
||||
"""
|
||||
@@ -46,7 +50,7 @@ def normal_scheduler(model, steps, s_min, s_max, sgm=False, floor=False):
|
||||
else:
|
||||
timesteps = torch.linspace(start, end, steps)
|
||||
|
||||
sigs = tuple(ms.sigma(timesteps[x]) for x in range(len(timesteps))) + (0.0,)
|
||||
sigs = (*(ms.sigma(timesteps[x]) for x in range(len(timesteps))), 0.0)
|
||||
return torch.FloatTensor(sigs)
|
||||
|
||||
|
||||
@@ -60,10 +64,12 @@ def get_sigmas_simple(model, n, s_min, s_max, device):
|
||||
max_idx = torch.argmin(torch.abs(ms.sigmas - s_max))
|
||||
sigmas_slice = ms.sigmas[min_idx:max_idx]
|
||||
ss = len(sigmas_slice) / n
|
||||
sigs = [float(s_max)]
|
||||
for x in range(1, n - 1):
|
||||
sigs += [float(sigmas_slice[-(1 + int(x * ss))])]
|
||||
sigs += [float(s_min), 0.0]
|
||||
sigs = (
|
||||
float(s_max),
|
||||
*(float(sigmas_slice[-(1 + int(x * ss))]) for x in range(1, n - 1)),
|
||||
float(s_min),
|
||||
0.0,
|
||||
)
|
||||
return torch.tensor(sigs, device=device)
|
||||
|
||||
|
||||
@@ -71,12 +77,7 @@ def get_sigmas_ddim_uniform(model, n, s_min, s_max, device):
|
||||
ms = model.model_sampling
|
||||
t_min, t_max = sigma_to_t(ms, torch.tensor([s_min, s_max], device=device))
|
||||
ddim_timesteps = torch.linspace(t_max, t_min, n, dtype=torch.int16, device=device)
|
||||
sigs = []
|
||||
for ts in ddim_timesteps:
|
||||
if ts > 999:
|
||||
ts = 999
|
||||
sigs.append(t_to_sigma(ms, ts))
|
||||
sigs += [0.0]
|
||||
sigs = (*(t_to_sigma(ms, min(ts, 999)) for ts in ddim_timesteps), 0.0)
|
||||
return torch.tensor(sigs, device=device)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user