Refactor and simplify
This commit is contained in:
@@ -1,13 +1,20 @@
|
||||
import os
|
||||
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
from .restart_sampling import (
|
||||
DEFAULT_SEGMENTS,
|
||||
SCHEDULER_MAPPING,
|
||||
VERBOSE,
|
||||
RestartPlan,
|
||||
RestartSampler,
|
||||
restart_sampling,
|
||||
)
|
||||
|
||||
INCLUDE_SELFTEST = (
|
||||
os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1"
|
||||
)
|
||||
|
||||
|
||||
def get_supported_samplers():
|
||||
samplers = comfy.samplers.KSampler.SAMPLERS.copy()
|
||||
@@ -302,7 +309,7 @@ class KRestartSamplerCustom:
|
||||
)
|
||||
|
||||
|
||||
class RestartScheduler:
|
||||
class RestartSchedulerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -336,8 +343,6 @@ class RestartScheduler:
|
||||
denoise,
|
||||
sigmas_opt=None,
|
||||
):
|
||||
# RestartPlan.self_test(model, max_steps=200)
|
||||
|
||||
plan = RestartPlan(
|
||||
model,
|
||||
steps,
|
||||
@@ -347,11 +352,12 @@ class RestartScheduler:
|
||||
denoise,
|
||||
sigmas=sigmas_opt,
|
||||
)
|
||||
plan.explain(chunked=True)
|
||||
if VERBOSE:
|
||||
plan.explain(chunked=True)
|
||||
return (plan.sigmas(),)
|
||||
|
||||
|
||||
class RestartSampler:
|
||||
class RestartSamplerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -366,39 +372,16 @@ class RestartSampler:
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
def go(self, sampler, chunked_mode):
|
||||
wrapped = comfy.samplers.KSAMPLER(
|
||||
lambda *args, **kwargs: self.sampler_function(
|
||||
sampler,
|
||||
chunked_mode,
|
||||
*args,
|
||||
**kwargs,
|
||||
),
|
||||
extra_options=sampler.extra_options,
|
||||
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 (wrapped,)
|
||||
|
||||
@staticmethod
|
||||
@torch.no_grad()
|
||||
def sampler_function(
|
||||
wrapped,
|
||||
chunked,
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args: list,
|
||||
**kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
plan = RestartPlan.from_sigmas(sigmas)
|
||||
return plan.sample(
|
||||
wrapped,
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args,
|
||||
restart_chunked=chunked,
|
||||
**kwargs,
|
||||
)
|
||||
return (restart_sampler,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -406,8 +389,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"KRestartSampler": KRestartSampler,
|
||||
"KRestartSamplerAdv": KRestartSamplerAdv,
|
||||
"KRestartSamplerCustom": KRestartSamplerCustom,
|
||||
"RestartScheduler": RestartScheduler,
|
||||
"RestartSampler": RestartSampler,
|
||||
"RestartScheduler": RestartSchedulerNode,
|
||||
"RestartSampler": RestartSamplerNode,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -416,3 +399,28 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"KRestartSamplerAdv": "KSampler With Restarts (Advanced)",
|
||||
"KRestartSamplerCustom": "KSampler With Restarts (Custom)",
|
||||
}
|
||||
|
||||
if os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1":
|
||||
|
||||
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
|
||||
|
||||
+189
-336
@@ -190,19 +190,16 @@ def restart_sampling(
|
||||
force_full_denoise=force_full_denoise,
|
||||
sigmas=sigmas,
|
||||
)
|
||||
plan = plan.to(model.load_device)
|
||||
### UNCOMMENT TO RUN SELF TEST
|
||||
# plan.self_test(
|
||||
# model,
|
||||
# min_steps=2,
|
||||
# schedules=SCHEDULER_MAPPING.keys(),
|
||||
# # schedules=("simple_test",),
|
||||
# restart_schedules=SCHEDULER_MAPPING.keys(),
|
||||
# )
|
||||
sigmas = plan.sigmas()
|
||||
|
||||
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,
|
||||
@@ -222,32 +219,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
|
||||
|
||||
restart_options = {
|
||||
"restart_chunked": chunked_mode,
|
||||
"restart_wrapped_sampler": sampler,
|
||||
"restart_custom_noise": custom_noise,
|
||||
}
|
||||
|
||||
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 | {},
|
||||
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, plan.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
|
||||
|
||||
@@ -352,35 +344,6 @@ class PlanItem(
|
||||
def s_max(self):
|
||||
return None if self.k < 1 else self.restart_sigmas[0].item()
|
||||
|
||||
# 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, skip_normal=False, next_pi=None):
|
||||
if not skip_normal:
|
||||
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
|
||||
)
|
||||
if next_pi and kidx == self.k - 1:
|
||||
sigmas = torch.cat((self.restart_sigmas[:-1], next_pi.sigmas)).to(
|
||||
self.restart_sigmas.device,
|
||||
)
|
||||
print("COMBINE", sigmas)
|
||||
else:
|
||||
sigmas = self.restart_sigmas
|
||||
x = sample(x, sigmas, kidx)
|
||||
return x
|
||||
|
||||
|
||||
class RestartPlan:
|
||||
def __init__(
|
||||
@@ -395,10 +358,6 @@ class RestartPlan:
|
||||
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
|
||||
ms = model.get_model_object("model_sampling")
|
||||
|
||||
effective_steps = (
|
||||
@@ -430,8 +389,6 @@ class RestartPlan:
|
||||
elif effective_steps != steps:
|
||||
sigmas = sigmas[-(steps + 1) :]
|
||||
|
||||
self.plain_sigmas = sigmas
|
||||
|
||||
restart_segments = prepare_restart_segments(restart_info, ms, sigmas)
|
||||
self.plan, self.total_steps = self.build_plan_items(
|
||||
model.model,
|
||||
@@ -444,9 +401,6 @@ class RestartPlan:
|
||||
def __repr__(self) -> str:
|
||||
return f"<RestartPlan: steps={self.total_steps}, plan={self.plan}>"
|
||||
|
||||
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.
|
||||
@@ -459,6 +413,7 @@ class RestartPlan:
|
||||
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
|
||||
@@ -473,14 +428,14 @@ class RestartPlan:
|
||||
s_max, k, n_restart = seg["t_max"], seg["k"], seg["n"]
|
||||
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]
|
||||
effsmin = max(float(model.model_sampling.sigma_min), sigmas[i + 1])
|
||||
if effsmin >= s_max:
|
||||
continue
|
||||
restart_sigmas = calc_sigmas(
|
||||
restart_scheduler,
|
||||
n_restart,
|
||||
effsmin,
|
||||
max(model_sigma_min, sigmas[i + 1]),
|
||||
s_max,
|
||||
model,
|
||||
device=device,
|
||||
@@ -496,199 +451,176 @@ class RestartPlan:
|
||||
plan.append(PlanItem(sigmas[range_start:]))
|
||||
return plan, sum(pi.total_steps for pi in plan)
|
||||
|
||||
@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]
|
||||
for idx in range(1, len(sigmas)):
|
||||
sigma = sigmas[idx]
|
||||
if last_sigma - sigma < -threshold:
|
||||
return sigmas[:idx]
|
||||
if sigma <= s_min:
|
||||
return sigmas[: idx + 1]
|
||||
last_sigma = sigma
|
||||
raise ValueError("Unexpected end of sigmas in a restart segment")
|
||||
|
||||
plain_sigmas = sigmas.clone().detach().cpu()
|
||||
plan = []
|
||||
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:]
|
||||
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:
|
||||
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:]
|
||||
plan.append(PlanItem(normal_sigmas, k, restart_sigmas))
|
||||
obj = cls.__new__(cls)
|
||||
obj.plan = plan
|
||||
obj.total_steps = sum(pi.total_steps for pi in plan)
|
||||
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()
|
||||
skip = False
|
||||
plan = self.plan
|
||||
planlen = len(plan)
|
||||
for idx in range(planlen):
|
||||
pi = plan[idx]
|
||||
nextpi = None if idx == planlen - 1 else plan[idx + 1]
|
||||
if not skip:
|
||||
yield pi.sigmas
|
||||
skip = False
|
||||
if pi.k == 0:
|
||||
continue
|
||||
for _ in range(pi.k - 1):
|
||||
yield pi.restart_sigmas
|
||||
if nextpi is None:
|
||||
yield pi.restart_sigmas
|
||||
continue
|
||||
skip = True
|
||||
yield pi.restart_sigmas[:-1]
|
||||
yield nextpi.sigmas
|
||||
|
||||
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.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:
|
||||
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" {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)})",
|
||||
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 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)
|
||||
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 = SCHEDULER_MAPPING.keys()
|
||||
if restart_schedules is None:
|
||||
restart_schedules = 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):
|
||||
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]:
|
||||
print(
|
||||
f"CALC NOISE: min={prev_seg[-1].item():.04}, max={seg[0].item():.04}",
|
||||
)
|
||||
noise_scale = ((seg[0] ** 2 - prev_seg[-1] ** 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:
|
||||
#
|
||||
# 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.
|
||||
# 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.
|
||||
#
|
||||
# make_noise_sampler:
|
||||
# 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 sample(
|
||||
self,
|
||||
ksampler,
|
||||
def sampler_function(
|
||||
cls,
|
||||
model,
|
||||
x,
|
||||
_sigmas,
|
||||
sigmas,
|
||||
*args: list,
|
||||
restart_wrapped_sampler=None,
|
||||
restart_chunked=True,
|
||||
restart_make_noise_sampler=None,
|
||||
restart_seed=None,
|
||||
extra_args=None,
|
||||
restart_custom_noise=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
**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
|
||||
|
||||
print("SAMPLING", sigmas)
|
||||
total_steps = len(sigmas - 1)
|
||||
step = 0
|
||||
if restart_seed is None:
|
||||
seed = (extra_args or {}).get("seed", 42)
|
||||
else:
|
||||
seed = restart_seed
|
||||
plan = self.plan
|
||||
|
||||
if VERBOSE:
|
||||
self.explain(restart_chunked)
|
||||
|
||||
def noise_sampler(*_args: list):
|
||||
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 restart_make_noise_sampler:
|
||||
return noise_sampler
|
||||
result = restart_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 = (
|
||||
@@ -706,117 +638,38 @@ class RestartPlan:
|
||||
if callback is not None:
|
||||
callback(cb_state)
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
def do_sample(x, sigs, _kidx=-1):
|
||||
if isinstance(sigs, (list, tuple)):
|
||||
sigs = torch.tensor(sigs, device=x.device)
|
||||
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)
|
||||
# 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.
|
||||
skip = False
|
||||
for idx in range(len(plan)):
|
||||
pi = plan[idx]
|
||||
nextpi = plan[idx + 1] if idx < len(plan) - 1 else None
|
||||
x = pi.execute(
|
||||
x,
|
||||
do_sample,
|
||||
get_noise_sampler,
|
||||
skip_normal=skip,
|
||||
next_pi=nextpi,
|
||||
)
|
||||
skip = pi.k > 0 and len(pi.restart_sigmas) > 2 and nextpi is not None
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def self_test(
|
||||
model,
|
||||
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 noise_scale, chunk_sigmas in cls.split_sigmas(sigmas):
|
||||
print(f"CHUNK: noise={noise_scale:.04}, sigmas={chunk_sigmas}")
|
||||
if noise_scale != 0:
|
||||
x += (
|
||||
restart_noise(
|
||||
x,
|
||||
chunk_sigmas[-1],
|
||||
chunk_sigmas[0],
|
||||
seed + noise_count,
|
||||
)(chunk_sigmas[0], chunk_sigmas[-1])
|
||||
* noise_scale
|
||||
)
|
||||
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}")
|
||||
raise
|
||||
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")
|
||||
noise_count += 1
|
||||
if restart_chunked:
|
||||
x = sampler(
|
||||
model,
|
||||
x,
|
||||
chunk_sigmas,
|
||||
*args,
|
||||
callback=cb_wrapper,
|
||||
disable=True,
|
||||
**kwargs,
|
||||
)
|
||||
continue
|
||||
for i in range(len(chunk_sigmas) - 1):
|
||||
x = sampler(
|
||||
model,
|
||||
x,
|
||||
chunk_sigmas[i : i + 2],
|
||||
*args,
|
||||
callback=cb_wrapper,
|
||||
disable=True,
|
||||
**kwargs,
|
||||
)
|
||||
return x
|
||||
|
||||
Reference in New Issue
Block a user