From 704196b0701982b295073cec1a5f3a47d568cfec Mon Sep 17 00:00:00 2001 From: ssit Date: Sat, 21 Oct 2023 18:04:58 -0400 Subject: [PATCH] Change match to ifs for older python combatibility Co-Authored-By: shiertier <90256303+shiertier@users.noreply.github.com> --- restart_sampling.py | 58 ++++++++++++++++++++++----------------------- 1 file changed, 28 insertions(+), 30 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index af6f36b..994152b 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -89,11 +89,10 @@ def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, _restart_scheduler = restart_scheduler _restart_segments = prepare_restart_segments(restart_info) - match sampler_name: - case "ddim": - sampler_wrapper = DDIMWrapper() - case _: - sampler_wrapper = KSamplerRestartWrapper(sampler_name) + if sampler_name == "ddim": + sampler_wrapper = DDIMWrapper() + else: + sampler_wrapper = KSamplerRestartWrapper(sampler_name) # Add the additional steps to the progress bar pbar_update_absolute = ProgressBar.update_absolute @@ -127,31 +126,30 @@ class OneStepSampler: self.denoise = denoise # Get the sampler function - match sampler: - case "ddim": - sampler = DDIMSampler(self.model, device=self.device) - sampler.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) - z_enc = sampler.stochastic_encode(latent_image, torch.tensor( - [len(timesteps) - 1] * noise.shape[0]).to(self.device), noise=noise, max_denoise=max_denoise) - samples, _ = sampler.sample_custom(ddim_timesteps=timesteps, - conditioning=positive, - batch_size=noise.shape[0], - shape=noise.shape[1:], - verbose=False, - unconditional_guidance_scale=cfg, - unconditional_conditioning=negative, - eta=0.0, - x_T=z_enc, - x0=latent_image, - img_callback=ddim_callback, - denoise_function=sampling_function, - extra_args=extra_args, - mask=noise_mask, - to_zero=sigmas[-1] == 0, - end_step=sigmas.shape[0] - 1, - disable_pbar=disable_pbar) - case _: - sample = getattr(k_diffusion_sampling, f"sample_{sampler}") + if sampler == "ddim": + sampler = DDIMSampler(self.model, device=self.device) + sampler.make_schedule_timesteps(ddim_timesteps=timesteps, verbose=False) + z_enc = sampler.stochastic_encode(latent_image, torch.tensor( + [len(timesteps) - 1] * noise.shape[0]).to(self.device), noise=noise, max_denoise=max_denoise) + samples, _ = sampler.sample_custom(ddim_timesteps=timesteps, + conditioning=positive, + batch_size=noise.shape[0], + shape=noise.shape[1:], + verbose=False, + unconditional_guidance_scale=cfg, + unconditional_conditioning=negative, + eta=0.0, + x_T=z_enc, + x0=latent_image, + img_callback=ddim_callback, + denoise_function=sampling_function, + extra_args=extra_args, + mask=noise_mask, + to_zero=sigmas[-1] == 0, + end_step=sigmas.shape[0] - 1, + disable_pbar=disable_pbar) + else: + sample = getattr(k_diffusion_sampling, f"sample_{sampler}") class RestartWrapper: