pulid inpaint, XYZ broken

This commit is contained in:
AI-Casanova
2024-11-08 10:04:02 -06:00
committed by Vladimir Mandic
parent 9b5c0c738d
commit 9fef22735f
2 changed files with 86 additions and 27 deletions
+30 -7
View File
@@ -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
+56 -20
View File
@@ -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