The struggle!

This commit is contained in:
blepping
2024-04-17 09:35:44 -06:00
parent e2dcfd4091
commit bda63be15c
3 changed files with 483 additions and 305 deletions
+28 -50
View File
@@ -1,12 +1,10 @@
import comfy
import torch
from . import restart_sampling as restart
from .restart_sampling import (
DEFAULT_SEGMENTS,
SCHEDULER_MAPPING,
KSamplerRestartWrapper,
rebuild_plan,
RestartPlan,
restart_sampling,
)
@@ -328,13 +326,6 @@ class RestartScheduler:
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/schedulers"
@staticmethod
def plan_sigmas(plan): # noqa: ANN205
for pi in plan:
yield pi.sigmas
for _ in range(pi.k):
yield pi.restart_sigmas
def go(
self,
model,
@@ -345,38 +336,19 @@ class RestartScheduler:
denoise,
sigmas_opt=None,
):
ms = model.get_model_object("model_sampling")
if sigmas_opt is None or len(sigmas_opt) < 2:
total_steps = steps
if denoise < 1.0:
if denoise <= 0.0:
return (torch.FloatTensor([]),)
total_steps = int(steps / denoise)
# RestartPlan.self_test(model, max_steps=200)
sigmas = restart.calc_sigmas(
scheduler,
total_steps,
float(ms.sigma_min),
float(ms.sigma_max),
model.model,
"cpu",
)
sigmas = sigmas[-(steps + 1) :]
else:
sigmas = sigmas_opt
prepared_segments = restart.prepare_restart_segments(segments, ms, sigmas)
plan, restart_steps = restart.build_plan(
model.model,
prepared_segments,
plan = RestartPlan(
model,
steps,
scheduler,
segments,
restart_scheduler,
sigmas,
"cpu",
denoise,
sigmas=sigmas_opt,
)
if restart.VERBOSE:
restart.explain_plan(plan, restart_steps, chunked=True)
restart_sigmas = torch.flatten(torch.cat(tuple(self.plan_sigmas(plan))))
print("MADE SIGMAS", restart_sigmas)
return (restart_sigmas,)
plan.explain(chunked=True)
return (plan.sigmas(),)
class RestartSampler:
@@ -408,19 +380,25 @@ class RestartSampler:
@staticmethod
@torch.no_grad()
def sampler_function(wrapped, chunked, model, x, sigmas, *args, **kwargs):
plan, total_steps = rebuild_plan(sigmas)
print("Rebuilt", total_steps, plan)
seed = kwargs.get("extra_args", {}).get("seed")
rw = KSamplerRestartWrapper(
def sampler_function(
wrapped,
chunked,
model,
x,
sigmas,
*args: list,
**kwargs: dict,
) -> torch.Tensor:
plan = RestartPlan.from_sigmas(sigmas)
return plan.sample(
wrapped,
None,
None,
None,
seed,
chunked=chunked,
model,
x,
sigmas,
*args,
restart_chunked=chunked,
**kwargs,
)
return rw.sample_plan(plan, total_steps, model, x, sigmas, *args, **kwargs)
NODE_CLASS_MAPPINGS = {
+454 -254
View File
@@ -1,3 +1,5 @@
from __future__ import annotations
import ast
import os
import warnings
@@ -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,46 @@ 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)
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 = []
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)
@@ -158,53 +179,19 @@ def restart_sampling(
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,
)
else:
sigmas = sigmas.detach().clone().to(model.load_device)
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:]
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,
)
sampler_wrapper = KSamplerRestartWrapper(
sampler,
real_model,
restart_scheduler,
restart_segments,
seed,
custom_noise,
chunked=chunked_mode,
denoise=denoise,
step_range=step_range,
force_full_denoise=force_full_denoise,
sigmas=sigmas,
)
plan = plan.to(model.load_device)
sigmas = plan.sigmas()
latent = latent_image
latent_image = latent["samples"]
@@ -227,12 +214,23 @@ 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,
sigmas.shape[-1] - 1,
x0_output,
)
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
sampler = KSAMPLER(
sampler_wrapper.ksampler_restart_wrapper,
ksampler = KSAMPLER(
lambda *args, **kwargs: plan.sample(
sampler,
*args,
restart_chunked=chunked_mode,
restart_make_noise_sampler=custom_noise,
restart_seed=seed,
**kwargs,
),
extra_options=sampler.extra_options | {},
inpaint_options=sampler.inpaint_options | {},
)
@@ -241,7 +239,7 @@ def restart_sampling(
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)
pbar_update_absolute(self, value, plan.total_steps, preview)
ProgressBar.update_absolute = pbar_update_absolute_wrapper
@@ -250,7 +248,7 @@ def restart_sampling(
model,
noise,
cfg,
sampler,
ksampler,
sigmas,
positive,
negative,
@@ -285,6 +283,42 @@ class PlanItem(
defaults=[None, 0, 0.0, 0.0, None],
),
):
def __new__(cls, *args: list, **kwargs: dict):
threshold = 1e-06
obj = super().__new__(cls, *args, **kwargs)
if len(obj.sigmas) < 2:
raise ValueError("PlanItem: invalid normal sigmas: too short")
if obj.k < 1:
return obj
if len(obj.restart_sigmas) < 2:
raise ValueError("PlanItem: invalid restart sigmas: too short")
if obj.s_min >= obj.s_max:
raise ValueError("PlanItem: invalid min/max: min >= max")
# if obj.sigmas[-1] >= obj.restart_sigmas[0]:
if obj.sigmas[-1] - obj.restart_sigmas[0] > threshold:
raise ValueError(
"PlanItem: invalid sigmas: last normal sigma >= first restart sigma",
)
# if obj.restart_sigmas[-1] < obj.sigmas[-1]:
if obj.sigmas[-1] - obj.restart_sigmas[-1] > 1e-02: # threshold:
errstr = (
f"PlanItem: invalid sigmas: last restart sigma {obj.restart_sigmas[-1]} < last normal sigma {obj.sigmas[-1]}",
)
raise ValueError(errstr)
t = obj.sigmas.sort(descending=True, stable=True)[0].unique_consecutive()
if not torch.equal(obj.sigmas, t):
raise ValueError(
"PlanItem: invalid normal sigmas: out of order or contains duplicates",
)
t = obj.restart_sigmas.sort(descending=True, stable=True)[
0
].unique_consecutive()
if not torch.equal(obj.restart_sigmas, t):
raise ValueError(
"PlanItem: invalid restart sigmas: out of order or contains duplicates",
)
return obj
# 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.
@@ -314,141 +348,281 @@ class PlanItem(
return x
# Builds a list of PlanItems and calculates the total number of steps. See the comments for PlanItem
# for more information about plans.
# Returns two values: the plan and the total steps.
@torch.no_grad()
def build_plan(model, restart_segments, restart_scheduler, sigmas, device):
segments = round_restart_segments(sigmas, restart_segments)
total_steps = len(sigmas) - 1 + calc_restart_steps(segments)
plan = []
range_start = -1
for i in range(len(sigmas) - 1):
if range_start == -1:
# Starting a new plan item - main sigmas start at the current index of i.
range_start = i
s_min = sigmas[i + 1].item()
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(
restart_scheduler,
n_restart,
s_min,
s_max,
model,
device=device,
class RestartPlan:
def __init__(
self,
model,
steps,
scheduler,
restart_info,
restart_scheduler,
denoise=1.0,
step_range=None,
force_full_denoise=False,
sigmas=None,
):
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)
)
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.
plan.append(PlanItem(sigmas[range_start:]))
return plan, total_steps
def rebuild_plan(sigmas):
def get_normal_segment(sigmas):
last_sigma = None
for idx in range(len(sigmas)):
sigma = sigmas[idx]
if last_sigma is not None and sigma >= last_sigma:
return sigmas[:idx]
last_sigma = sigma
return sigmas
def get_restart_segment(sigmas, s_min):
last_sigma = None
for idx in range(len(sigmas)):
sigma = sigmas[idx]
if (last_sigma is not None and sigma >= last_sigma) or sigma < s_min:
return sigmas[:idx]
last_sigma = sigma
raise ValueError("Unexpected end of sigmas in a restart segment")
plan = []
total_steps = 0
while len(sigmas) > 0:
normal_sigmas = get_normal_segment(sigmas)
nslen = len(normal_sigmas)
sigmas = sigmas[nslen:]
total_steps += nslen - 1
if len(sigmas) == 0:
plan.append(PlanItem(normal_sigmas))
break
restart_sigmas = get_restart_segment(sigmas, normal_sigmas[-1])
rslen = len(restart_sigmas)
sigmas = sigmas[rslen:]
k = 1
while len(sigmas) > 0 and torch.equal(sigmas[:rslen], restart_sigmas):
k += 1
sigmas = sigmas[rslen:]
total_steps += (rslen - 1) * k
plan.append(
PlanItem(
normal_sigmas,
k,
normal_sigmas[-1],
restart_sigmas[0],
restart_sigmas,
),
)
return plan, total_steps
# Dumps information about the plan to the console. It uses the normal plan execute
# logic.
def explain_plan(plan, total_steps, 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}):")
for pi in plan:
print(
f"\n{pi.sigmas[-1].item():.04} .. {pi.sigmas[0].item():.04} ({len(pi.sigmas)})",
)
if pi.k > 0:
print(
f" {pi.restart_sigmas[-1].item():.04} ({pi.s_min:.04}) .. {pi.restart_sigmas[0].item():.04} ({pi.s_max:.04}): k={pi.k} ({len(pi.restart_sigmas)})",
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,
"cpu",
# model.load_device,
)
step = 0
last_kidx = -1
else:
sigmas = sigmas.detach().cpu().clone()
if step_range is not None:
start_step, last_step = step_range
# 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 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)}",
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) :]
self.plain_sigmas = sigmas
restart_segments = prepare_restart_segments(
restart_info,
real_model.model_sampling,
sigmas,
)
self.plan, self.total_steps = self.build_plan_items(
model.model,
restart_segments,
restart_scheduler,
sigmas,
"cpu",
)
step += chunk_size
return x
# Stub function to satisfy PlanItem.execute
def get_noise_sampler(*_args):
return lambda *_args: 0.0
def __repr__(self) -> str:
return f"<RestartPlan: steps={self.total_steps}, plan={self.plan}>"
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.",
)
def __len__(self) -> int:
return self.total_steps
# 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_items(
model,
restart_segments,
restart_scheduler,
sigmas,
device,
) -> tuple[list, int]:
segments = round_restart_segments(sigmas, restart_segments)
total_steps = len(sigmas) - 1 + calc_restart_steps(segments)
plan = []
range_start = -1
for i in range(len(sigmas) - 1):
if range_start == -1:
# Starting a new plan item - main sigmas start at the current index of i.
range_start = i
s_min = sigmas[i + 1].item()
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(
restart_scheduler,
n_restart,
s_min,
s_max,
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.
plan.append(PlanItem(sigmas[range_start:]))
return plan, total_steps
@classmethod
def from_sigmas(cls, sigmas, threshold=1e-06):
def get_normal_segment(sigmas):
# 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 last_sigma - sigma < threshold:
return sigmas[:idx]
last_sigma = sigma
return sigmas
def get_restart_segment(sigmas, s_min):
# s_min here is the last sigma of the previous normal segment. A restart segment
# ends when:
# 1. We reach the end of the list, or
# 2. We hit a sigma greater or equal to the last sigma, or
# 3. We hit a sigma less than s_min
last_sigma = sigmas[0]
if s_min - last_sigma > threshold:
return sigmas[:2]
for idx in range(1, len(sigmas)):
sigma = sigmas[idx]
# sigma > last_sigma
if last_sigma - sigma < threshold:
return sigmas[:idx]
# sigma < s_min
if s_min - sigma > -threshold:
# TODO: Document this part
if idx < len(sigmas) - 2 and sigmas[idx + 1] - s_min < -threshold:
return sigmas[:idx]
return sigmas[: idx + 1]
last_sigma = sigma
raise ValueError("Unexpected end of sigmas in a restart segment")
plain_sigmas = sigmas.detach().cpu().clone()
plan = []
total_steps = 0
while len(sigmas) > 0:
# Get the normal segment - a restart segment can never be first.
normal_sigmas = get_normal_segment(sigmas)
nslen = len(normal_sigmas)
if nslen < 2:
print(sigmas)
raise ValueError(
"Encountered invalid normal segment rebuilding sigmas: too short",
)
sigmas = sigmas[nslen:]
total_steps += nslen - 1
if len(sigmas) == 0:
# No restart segments follow the normal segment so we're done.
plan.append(PlanItem(normal_sigmas))
break
# If we're here there has to be a restart segment; get it.
restart_sigmas = get_restart_segment(sigmas, normal_sigmas[-1])
rslen = len(restart_sigmas)
if rslen < 2:
print(restart_sigmas)
raise ValueError(
"Encountered invalid normal segment rebuilding sigmas: too short",
)
sigmas = sigmas[rslen:]
k = 1
# The restart segment may be repeated multiple times. If so, count the
# repeats and trim the sigmas list.
while len(sigmas) > 0 and torch.equal(sigmas[:rslen], restart_sigmas):
k += 1
sigmas = sigmas[rslen:]
total_steps += (rslen - 1) * k
plan.append(
PlanItem(
normal_sigmas,
k,
normal_sigmas[-1],
restart_sigmas[0],
restart_sigmas,
),
)
obj = cls.__new__(cls)
obj.plan = plan
obj.total_steps = total_steps
obj.plain_sigmas = plain_sigmas
return obj
def sigmas(self) -> torch.Tensor:
def sigmas_generator():
for pi in self.plan:
yield pi.sigmas.cpu()
for _ in range(pi.k):
yield pi.restart_sigmas.cpu()
return torch.flatten(torch.cat(tuple(sigmas_generator())))
def to(self, device):
obj = self.__class__.__new__(self.__class__)
obj.plain_sigmas = self.plain_sigmas.to(device)
obj.total_steps = self.total_steps
items = obj.plan = []
for pi in self.plan:
sigmas = pi.sigmas.to(device)
if pi.k < 1:
items.append(PlanItem(sigmas))
continue
items.append(
PlanItem(
sigmas,
pi.k,
pi.s_min,
pi.s_max,
pi.restart_sigmas.to(device),
),
)
return obj
# Dumps information about the plan to the console. It uses the normal plan execute
# logic.
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 {self.total_steps}):")
for pi in self.plan:
print(
f"\n{pi.sigmas[-1].item():.04} .. {pi.sigmas[0].item():.04} ({len(pi.sigmas)})",
)
if pi.k > 0:
print(
f" {pi.restart_sigmas[-1].item():.04} ({pi.s_min:.04}) .. {pi.restart_sigmas[0].item():.04} ({pi.s_max:.04}): k={pi.k} ({len(pi.restart_sigmas)})",
)
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 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: list):
return lambda *_args: 0.0
for pi in self.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.",
)
class KSamplerRestartWrapper:
# Some extra explanation for a couple of these arguments:
#
# chunked:
@@ -461,48 +635,33 @@ 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,
):
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
@torch.no_grad()
def sample_plan(
def sample(
self,
plan,
total_steps,
ksampler,
model,
x,
sigmas,
*args,
_sigmas,
*args: list,
restart_chunked=True,
restart_make_noise_sampler=None,
restart_seed=None,
extra_args=None,
callback=None,
disable=None,
**kwargs,
**kwargs: dict,
):
self.total_steps = total_steps
ksampler = self.ksampler
step = 0
seed = self.seed
if restart_seed is None:
seed = (extra_args or {}).get("seed", 42)
else:
seed = restart_seed
plan = self.plan
if VERBOSE:
explain_plan(plan, self.total_steps, chunked=self.chunked)
self.explain(restart_chunked)
def noise_sampler(*_args):
def noise_sampler(*_args: list):
return torch.randn_like(x)
# Passed to the PlanItem .execute method. Most of the time, self.make_noise_sampler
@@ -511,9 +670,9 @@ class KSamplerRestartWrapper:
# don't all use the same noise.
def get_noise_sampler(x, s_min, s_max):
nonlocal seed
if not self.make_noise_sampler:
if not restart_make_noise_sampler:
return noise_sampler
result = self.make_noise_sampler(x, s_min, s_max, seed)
result = restart_make_noise_sampler(x, s_min, s_max, seed)
seed += 1
return result
@@ -554,7 +713,7 @@ class KSamplerRestartWrapper:
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 restart_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)
@@ -566,37 +725,78 @@ class KSamplerRestartWrapper:
# Execute the plan items in sequence.
for pi in plan:
x = pi.execute(x, do_sample, get_noise_sampler)
return x
@torch.no_grad()
def ksampler_restart_wrapper(
self,
@staticmethod
def self_test(
model,
x,
sigmas,
*args,
extra_args=None,
callback=None,
disable=None,
**kwargs,
):
plan, total_steps = build_plan(
self.real_model,
self.restart_segments,
self.restart_scheduler,
sigmas,
x.device,
)
return self.sample_plan(
plan,
total_steps,
model,
x,
sigmas,
*args,
extra_args=extra_args,
callback=callback,
disable=disable,
**kwargs,
)
schedules=None,
restart_schedules=None,
segments=None,
min_steps=2,
max_steps=100,
) -> None:
if schedules is None:
schedules = SCHEDULER_MAPPING.keys() - {"simple_test"}
if restart_schedules is None:
restart_schedules = SCHEDULER_MAPPING.keys() - {"simple_test"}
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:
p1 = RestartPlan(
model,
tsteps,
schname,
tsegs,
rschname,
1.0,
)
except ValueError as err:
print(f"{label}\n\t!! FAIL: {err}")
continue
try:
p2 = RestartPlan.from_sigmas(p1.sigmas())
except ValueError:
print(label)
p1.explain(chunked=True)
raise
fail = None
if len(p1) != len(p2):
fail = "steps"
if not fail:
for idx in range(len(p1.plan)):
pi1, pi2 = p1.plan[idx], p2.plan[idx]
if not torch.equal(
torch.round(pi1.sigmas, decimals=5),
torch.round(pi2.sigmas, decimals=5),
):
fail = "normal"
break
if pi1.k != pi2.k:
fail = "k"
break
if pi1.k < 1:
continue
if not torch.equal(
torch.round(pi1.restart_sigmas, decimals=5),
torch.round(pi2.restart_sigmas, decimals=5),
):
fail = "restart"
break
if fail:
print(label)
print("!!!", fail)
p1.explain()
print("====")
p2.explain()
raise ValueError("Failed rebuilding restart plan")
print("\n|| Done test")
+1 -1
View File
@@ -7,7 +7,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: