Make restart a sampler, add node to generate sigmas: phase 1

This commit is contained in:
blepping
2024-04-15 10:35:00 -06:00
parent a73142188a
commit 3752a2daf9
2 changed files with 229 additions and 78 deletions
+143 -1
View File
@@ -1,6 +1,12 @@
import comfy
import torch
from .restart_sampling import DEFAULT_SEGMENTS, SCHEDULER_MAPPING, restart_sampling
from . import restart_sampling as restart
from .restart_sampling import (
DEFAULT_SEGMENTS,
SCHEDULER_MAPPING,
restart_sampling,
)
def get_supported_samplers():
@@ -296,11 +302,147 @@ class KRestartSamplerCustom:
)
class RestartScheduler:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"scheduler": (tuple(SCHEDULER_MAPPING.keys()),),
"segments": (
"STRING",
{"default": DEFAULT_SEGMENTS, "multiline": False},
),
"restart_scheduler": (get_supported_restart_schedulers(),),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
},
"optional": {
"sigmas_opt": ("SIGMAS",),
},
}
RETURN_TYPES = ("SIGMAS",)
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 * -1
def go(
self,
model,
steps,
scheduler,
segments,
restart_scheduler,
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)
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,
restart_scheduler,
sigmas,
"cpu",
)
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,)
class RestartSampler:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sampler": ("SAMPLER",),
},
}
RETURN_TYPES = ("SAMPLER",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/samplers"
def go(self, sampler):
wrapped = comfy.samplers.KSAMPLER(
lambda *args, **kwargs: self.sampler_function(sampler, *args, **kwargs),
extra_options=sampler.extra_options,
inpaint_options=sampler.inpaint_options,
)
return (wrapped,)
@staticmethod
@torch.no_grad()
def sampler_function(wrapped, model, x, sigmas, *args, **kwargs):
last_sigma = None
chunks = []
while len(sigmas) > 0:
last_sigma = None
for idx in range(len(sigmas) - 1):
curr_sigma = sigmas[idx + 1]
if last_sigma is None or (
curr_sigma.sign() == last_sigma.sign()
and curr_sigma.abs() < last_sigma.abs()
):
last_sigma = curr_sigma
continue
break
if idx == len(sigmas) - 2:
chunks.append(sigmas)
break
chunks.append(
sigmas[: idx + 1] * -1 if sigmas[0] < 0 else sigmas[: idx + 1],
)
sigmas = sigmas[idx + 1 :]
print("CHUNKS", chunks)
chunks = [chunk for chunk in chunks if len(chunk) > 1]
for idx, chunk_sigmas in enumerate(chunks):
print(">>>", idx, chunk_sigmas)
if idx > 0 and chunk_sigmas[0] > chunks[idx - 1][-1]:
print("NOISE", chunk_sigmas[0], chunk_sigmas[-1])
x += (
torch.randn_like(x)
* (chunk_sigmas[0] ** 2 - chunk_sigmas[-1] ** 2) ** 0.5
)
x = wrapped.sampler_function(model, x, chunk_sigmas, *args, **kwargs)
return x
NODE_CLASS_MAPPINGS = {
"KRestartSamplerSimple": KRestartSamplerSimple,
"KRestartSampler": KRestartSampler,
"KRestartSamplerAdv": KRestartSamplerAdv,
"KRestartSamplerCustom": KRestartSamplerCustom,
"RestartScheduler": RestartScheduler,
"RestartSampler": RestartSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+86 -77
View File
@@ -314,6 +314,84 @@ 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,
)
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
# 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(plan)
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.
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):
return lambda *_args: 0.0
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.",
)
class KSamplerRestartWrapper:
# Some extra explanation for a couple of these arguments:
#
@@ -346,81 +424,6 @@ class KSamplerRestartWrapper:
self.make_noise_sampler = make_noise_sampler
self.chunked = chunked
# 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(self, sigmas, device):
segments = round_restart_segments(sigmas, self.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(
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.
plan.append(PlanItem(sigmas[range_start:]))
return plan, total_steps
# Dumps information about the plan to the console. It uses the normal plan execute
# logic.
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.
def do_sample(x, sigs, kidx=-1):
nonlocal step, last_kidx
rlabel = f"R{kidx+1:>3}" if kidx > last_kidx else " "
last_kidx = kidx
if not self.chunked:
for i in range(len(sigs) - 1):
step += 1
print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigs[i:i+2])}")
return x
chunk_size = len(sigs) - 2
step += 1
print(
f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigs)}",
)
step += chunk_size
return x
# Stub function to satisfy PlanItem.execute
def get_noise_sampler(*_args):
return lambda *_args: 0.0
for pi in plan:
pi.execute(0.0, do_sample, get_noise_sampler)
print(
"** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.",
)
@torch.no_grad()
def ksampler_restart_wrapper(
self,
@@ -436,10 +439,16 @@ class KSamplerRestartWrapper:
ksampler = self.ksampler
step = 0
seed = self.seed
plan, self.total_steps = self.build_plan(sigmas, x.device)
plan, self.total_steps = build_plan(
self.real_model,
self.restart_segments,
self.restart_scheduler,
sigmas,
x.device,
)
if VERBOSE:
self.explain_plan(plan, self.total_steps)
explain_plan(plan, self.total_steps, chunked=self.chunked)
def noise_sampler(*_args):
return torch.randn_like(x)