sampler fixes
This commit is contained in:
+5
-4
@@ -327,11 +327,11 @@ class SUPIR_sample:
|
||||
"steps": ("INT", {"default": 45, "min": 3, "max": 4096, "step": 1}),
|
||||
"cfg_scale_start": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 9.0, "step": 0.05}),
|
||||
"cfg_scale_end": ("FLOAT", {"default": 4.0, "min": 0, "max": 20, "step": 0.01}),
|
||||
"s_churn": ("INT", {"default": 5, "min": 0, "max": 40, "step": 1}),
|
||||
"EDM_s_churn": ("INT", {"default": 5, "min": 0, "max": 40, "step": 1}),
|
||||
"s_noise": ("FLOAT", {"default": 1.003, "min": 1.0, "max": 1.1, "step": 0.001}),
|
||||
"control_scale_start": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
|
||||
"control_scale_end": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
|
||||
"restore_cfg": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 6.0, "step": 0.05}),
|
||||
"restore_cfg": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 20.0, "step": 0.05}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
"sampler": (
|
||||
[
|
||||
@@ -339,6 +339,7 @@ class SUPIR_sample:
|
||||
'RestoreEDMSampler',
|
||||
'TiledRestoreDPMPP2MSampler',
|
||||
'TiledRestoreEDMSampler',
|
||||
'EulerAncestralSampler'
|
||||
], {
|
||||
"default": 'RestoreEDMSampler'
|
||||
}),
|
||||
@@ -355,7 +356,7 @@ class SUPIR_sample:
|
||||
DESCRIPTION="Samples using SUPIR's modified diffusion."
|
||||
CATEGORY = "SUPIR"
|
||||
|
||||
def sample(self, SUPIR_model, latents, steps, seed, cfg_scale_end, s_churn, s_noise, positive, negative,
|
||||
def sample(self, SUPIR_model, latents, steps, seed, cfg_scale_end, EDM_s_churn, s_noise, positive, negative,
|
||||
cfg_scale_start, control_scale_start, control_scale_end, restore_cfg, keep_model_loaded,
|
||||
sampler, sampler_tile_size=1024, sampler_tile_stride=512):
|
||||
|
||||
@@ -369,7 +370,7 @@ class SUPIR_sample:
|
||||
'params': {
|
||||
'num_steps': steps,
|
||||
'restore_cfg': restore_cfg,
|
||||
's_churn': s_churn,
|
||||
's_churn': EDM_s_churn,
|
||||
's_noise': s_noise,
|
||||
'discretization_config': {
|
||||
'target': '.sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization'
|
||||
|
||||
@@ -560,6 +560,9 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler):
|
||||
restore_cfg_s_tmin=0.05, eta=1., *args, **kwargs):
|
||||
self.s_noise = s_noise
|
||||
self.eta = eta
|
||||
self.restore_cfg = restore_cfg
|
||||
self.restore_cfg_s_tmin = restore_cfg_s_tmin
|
||||
self.sigma_max = 14.6146
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def denoise(self, x, denoiser, sigma, cond, uc, control_scale=1.0):
|
||||
@@ -591,10 +594,20 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler):
|
||||
cond,
|
||||
uc=None,
|
||||
eps_noise=None,
|
||||
x_center=None,
|
||||
control_scale=1.0,
|
||||
use_linear_control_scale=False,
|
||||
control_scale_start=0.0
|
||||
):
|
||||
if use_linear_control_scale:
|
||||
control_scale = (sigma[0].item() / self.sigma_max) * (control_scale_start - control_scale) + control_scale
|
||||
|
||||
denoised = self.denoise(x, denoiser, sigma, cond, uc, control_scale=control_scale)
|
||||
|
||||
if (next_sigma[0] > self.restore_cfg_s_tmin) and (self.restore_cfg > 0):
|
||||
d_center = (denoised - x_center)
|
||||
denoised = denoised - d_center * ((sigma.view(-1, 1, 1, 1) / self.sigma_max) ** self.restore_cfg)
|
||||
|
||||
h, r, t, t_next = self.get_variables(sigma, next_sigma, previous_sigma)
|
||||
eta_h = self.eta * h
|
||||
mult = [
|
||||
@@ -619,7 +632,8 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler):
|
||||
|
||||
return x, denoised
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, control_scale=1.0, **kwargs):
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, x_center=None, control_scale=1.0,
|
||||
use_linear_control_scale=False, control_scale_start=0.0, **kwargs):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
@@ -647,6 +661,9 @@ class RestoreDPMPP2MSampler(DPMPP2MSampler):
|
||||
uc=uc,
|
||||
eps_noise=eps_noise,
|
||||
control_scale=control_scale,
|
||||
x_center=x_center,
|
||||
use_linear_control_scale=use_linear_control_scale,
|
||||
control_scale_start=control_scale_start,
|
||||
)
|
||||
pbar_comfy.update(1)
|
||||
|
||||
@@ -725,4 +742,54 @@ class TiledRestoreDPMPP2MSampler(RestoreDPMPP2MSampler):
|
||||
x = x_next
|
||||
old_denoised = old_denoised_next
|
||||
pbar_comfy.update(1)
|
||||
return x
|
||||
|
||||
class SubstepSampler(EulerAncestralSampler):
|
||||
def __init__(self, s_churn=0.0, s_tmin=0.0, s_tmax=float("inf"), s_noise=1.0, restore_cfg=4.0,
|
||||
restore_cfg_s_tmin=0.05, eta=1., n_sample_steps=4, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.n_sample_steps = n_sample_steps
|
||||
self.steps_subset = [0, 100, 200, 300, 1000]
|
||||
|
||||
def prepare_sampling_loop(self, x, cond, uc=None, num_steps=None):
|
||||
sigmas = self.discretization(1000, device=self.device)
|
||||
sigmas = sigmas[
|
||||
self.steps_subset[: self.num_steps] + self.steps_subset[-1:]
|
||||
]
|
||||
print(sigmas)
|
||||
# uc = cond
|
||||
x *= torch.sqrt(1.0 + sigmas[0] ** 2.0)
|
||||
num_sigmas = len(sigmas)
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
return x, s_in, sigmas, num_sigmas, cond, uc
|
||||
|
||||
def denoise(self, x, denoiser, sigma, cond, uc, control_scale=1.0):
|
||||
denoised = denoiser(*self.guider.prepare_inputs(x, sigma, cond, uc), control_scale)
|
||||
denoised = self.guider(denoised, sigma)
|
||||
return denoised
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, control_scale=1.0, *args, **kwargs):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
for i in self.get_sigma_gen(num_sigmas):
|
||||
x = self.sampler_step(
|
||||
s_in * sigmas[i],
|
||||
s_in * sigmas[i + 1],
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc,
|
||||
control_scale=control_scale,
|
||||
)
|
||||
|
||||
return x
|
||||
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc, control_scale=1.0):
|
||||
sigma_down, sigma_up = get_ancestral_step(sigma, next_sigma, eta=self.eta)
|
||||
denoised = self.denoise(x, denoiser, sigma, cond, uc, control_scale=control_scale)
|
||||
x = self.ancestral_euler_step(x, denoised, sigma, sigma_down)
|
||||
x = self.ancestral_step(x, sigma, next_sigma, sigma_up)
|
||||
|
||||
return x
|
||||
Reference in New Issue
Block a user