fix sampler eta

This commit is contained in:
Vladimir Mandic
2023-10-09 07:27:00 -04:00
parent 57ce239027
commit 6b4df05b5c
4 changed files with 24 additions and 26 deletions
+3 -3
View File
@@ -1028,7 +1028,7 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
self.ops.append('txt2img')
hypertile_set(self)
self.sampler = modules.sd_samplers.create_sampler(self.sampler_name, self.sd_model)
self.sampler = modules.sd_samplers.create_sampler(self.sampler_name, self.sd_model, self)
x = create_random_tensors([4, self.height // 8, self.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=self.subseed_strength, seed_resize_from_h=self.seed_resize_from_h, seed_resize_from_w=self.seed_resize_from_w, p=self)
samples = self.sampler.sample(self, x, conditioning, unconditional_conditioning, image_conditioning=self.txt2img_image_conditioning(x))
if not self.enable_hr or shared.state.interrupted or shared.state.skipped:
@@ -1076,7 +1076,7 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
if self.denoising_strength > 0:
self.ops.append('hires')
devices.torch_gc() # GC now before running the next img2img to prevent running out of memory
self.sampler = modules.sd_samplers.create_sampler(self.latent_sampler or self.sampler_name, self.sd_model)
self.sampler = modules.sd_samplers.create_sampler(self.latent_sampler or self.sampler_name, self.sd_model, self)
samples = samples[:, :, self.truncate_y//2:samples.shape[2]-(self.truncate_y+1)//2, self.truncate_x//2:samples.shape[3]-(self.truncate_x+1)//2]
noise = create_random_tensors(samples.shape[1:], seeds=seeds, subseeds=subseeds, subseed_strength=subseed_strength, p=self)
modules.sd_models.apply_token_merging(self.sd_model, self.get_token_merging_ratio(for_hr=True))
@@ -1130,7 +1130,7 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
if self.sampler_name == "PLMS":
self.sampler_name = 'UniPC'
self.sampler = modules.sd_samplers.create_sampler(self.sampler_name, self.sd_model)
self.sampler = modules.sd_samplers.create_sampler(self.sampler_name, self.sd_model, self)
if self.image_mask is not None:
self.ops.append('inpaint')
+2 -2
View File
@@ -34,7 +34,7 @@ def find_sampler_config(name):
return config
def create_sampler(name, model):
def create_sampler(name, model, p=None):
if name == 'Default' and hasattr(model, 'scheduler'):
config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_')}
shared.log.debug(f'Sampler default {type(model.scheduler).__name__}: {config}')
@@ -46,7 +46,7 @@ def create_sampler(name, model):
if shared.backend == shared.Backend.ORIGINAL:
sampler = config.constructor(model)
sampler.config = config
sampler.initialize(p=None)
sampler.initialize(p)
sampler.name = name
shared.log.debug(f'Sampler: sampler={sampler.name} config={sampler.config.options}')
return sampler
+18 -20
View File
@@ -128,26 +128,24 @@ class VanillaStableDiffusionSampler:
self.update_step(x)
def initialize(self, p):
if self.is_ddim:
self.eta = p.eta if p.eta is not None else shared.opts.scheduler_eta
else:
self.eta = 0.0
if self.eta != 0.0:
p.extra_generation_params["Sampler Eta"] = self.eta
if self.is_unipc:
keys = [
('Solver order', 'schedulers_solver_order'),
('Sampler low order', 'schedulers_use_loworder'),
('UniPC variant', 'uni_pc_variant'),
('UniPC skip type', 'uni_pc_skip_type'),
]
for name, key in keys:
v = getattr(shared.opts, key)
if v != shared.opts.get_default(key):
p.extra_generation_params[name] = v
if p is not None:
if self.is_ddim:
self.eta = p.eta if p.eta is not None else shared.opts.scheduler_eta
else:
self.eta = 0.0
if self.eta != 0.0:
p.extra_generation_params["Sampler Eta"] = self.eta
if self.is_unipc:
keys = [
('Solver order', 'schedulers_solver_order'),
('Sampler low order', 'schedulers_use_loworder'),
('UniPC variant', 'uni_pc_variant'),
('UniPC skip type', 'uni_pc_skip_type'),
]
for name, key in keys:
v = getattr(shared.opts, key)
if v != shared.opts.get_default(key):
p.extra_generation_params[name] = v
for fieldname in ['p_sample_ddim', 'p_sample_plms']:
if hasattr(self.sampler, fieldname):
setattr(self.sampler, fieldname, self.p_sample_ddim_hook)