update inpaint model
This commit is contained in:
@@ -22,6 +22,7 @@ from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers import DiffusionPipeline, ImagePipelineOutput
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.models import AutoencoderKL
|
||||
@@ -108,16 +109,16 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
text_encoder: T5EncoderModel,
|
||||
vae: AutoencoderKL,
|
||||
transformer: Transformer3DModel,
|
||||
image_processor,
|
||||
image_encoder,
|
||||
scheduler: DPMSolverMultistepScheduler,
|
||||
# image_processor,
|
||||
# image_encoder,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer,
|
||||
image_processor=image_processor, image_encoder=image_encoder,
|
||||
scheduler=scheduler
|
||||
scheduler=scheduler,
|
||||
# image_processor=image_processor, image_encoder=image_encoder,
|
||||
)
|
||||
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
@@ -630,7 +631,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
video.append(latents_bs)
|
||||
video = torch.cat(video, 2)
|
||||
video = video.clamp(-1, 1)
|
||||
video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
|
||||
# video = self.smooth_output(video, mini_batch_encoder, mini_batch_decoder).cpu().clamp(-1, 1)
|
||||
else:
|
||||
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
video = []
|
||||
@@ -657,6 +658,16 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
|
||||
return image_latents
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.StableDiffusionImg2ImgPipeline.get_timesteps
|
||||
def get_timesteps(self, num_inference_steps, strength, device):
|
||||
# get the original timestep using init_timestep
|
||||
init_timestep = min(int(num_inference_steps * strength), num_inference_steps)
|
||||
|
||||
t_start = max(num_inference_steps - init_timestep, 0)
|
||||
timesteps = self.scheduler.timesteps[t_start * self.scheduler.order :]
|
||||
|
||||
return timesteps, num_inference_steps - t_start
|
||||
|
||||
def prepare_mask_latents(
|
||||
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance
|
||||
):
|
||||
@@ -878,11 +889,13 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
# 4. set timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
timesteps, num_inference_steps = self.get_timesteps(
|
||||
num_inference_steps=num_inference_steps, strength=strength, device=device
|
||||
)
|
||||
# at which timestep to set the initial noise (n.b. 50% if strength is 0.5)
|
||||
latent_timestep = timesteps[:1].repeat(batch_size)
|
||||
latent_timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
|
||||
# create a boolean to check if the strength is set to 1. if so then initialise the latents with pure noise
|
||||
is_strength_max = strength == 1.0
|
||||
|
||||
@@ -930,6 +943,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
if num_channels_transformer == 12:
|
||||
mask_condition = torch.tile(mask_condition, [1, 3, 1, 1, 1])
|
||||
if masked_video_latents is None:
|
||||
masked_video = init_video * (mask_condition < 0.5) + torch.ones_like(init_video) * (mask_condition > 0.5) * -1
|
||||
else:
|
||||
@@ -947,10 +961,18 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
else:
|
||||
mask = mask_condition
|
||||
mask = torch.tile(mask_condition, [1, num_channels_transformer, 1, 1, 1])
|
||||
mask = F.interpolate(mask, size=latents.size()[-3:], mode='trilinear', align_corners=True).to(latents.device, latents.dtype)
|
||||
|
||||
masked_video_latents = torch.zeros_like(latents).to(latents.device, latents.dtype)
|
||||
else:
|
||||
mask = torch.zeros_like(latents).to(latents.device, latents.dtype)
|
||||
if num_channels_transformer == 12:
|
||||
mask = torch.zeros_like(latents).to(latents.device, latents.dtype)
|
||||
else:
|
||||
mask = torch.zeros_like(init_video[:, :1])
|
||||
mask = torch.tile(mask, [1, num_channels_transformer, 1, 1, 1])
|
||||
mask = F.interpolate(mask, size=latents.size()[-3:], mode='trilinear', align_corners=True).to(latents.device, latents.dtype)
|
||||
|
||||
masked_video_latents = torch.zeros_like(latents).to(latents.device, latents.dtype)
|
||||
|
||||
# Check that sizes of mask, masked image and latents match
|
||||
@@ -971,7 +993,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
f"The transformer {self.transformer.__class__} should have 9 input channels, not {self.transformer.config.in_channels}."
|
||||
)
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
# 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
# 6.1 Prepare micro-conditions.
|
||||
@@ -988,9 +1010,9 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
|
||||
added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio}
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
|
||||
# 10. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
@@ -1001,7 +1023,9 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
masked_video_latents_input = (
|
||||
torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents
|
||||
)
|
||||
inpaint_latents = torch.cat([mask_input, masked_video_latents_input], dim=1)
|
||||
inpaint_latents = torch.cat([mask_input, masked_video_latents_input], dim=1).to(latent_model_input.dtype)
|
||||
else:
|
||||
inpaint_latents = None
|
||||
|
||||
current_timestep = t
|
||||
if not torch.is_tensor(current_timestep):
|
||||
@@ -1025,7 +1049,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
timestep=current_timestep,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
inpaint_latents=inpaint_latents.to(latent_model_input.dtype),
|
||||
inpaint_latents=inpaint_latents,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
@@ -1042,7 +1066,7 @@ class EasyAnimateInpaintPipeline(DiffusionPipeline):
|
||||
|
||||
if num_channels_transformer == 4:
|
||||
init_latents_proper = image_latents
|
||||
init_mask = F.interpolate(mask, size=init_latents_proper.size(), mode='trilinear', align_corners=True)
|
||||
init_mask = mask
|
||||
if i < len(timesteps) - 1:
|
||||
noise_timestep = timesteps[i + 1]
|
||||
init_latents_proper = self.scheduler.add_noise(
|
||||
|
||||
Reference in New Issue
Block a user