Skip to content

Commit a77ac2e

Browse files
Merge pull request #7730 from CCRcmcpe/fix-dpm-sde-batch
Fix DPM++ SDE not deterministic across different batch sizes (#5210)
2 parents a742fac + f55a7e0 commit a77ac2e

2 files changed

Lines changed: 31 additions & 8 deletions

File tree

‎modules/sd_samplers_kdiffusion.py‎

Lines changed: 30 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -269,6 +269,16 @@ def get_sigmas(self, p, steps):
269269

270270
return sigmas
271271

272+
def create_noise_sampler(self, x, sigmas, p):
273+
"""For DPM++ SDE: manually create noise sampler to enable deterministic results across different batch sizes"""
274+
if shared.opts.no_dpmpp_sde_batch_determinism:
275+
return None
276+
277+
from k_diffusion.sampling import BrownianTreeNoiseSampler
278+
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
279+
current_iter_seeds = p.all_seeds[p.iteration * p.batch_size:(p.iteration + 1) * p.batch_size]
280+
return BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=current_iter_seeds)
281+
272282
def sample_img2img(self, p, x, noise, conditioning, unconditional_conditioning, steps=None, image_conditioning=None):
273283
steps, t_enc = sd_samplers_common.setup_img2img_steps(p, steps)
274284

@@ -278,18 +288,24 @@ def sample_img2img(self, p, x, noise, conditioning, unconditional_conditioning,
278288
xi = x + noise * sigma_sched[0]
279289

280290
extra_params_kwargs = self.initialize(p)
281-
if 'sigma_min' in inspect.signature(self.func).parameters:
291+
parameters = inspect.signature(self.func).parameters
292+
293+
if 'sigma_min' in parameters:
282294
## last sigma is zero which isn't allowed by DPM Fast & Adaptive so taking value before last
283295
extra_params_kwargs['sigma_min'] = sigma_sched[-2]
284-
if 'sigma_max' in inspect.signature(self.func).parameters:
296+
if 'sigma_max' in parameters:
285297
extra_params_kwargs['sigma_max'] = sigma_sched[0]
286-
if 'n' in inspect.signature(self.func).parameters:
298+
if 'n' in parameters:
287299
extra_params_kwargs['n'] = len(sigma_sched) - 1
288-
if 'sigma_sched' in inspect.signature(self.func).parameters:
300+
if 'sigma_sched' in parameters:
289301
extra_params_kwargs['sigma_sched'] = sigma_sched
290-
if 'sigmas' in inspect.signature(self.func).parameters:
302+
if 'sigmas' in parameters:
291303
extra_params_kwargs['sigmas'] = sigma_sched
292304

305+
if self.funcname == 'sample_dpmpp_sde':
306+
noise_sampler = self.create_noise_sampler(x, sigmas, p)
307+
extra_params_kwargs['noise_sampler'] = noise_sampler
308+
293309
self.model_wrap_cfg.init_latent = x
294310
self.last_latent = x
295311
extra_args={
@@ -303,22 +319,28 @@ def sample_img2img(self, p, x, noise, conditioning, unconditional_conditioning,
303319

304320
return samples
305321

306-
def sample(self, p, x, conditioning, unconditional_conditioning, steps=None, image_conditioning = None):
322+
def sample(self, p, x, conditioning, unconditional_conditioning, steps=None, image_conditioning=None):
307323
steps = steps or p.steps
308324

309325
sigmas = self.get_sigmas(p, steps)
310326

311327
x = x * sigmas[0]
312328

313329
extra_params_kwargs = self.initialize(p)
314-
if 'sigma_min' in inspect.signature(self.func).parameters:
330+
parameters = inspect.signature(self.func).parameters
331+
332+
if 'sigma_min' in parameters:
315333
extra_params_kwargs['sigma_min'] = self.model_wrap.sigmas[0].item()
316334
extra_params_kwargs['sigma_max'] = self.model_wrap.sigmas[-1].item()
317-
if 'n' in inspect.signature(self.func).parameters:
335+
if 'n' in parameters:
318336
extra_params_kwargs['n'] = steps
319337
else:
320338
extra_params_kwargs['sigmas'] = sigmas
321339

340+
if self.funcname == 'sample_dpmpp_sde':
341+
noise_sampler = self.create_noise_sampler(x, sigmas, p)
342+
extra_params_kwargs['noise_sampler'] = noise_sampler
343+
322344
self.last_latent = x
323345
samples = self.launch_sampling(steps, lambda: self.func(self.model_wrap_cfg, x, extra_args={
324346
'cond': conditioning,

‎modules/shared.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -414,6 +414,7 @@ def list_samplers():
414414
options_templates.update(options_section(('compatibility', "Compatibility"), {
415415
"use_old_emphasis_implementation": OptionInfo(False, "Use old emphasis implementation. Can be useful to reproduce old seeds."),
416416
"use_old_karras_scheduler_sigmas": OptionInfo(False, "Use old karras scheduler sigmas (0.1 to 10)."),
417+
"no_dpmpp_sde_batch_determinism": OptionInfo(False, "Do not make DPM++ SDE deterministic across different batch sizes."),
417418
"use_old_hires_fix_width_height": OptionInfo(False, "For hires fix, use width/height sliders to set final resolution rather than first pass (disables Upscale by, Resize width/height to)."),
418419
}))
419420

0 commit comments

Comments
 (0)