Refactor and simplify

This commit is contained in:
blepping
2024-04-19 05:58:05 -06:00
parent 471b5972c9
commit f5c7a877e5
2 changed files with 236 additions and 375 deletions
+47 -39
View File
@@ -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
View File
@@ -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