From 6b4df05b5cedf4b716220139e2c213ab301c2571 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 9 Oct 2023 07:27:00 -0400 Subject: [PATCH] fix sampler eta --- modules/lora | 2 +- modules/processing.py | 6 +++--- modules/sd_samplers.py | 4 ++-- modules/sd_samplers_compvis.py | 38 ++++++++++++++++------------------ 4 files changed, 24 insertions(+), 26 deletions(-) diff --git a/modules/lora b/modules/lora index 2d87bb648..33ee0acd3 160000 --- a/modules/lora +++ b/modules/lora @@ -1 +1 @@ -Subproject commit 2d87bb648f30adab00ceb38a0da786cd548d5ce7 +Subproject commit 33ee0acd358bb23df73cfa1a45e63b37bd17c28d diff --git a/modules/processing.py b/modules/processing.py index 6756b8040..536dd97bf 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -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') diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 103d7a93c..4a7d4c352 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -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 diff --git a/modules/sd_samplers_compvis.py b/modules/sd_samplers_compvis.py index 548479ccc..c5ec87a1c 100644 --- a/modules/sd_samplers_compvis.py +++ b/modules/sd_samplers_compvis.py @@ -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)