Merge pull request #13 from blepping/restart_custom_noise
Allow chunked restart sampling
This commit is contained in:
@@ -18,6 +18,9 @@ git clone https://github.com/ssitu/ComfyUI_restart_sampling
|
||||
|
||||
The Restart sampler nodes can be found in the node menu under `sampling`.
|
||||
|
||||
If you set the environment variable `COMFYUI_VERBOSE_RESTART_SAMPLING` to `1`, restart sampling will dump
|
||||
information about the steps it's going to run to the console.
|
||||
|
||||
### Nodes
|
||||
|
||||
|Node|Image|Description|
|
||||
@@ -39,6 +42,21 @@ Both $t_{\textrm{min}}$ and $t_{\textrm{max}}$ within a segment definition may b
|
||||
|
||||
You may freely mix the different formats. For example, `[2, 2, -500, "10%"], [3, 2, 5.3, -3]` would be a valid sequence. Note: Random numbers used for example only, not recommended.
|
||||
|
||||
**Special segment values**:
|
||||
|
||||
* Enter `default` to use the default segment list.
|
||||
* Enter `a1111` to emulate A1111 WebUI's segment calculation behavior.
|
||||
For full emulation, enabled chunked mode, set both schedulers to `karras` and the sampler to `heun`.
|
||||
|
||||
### Chunked Mode
|
||||
|
||||
When chunked mode is enabled, the sampler is called with as many steps as possible up to the next segment. When disabled, the sampler
|
||||
is only called with a single step at a time. Some samplers such as SDE samplers, momentum samplers, second order samplers
|
||||
like dpmpp_2m use state from previous steps - when called step-by-step, this state is lost. Using chunked mode may make those
|
||||
samplers more accurate.
|
||||
|
||||
*Note*: Using SDE or momentum samplers with restart is likely not an improvement over normal sampling.
|
||||
|
||||
## Visual Example
|
||||
|
||||
Consider the default segments of `[3,2,0.06,0.30],[3,1,0.30,0.59]`.
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import comfy
|
||||
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING
|
||||
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING, DEFAULT_SEGMENTS
|
||||
|
||||
|
||||
def get_supported_samplers():
|
||||
@@ -24,8 +24,6 @@ def get_supported_samplers():
|
||||
def get_supported_restart_schedulers():
|
||||
return list(SCHEDULER_MAPPING.keys())
|
||||
|
||||
DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]"
|
||||
|
||||
|
||||
class KRestartSamplerSimple:
|
||||
@classmethod
|
||||
@@ -42,7 +40,7 @@ class KRestartSamplerSimple:
|
||||
"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}),
|
||||
"segments": ("STRING", {"default": "default", "multiline": False}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -71,6 +69,7 @@ class KRestartSampler:
|
||||
"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}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,8 +77,8 @@ class KRestartSampler:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler):
|
||||
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise)
|
||||
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:
|
||||
@@ -103,6 +102,7 @@ class KRestartSamplerAdv:
|
||||
"return_with_leftover_noise": (["disable", "enable"], ),
|
||||
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(), ),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -110,10 +110,10 @@ class KRestartSamplerAdv:
|
||||
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):
|
||||
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)
|
||||
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:
|
||||
@@ -137,6 +137,7 @@ class KRestartSamplerCustom:
|
||||
"return_with_leftover_noise": (["disable", "enable"], ),
|
||||
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(), ),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -145,11 +146,10 @@ class KRestartSamplerCustom:
|
||||
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):
|
||||
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)
|
||||
|
||||
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,
|
||||
|
||||
+201
-50
@@ -1,4 +1,6 @@
|
||||
import ast
|
||||
from collections import namedtuple
|
||||
import os
|
||||
import warnings
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
@@ -9,6 +11,10 @@ from comfy.samplers import KSAMPLER, sampler_object
|
||||
from comfy.utils import ProgressBar
|
||||
from .restart_schedulers import SCHEDULER_MAPPING
|
||||
|
||||
VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1"
|
||||
|
||||
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:
|
||||
@@ -33,12 +39,33 @@ def resolve_t_value(val, ms):
|
||||
raise ValueError("bad t_min or t_max value")
|
||||
|
||||
|
||||
def prepare_restart_segments(restart_info, ms):
|
||||
try:
|
||||
restart_arrays = ast.literal_eval(f"[{restart_info}]")
|
||||
except SyntaxError as e:
|
||||
print("Ill-formed restart segments")
|
||||
raise e
|
||||
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":
|
||||
# Emulate A1111 WebUI's restart sampler behavior.
|
||||
steps = len(sigmas) - 1
|
||||
if steps < 20:
|
||||
# Less than 20 steps - no restarts.
|
||||
return []
|
||||
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]]
|
||||
if restart_arrays is None:
|
||||
try:
|
||||
restart_arrays = ast.literal_eval(f"[{restart_info}]")
|
||||
except SyntaxError as e:
|
||||
print("Ill-formed restart segments")
|
||||
raise e
|
||||
restart_segments = []
|
||||
for arr in restart_arrays:
|
||||
if len(arr) != 4:
|
||||
@@ -87,23 +114,23 @@ def calc_restart_steps(restart_segments):
|
||||
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):
|
||||
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)
|
||||
|
||||
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
real_model = model
|
||||
while hasattr(real_model, "model"):
|
||||
real_model = real_model.model
|
||||
|
||||
restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling)
|
||||
|
||||
effective_steps = steps if step_range is not None or denoise > 0.9999 else int(steps / denoise)
|
||||
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,
|
||||
)
|
||||
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
|
||||
|
||||
@@ -117,8 +144,9 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
elif effective_steps != steps:
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
total_steps = [0] # Updated in the wrapper.
|
||||
sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps)
|
||||
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)
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
@@ -148,7 +176,7 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega
|
||||
pbar_update_absolute = ProgressBar.update_absolute
|
||||
|
||||
def pbar_update_absolute_wrapper(self, value, total=None, preview=None):
|
||||
pbar_update_absolute(self, value, total_steps[0], preview)
|
||||
pbar_update_absolute(self, value, sampler_wrapper.total_steps, preview)
|
||||
|
||||
ProgressBar.update_absolute = pbar_update_absolute_wrapper
|
||||
|
||||
@@ -174,49 +202,172 @@ 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])):
|
||||
# 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.
|
||||
# restart_sigmas: Sigmas for the restart segment if it exists, otherwise None.
|
||||
# Note: n_restart is not included as it can be calculated from the length of restart_sigmas.
|
||||
__slots__ = ()
|
||||
|
||||
# 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):
|
||||
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
|
||||
x = sample(x, self.restart_sigmas, kidx)
|
||||
return x
|
||||
|
||||
|
||||
class KSamplerRestartWrapper:
|
||||
|
||||
ksampler = None
|
||||
|
||||
def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps):
|
||||
# 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.
|
||||
# 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:
|
||||
# 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 = total_steps
|
||||
self.total_steps = 0
|
||||
self.seed = seed
|
||||
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, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs):
|
||||
ksampler = self.ksampler
|
||||
segments = round_restart_segments(sigmas, self.restart_segments)
|
||||
self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments)
|
||||
step = 0
|
||||
seed = self.seed
|
||||
plan, self.total_steps = self.build_plan(sigmas, x.device)
|
||||
|
||||
def callback_wrapper(x):
|
||||
x["i"] = step
|
||||
if callback is not None:
|
||||
callback(x)
|
||||
with trange(self.total_steps[0], disable=disable) as pbar:
|
||||
for i in range(len(sigmas) - 1):
|
||||
x = ksampler.sampler_function(
|
||||
model, x, torch.tensor([sigmas[i], sigmas[i + 1]],
|
||||
device=x.device), *args, extra_args=extra_args, callback=callback_wrapper, disable=True,
|
||||
**kwargs)
|
||||
pbar.update(1)
|
||||
if VERBOSE:
|
||||
self.explain_plan(plan, self.total_steps)
|
||||
|
||||
def noise_sampler(*_args):
|
||||
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 self.make_noise_sampler:
|
||||
return noise_sampler
|
||||
result = self.make_noise_sampler(x, s_min, s_max, seed)
|
||||
seed += 1
|
||||
return result
|
||||
|
||||
with trange(self.total_steps, disable=disable) as pbar:
|
||||
def callback_wrapper(x):
|
||||
nonlocal step
|
||||
step += 1
|
||||
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=x.device)
|
||||
for _ in range(k):
|
||||
x += torch.randn_like(x) * (s_max ** 2 - s_min ** 2) ** 0.5
|
||||
for j in range(n_restart - 1):
|
||||
x = ksampler.sampler_function(model, x, torch.tensor(
|
||||
[seg_sigmas[j], seg_sigmas[j + 1]], device=x.device), *args, extra_args=extra_args,
|
||||
callback=callback_wrapper, disable=True, **kwargs)
|
||||
pbar.update(1)
|
||||
step += 1
|
||||
pbar.update(1)
|
||||
x["i"] = step
|
||||
if callback is not None:
|
||||
callback(x)
|
||||
|
||||
# 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 self.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.
|
||||
for pi in plan:
|
||||
x = pi.execute(x, do_sample, get_noise_sampler)
|
||||
|
||||
return x
|
||||
|
||||
Reference in New Issue
Block a user