This commit is contained in:
Seunghoon Lee
2023-10-29 22:57:16 +09:00
parent a99c7d3d53
commit 0fb95df44d
+7 -4
View File
@@ -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