sampler fixes

This commit is contained in:
Kijai
2024-03-13 18:24:49 +02:00
parent 24c06c19e6
commit 07d30c74ab
2 changed files with 73 additions and 5 deletions
+5 -4
View File
@@ -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'
+68 -1
View File
@@ -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