mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
fix sampler eta
This commit is contained in:
+1
-1
Submodule modules/lora updated: 2d87bb648f...33ee0acd35
@@ -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')
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user