diff --git a/modules/pulid/pulid_sampling.py b/modules/pulid/pulid_sampling.py index 9996f035a..e319c0d27 100644 --- a/modules/pulid/pulid_sampling.py +++ b/modules/pulid/pulid_sampling.py @@ -67,6 +67,15 @@ def default_noise_sampler(x): return lambda sigma, sigma_next: torch.randn_like(x) +def inpaint_mask(x, i, steps, mask_args): + noised_original = mask_args["latent"].clone().to(x) + latent_mask = mask_args["latent_mask"].to(x) + if i < steps: + noised_original += mask_args["noise"].to(x) * mask_args["sigmas"][i+1].to(x) + x = (latent_mask * x) + ((1 - latent_mask) * noised_original.to(x)) + return x + + class BatchedBrownianTree: """A wrapper around torchsde.BrownianTree that enables batches of entropy.""" @@ -120,7 +129,7 @@ class BrownianTreeNoiseSampler: @torch.no_grad() -def sample_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.): +def sample_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1., mask_args=None): """Implements Algorithm 2 (Euler steps) from Karras et al. (2022).""" extra_args = {} if extra_args is None else extra_args s_in = x.new_ones([x.shape[0]]) @@ -137,11 +146,13 @@ def sample_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, dt = sigmas[i + 1] - sigma_hat # Euler method x = x + (d * dt).to(x.dtype) + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @torch.no_grad() -def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None): +def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): """Ancestral sampling with Euler method steps.""" extra_args = {} if extra_args is None else extra_args noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler @@ -157,6 +168,8 @@ def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, dis x = x + (d * dt).to(x.dtype) if sigmas[i + 1] > 0: x = x + (noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up).to(x.dtype) + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @@ -375,7 +388,7 @@ class DPMSolver(nn.Module): @torch.no_grad() -def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None): +def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): """Ancestral sampling with DPM-Solver++(2S) second-order steps.""" extra_args = {} if extra_args is None else extra_args noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler @@ -405,11 +418,13 @@ def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, # Noise addition if sigmas[i + 1] > 0: x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @torch.no_grad() -def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2): +def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2, mask_args=None): """DPM-Solver++ (stochastic).""" sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) if noise_sampler is None else noise_sampler @@ -447,11 +462,13 @@ def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=N denoised_d = (1 - fac) * denoised + fac * denoised_2 x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @torch.no_grad() -def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None): +def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None, mask_args=None): """DPM-Solver++(2M).""" extra_args = {} if extra_args is None else extra_args s_in = x.new_ones([x.shape[0]]) @@ -473,11 +490,13 @@ def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=No denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_d old_denoised = denoised + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @torch.no_grad() -def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint'): +def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint', mask_args=None): """DPM-Solver++(2M) SDE.""" if solver_type not in {'heun', 'midpoint'}: @@ -518,11 +537,13 @@ def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disabl old_denoised = denoised h_last = h + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x @torch.no_grad() -def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None): +def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): """DPM-Solver++(3M) SDE.""" sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() @@ -568,4 +589,6 @@ def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disabl denoised_1, denoised_2 = denoised, denoised_1 h_1, h_2 = h, h_1 + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) return x diff --git a/modules/pulid/pulid_sdxl.py b/modules/pulid/pulid_sdxl.py index 2392219c0..fade7509d 100644 --- a/modules/pulid/pulid_sdxl.py +++ b/modules/pulid/pulid_sdxl.py @@ -260,22 +260,40 @@ class StableDiffusionXLPuLIDPipeline: self.callback_on_step_end(self.pipe, step=self.step, timestep=t, kwargs={ 'latents': latent }) return latent - def init_latent(self, seed, size, image, strength): # pylint: disable=unused-argument + def init_latent(self, seed, size, image, mask_image, strength, width, height): # pylint: disable=unused-argument # standard txt2img will full noise noise = torch.randn((size[0], 4, size[1] // 8, size[2] // 8), device="cpu", generator=torch.manual_seed(seed)) noise = noise.to(dtype=self.pipe.unet.dtype, device=self.device) - if image is not None and strength > 0: + if strength > 0 and image is not None: image = self.pipe.image_processor.preprocess(image) - latents = self.pipe.prepare_latents( - image, - None, # timestep (not needed) - 1, # batch_size - 1, # num_images_per_prompt - noise.dtype, - noise.device, - None, # generator - False, # add_noise - ) + if mask_image is not None: # Inpaint + latents = self.pipe.prepare_latents(1, # batch_size, + self.pipe.vae.config.latent_channels, # num_channels_latents + height, + width, + noise.dtype, + noise.device, + None, # generator + latents=None, + image=image, + timestep=1000, + is_strength_max=False, + add_noise=False, + return_noise=False, + return_image_latents=False, + ) + latents = latents[0] + else: # img2img + latents = self.pipe.prepare_latents(image, + None, # timestep (not needed) + 1, # batch_size + 1, # num_images_per_prompt + noise.dtype, + noise.device, + None, # generator + False, # add_noise + ) + else: latents = torch.zeros_like(noise) @@ -309,8 +327,8 @@ class StableDiffusionXLPuLIDPipeline: # latents - latents, noise = self.init_latent(seed, size, image, strength) - latents = latents + noise * sigmas[0].to(noise) + latent, noise = self.init_latent(seed, size, image, mask_image, strength, width, height) + noisy_latent = latent + noise * sigmas[0].to(noise) ( prompt_embeds, @@ -339,17 +357,35 @@ class StableDiffusionXLPuLIDPipeline: cross_attention_kwargs={'id_embedding': uncond_id_embedding, 'id_scale': id_scale}, ), ) + if mask_image is not None: + latent_mask = torch.Tensor(np.asarray(mask_image.convert("L").resize((noisy_latent.shape[-1], noisy_latent.shape[-2])))).reshape((noisy_latent.shape[-2], noisy_latent.shape[-1])) + latent_mask /= latent_mask.max() + mask_args = dict( + latent=latent, + latent_mask=latent_mask, + noise=noise, + sigmas=sigmas, + ) + else: + mask_args = None - latents = self.sampler(self.sample, latents, sigmas, extra_args=sampler_kwargs, disable=False) + latents = self.sampler(self.sample, noisy_latent, sigmas, extra_args=sampler_kwargs, disable=False, mask_args=mask_args) latents = latents.to(dtype=self.pipe.vae.dtype, device=self.device) / self.pipe.vae.config.scaling_factor images = self.pipe.vae.decode(latents).sample images = self.pipe.image_processor.postprocess(images, output_type='pil') - if mask_image is not None: - # TODO: pulid inpaint - # easiest inpaint is to use normal img2img and then combine output with input using mask - # note that mask can be binary or grayscale (soft mask) - raise NotImplementedError(f'PuLID: task=inpaint class={self.__class__.__name__} pipe={self.pipe.__class__.__name__} mask_image={mask_image}') + # Pixel space final mask + # if mask_image is not None: + # # TODO: Fix XYZ + # from PIL import Image + # mask_image = np.asarray(mask_image.convert("L")) + # mask_image = mask_image / mask_image.max() + # mask_image = mask_image.reshape(1,mask_image.shape[0],mask_image.shape[1],1) + # image = np.asarray(image).astype(mask_image.dtype) + # images = np.asarray(images).astype(mask_image.dtype) + # images = ((1 - mask_image) * image) + (mask_image * images) + # images = images[0].round().astype(np.uint8) + # images = [Image.fromarray(images)] return images