diff --git a/modules/sd_samplers_kdiffusion.py b/modules/sd_samplers_kdiffusion.py index ff8d63bab..6617a0575 100644 --- a/modules/sd_samplers_kdiffusion.py +++ b/modules/sd_samplers_kdiffusion.py @@ -173,10 +173,12 @@ class CFGDenoiser(torch.nn.Module): else: denoised = self.combine_denoised(x_out, conds_list, uncond, cond_scale) if self.mask is not None: - init_latent = self.init_latent if devices.backend == "directml": - init_latent = init_latent.float() - denoised = init_latent * self.mask + self.nmask * denoised + self.init_latent = self.init_latent.float() + denoised = self.init_latent * self.mask + self.nmask * denoised + self.init_latent = self.init_latent.half() + else: + denoised = self.init_latent * self.mask + self.nmask * denoised after_cfg_callback_params = AfterCFGCallbackParams(denoised, shared.state.sampling_step, shared.state.sampling_steps) cfg_after_cfg_callback(after_cfg_callback_params) denoised = after_cfg_callback_params.x @@ -336,7 +338,8 @@ class KDiffusionSampler: 's_min_uncond': self.s_min_uncond } samples = self.launch_sampling(t_enc + 1, lambda: self.func(self.model_wrap_cfg, xi, extra_args=extra_args, disable=False, callback=self.callback_state, **extra_params_kwargs)) - return samples.type(devices.dtype) + samples = samples.type(devices.dtype) + return samples def sample(self, p, x, conditioning, unconditional_conditioning, steps=None, image_conditioning=None): steps = steps or p.steps