mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
pulid inpaint, XYZ broken
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user