support float16 && add reference && fix bug in training

This commit is contained in:
bubbliiiing
2024-07-24 13:19:03 +08:00
parent 1a5bab2234
commit 039d67acf1
12 changed files with 143 additions and 68 deletions
+1
View File
@@ -409,6 +409,7 @@ For more details, please refer to [arxiv](https://arxiv.org/abs/2405.18991).
- Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
- Open-Sora: https://github.com/hpcaitech/Open-Sora
- Animatediff: https://github.com/guoyww/AnimateDiff
- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
# License
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
+1
View File
@@ -406,6 +406,7 @@ EasyAnimateV3:
- Open-Sora-Plan: https://github.com/PKU-YuanGroup/Open-Sora-Plan
- Open-Sora: https://github.com/hpcaitech/Open-Sora
- Animatediff: https://github.com/guoyww/AnimateDiff
- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
# 许可证
本项目采用 [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
+7 -2
View File
@@ -1,4 +1,5 @@
import time
import torch
from easyanimate.api.api import infer_forward_api, update_diffusion_transformer_api, update_edition_api
from easyanimate.ui.ui import ui_modelscope, ui_eas, ui
@@ -9,6 +10,10 @@ if __name__ == "__main__":
# Low gpu memory mode, this is used when the GPU memory is under 16GB
low_gpu_memory_mode = False
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Server ip
server_name = "0.0.0.0"
server_port = 7860
@@ -20,11 +25,11 @@ if __name__ == "__main__":
savedir_sample = "samples"
if ui_mode == "modelscope":
demo, controller = ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode)
demo, controller = ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode, weight_dtype)
elif ui_mode == "eas":
demo, controller = ui_eas(edition, config_path, model_name, savedir_sample)
else:
demo, controller = ui(low_gpu_memory_mode)
demo, controller = ui(low_gpu_memory_mode, weight_dtype)
# launch gradio
app, _, _ = demo.queue(status_update_rate=1).launch(
+2
View File
@@ -1,3 +1,5 @@
"""Modified from https://github.com/kijai/ComfyUI-EasyAnimateWrapper/blob/main/nodes.py
"""
import gc
import os
+3 -1
View File
@@ -291,7 +291,9 @@ class ImageVideoDataset(Dataset):
clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
sample["clip_pixel_values"] = clip_pixel_values
ref_pixel_values = torch.tile(sample["pixel_values"][0].unsqueeze(0), [sample["pixel_values"].size()[0], 1, 1, 1])
ref_pixel_values = sample["pixel_values"][0].unsqueeze(0)
if (mask == 1).all():
ref_pixel_values = torch.ones_like(ref_pixel_values) * -1
sample["ref_pixel_values"] = ref_pixel_values
return sample
+14 -3
View File
@@ -101,6 +101,7 @@ class AutoencoderKLMagvit(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
use_tiling=False,
mini_batch_encoder=9,
mini_batch_decoder=3,
upcast_vae=False,
):
super().__init__()
down_block_types = str_eval(down_block_types)
@@ -152,6 +153,7 @@ class AutoencoderKLMagvit(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
self.mini_batch_decoder = mini_batch_decoder
self.use_slicing = False
self.use_tiling = use_tiling
self.upcast_vae = upcast_vae
self.tile_sample_min_size = 384
self.tile_overlap_factor = 0.25
self.tile_latent_min_size = int(self.tile_sample_min_size / (2 ** (len(ch_mult) - 1)))
@@ -253,8 +255,13 @@ class AutoencoderKLMagvit(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
The latent representations of the encoded images. If `return_dict` is True, a
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
"""
if self.use_tiling and (x.shape[-1] > self.tile_sample_min_size and x.shape[-2] > self.tile_sample_min_size):
return self.tiled_encode(x, return_dict=return_dict)
if self.upcast_vae:
x = x.float()
self.encoder = self.encoder.float()
self.quant_conv = self.quant_conv.float()
if self.use_tiling and (x.shape[-1] > self.tile_sample_min_size or x.shape[-2] > self.tile_sample_min_size):
x = self.tiled_encode(x, return_dict=return_dict)
return x
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self.encoder(x_slice) for x_slice in x.split(1)]
@@ -271,7 +278,11 @@ class AutoencoderKLMagvit(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size and z.shape[-2] > self.tile_latent_min_size):
if self.upcast_vae:
z = z.float()
self.decoder = self.decoder.float()
self.post_quant_conv = self.post_quant_conv.float()
if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size):
return self.tiled_decode(z, return_dict=return_dict)
z = self.post_quant_conv(z)
dec = self.decoder(z)
+28 -12
View File
@@ -56,7 +56,7 @@ css = """
"""
class EasyAnimateController:
def __init__(self, low_gpu_memory_mode):
def __init__(self, low_gpu_memory_mode, weight_dtype):
# config dirs
self.basedir = os.getcwd()
self.config_dir = os.path.join(self.basedir, "config")
@@ -88,7 +88,7 @@ class EasyAnimateController:
self.lora_model_path = "none"
self.low_gpu_memory_mode = low_gpu_memory_mode
self.weight_dtype = torch.bfloat16
self.weight_dtype = weight_dtype
def refresh_diffusion_transformer(self):
self.diffusion_transformer_list = sorted(glob(os.path.join(self.diffusion_transformer_dir, "*/")))
@@ -132,10 +132,16 @@ class EasyAnimateController:
diffusion_transformer_dropdown,
subfolder="vae",
).to(self.weight_dtype)
if OmegaConf.to_container(self.inference_config['vae_kwargs'])['enable_magvit'] and self.weight_dtype == torch.float16:
self.vae.upcast_vae = True
transformer_additional_kwargs = OmegaConf.to_container(self.inference_config['transformer_additional_kwargs'])
if self.weight_dtype == torch.float16:
transformer_additional_kwargs["upcast_attention"] = True
self.transformer = Transformer3DModel.from_pretrained_2d(
diffusion_transformer_dropdown,
subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(self.inference_config.transformer_additional_kwargs)
transformer_additional_kwargs=transformer_additional_kwargs
).to(self.weight_dtype)
self.tokenizer = T5Tokenizer.from_pretrained(diffusion_transformer_dropdown, subfolder="tokenizer")
self.text_encoder = T5EncoderModel.from_pretrained(diffusion_transformer_dropdown, subfolder="text_encoder", torch_dtype=self.weight_dtype)
@@ -471,8 +477,8 @@ class EasyAnimateController:
return gr.Image.update(visible=False, value=None), gr.Video.update(value=save_sample_path, visible=True), "Success"
def ui(low_gpu_memory_mode):
controller = EasyAnimateController(low_gpu_memory_mode)
def ui(low_gpu_memory_mode, weight_dtype):
controller = EasyAnimateController(low_gpu_memory_mode, weight_dtype)
with gr.Blocks(css=css) as demo:
gr.Markdown(
@@ -712,10 +718,7 @@ def ui(low_gpu_memory_mode):
class EasyAnimateController_Modelscope:
def __init__(self, edition, config_path, model_name, savedir_sample, low_gpu_memory_mode):
# Weight Dtype
weight_dtype = torch.bfloat16
def __init__(self, edition, config_path, model_name, savedir_sample, low_gpu_memory_mode, weight_dtype):
# Basic dir
self.basedir = os.getcwd()
self.personalized_model_dir = os.path.join(self.basedir, "models", "Personalized_Model")
@@ -728,10 +731,13 @@ class EasyAnimateController_Modelscope:
self.edition = edition
self.inference_config = OmegaConf.load(config_path)
# Get Transformer
transformer_additional_kwargs = OmegaConf.to_container(self.inference_config['transformer_additional_kwargs'])
if weight_dtype == torch.float16:
transformer_additional_kwargs["upcast_attention"] = True
self.transformer = Transformer3DModel.from_pretrained_2d(
model_name,
subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(self.inference_config['transformer_additional_kwargs'])
transformer_additional_kwargs=transformer_additional_kwargs
).to(weight_dtype)
if OmegaConf.to_container(self.inference_config['vae_kwargs'])['enable_magvit']:
Choosen_AutoencoderKL = AutoencoderKLMagvit
@@ -741,6 +747,8 @@ class EasyAnimateController_Modelscope:
model_name,
subfolder="vae"
).to(weight_dtype)
if OmegaConf.to_container(self.inference_config['vae_kwargs'])['enable_magvit'] and weight_dtype == torch.float16:
self.vae.upcast_vae = True
self.tokenizer = T5Tokenizer.from_pretrained(
model_name,
subfolder="tokenizer"
@@ -814,6 +822,8 @@ class EasyAnimateController_Modelscope:
base_resolution,
generation_method,
length_slider,
overlap_video_length,
partial_video_length,
cfg_scale_slider,
start_image,
end_image,
@@ -942,8 +952,8 @@ class EasyAnimateController_Modelscope:
return gr.Image.update(visible=False, value=None), gr.Video.update(value=save_sample_path, visible=True), "Success"
def ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode):
controller = EasyAnimateController_Modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode)
def ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode, weight_dtype):
controller = EasyAnimateController_Modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memory_mode, weight_dtype)
with gr.Blocks(css=css) as demo:
gr.Markdown(
@@ -1029,6 +1039,8 @@ def ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memo
visible=False,
)
length_slider = gr.Slider(label="Animation length (视频帧数)", value=80, minimum=40, maximum=96, step=1)
overlap_video_length = gr.Slider(label="Overlap length (视频续写的重叠帧数)", value=4, minimum=1, maximum=4, step=1, visible=False)
partial_video_length = gr.Slider(label="Partial video generation length (每个部分的视频生成帧数)", value=72, minimum=8, maximum=144, step=8, visible=False)
cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=6.0, minimum=0, maximum=20)
else:
resize_method = gr.Radio(
@@ -1058,6 +1070,8 @@ def ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memo
visible=True,
)
length_slider = gr.Slider(label="Animation length (视频帧数)", value=48, minimum=8, maximum=48, step=8)
overlap_video_length = gr.Slider(label="Overlap length (视频续写的重叠帧数)", value=4, minimum=1, maximum=4, step=1, visible=False)
partial_video_length = gr.Slider(label="Partial video generation length (每个部分的视频生成帧数)", value=72, minimum=8, maximum=144, step=8, visible=False)
with gr.Accordion("Image to Video (图片到视频)", open=True):
with gr.Row():
@@ -1146,6 +1160,8 @@ def ui_modelscope(edition, config_path, model_name, savedir_sample, low_gpu_memo
base_resolution,
generation_method,
length_slider,
overlap_video_length,
partial_video_length,
cfg_scale_slider,
start_image,
end_image,
@@ -67,6 +67,8 @@ class CausalConv3d(nn.Conv3d):
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, C, T, H, W)
dtype = x.dtype
x = x.float()
if self.padding_flag == 0:
x = F.pad(
x,
@@ -78,6 +80,7 @@ class CausalConv3d(nn.Conv3d):
x,
pad=(0, 0, 0, 0, self.temporal_padding_origin, self.temporal_padding_origin),
)
x = x.to(dtype=dtype)
return super().forward(x)
def set_padding_one_frame(self):
+11 -3
View File
@@ -47,6 +47,8 @@ fps = 24
partial_video_length = None
overlap_video_length = 4
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
validation_image_start = "asset/1.png"
@@ -64,15 +66,19 @@ save_path = "samples/easyanimate-videos_i2v"
config = OmegaConf.load(config_path)
# Get Transformer
if config['enable_multi_text_encoder']:
if config.get('enable_multi_text_encoder', False):
Choosen_Transformer3DModel = HunyuanTransformer3DModel
else:
Choosen_Transformer3DModel = Transformer3DModel
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
if weight_dtype == torch.float16:
transformer_additional_kwargs["upcast_attention"] = True
transformer = Choosen_Transformer3DModel.from_pretrained_2d(
model_name,
subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs=transformer_additional_kwargs
).to(weight_dtype)
if transformer_path is not None:
@@ -108,6 +114,8 @@ vae = Choosen_AutoencoderKL.from_pretrained(
model_name,
subfolder="vae",
).to(weight_dtype)
if OmegaConf.to_container(config['vae_kwargs'])['enable_magvit'] and weight_dtype == torch.float16:
vae.upcast_vae = True
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
@@ -137,7 +145,7 @@ Choosen_Scheduler = scheduler_dict = {
"DDIM": DDIMScheduler,
}[sampler_name]
if config['enable_multi_text_encoder']:
if config.get('enable_multi_text_encoder', False):
scheduler = Choosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
+11 -3
View File
@@ -47,6 +47,8 @@ sample_size = [384, 672]
video_length = 144
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
negative_prompt = "The video is not of a high quality, it has a low resolution, and the audio quality is not clear. Strange motion trajectory, a poor composition and deformed video, low resolution, duplicate and ugly, strange body structure, long and strange neck, bad teeth, bad eyes, bad limbs, bad hands, rotating camera, blurry camera, shaking camera. Deformation, low-resolution, blurry, ugly, distortion. "
@@ -59,15 +61,19 @@ save_path = "samples/easyanimate-videos"
config = OmegaConf.load(config_path)
# Get Transformer
if config['enable_multi_text_encoder']:
if config.get('enable_multi_text_encoder', False):
Choosen_Transformer3DModel = HunyuanTransformer3DModel
else:
Choosen_Transformer3DModel = Transformer3DModel
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
if weight_dtype == torch.float16:
transformer_additional_kwargs["upcast_attention"] = True
transformer = Choosen_Transformer3DModel.from_pretrained_2d(
model_name,
subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs=transformer_additional_kwargs
).to(weight_dtype)
if transformer_path is not None:
@@ -103,6 +109,8 @@ vae = Choosen_AutoencoderKL.from_pretrained(
model_name,
subfolder="vae"
).to(weight_dtype)
if OmegaConf.to_container(config['vae_kwargs'])['enable_magvit'] and weight_dtype == torch.float16:
vae.upcast_vae = True
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
@@ -141,7 +149,7 @@ scheduler = Choosen_Scheduler.from_pretrained(
)
# scheduler = Choosen_Scheduler(**OmegaConf.to_container(config['noise_scheduler_kwargs']))
if config['enable_multi_text_encoder']:
if config.get('enable_multi_text_encoder', False):
if transformer.config.in_channels != vae.config.latent_channels:
pipeline = EasyAnimatePipeline_Multi_Text_Encoder_Inpaint.from_pretrained(
model_name,
+5 -5
View File
@@ -3,8 +3,7 @@ einops
safetensors
timm
tomesd
accelerate
torch>=2.2.0
torch>=2.1.2
torchdiffeq
torchsde
xformers
@@ -19,8 +18,9 @@ albumentations
imageio[ffmpeg]
imageio[pyav]
tensorboard
gradio==3.41.2
diffusers==0.28.2
transformers==4.37.2
beautifulsoup4
ftfy
accelerate>=0.25.0
gradio>=3.41.2
diffusers>=0.28.2
transformers>=4.37.2
+57 -39
View File
@@ -1339,15 +1339,18 @@ def main():
if vae.quant_conv.weight.ndim==5:
# This way is quicker when batch grows up
if vae.slice_compression_vae:
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
ref_pixel_values_bs = ref_pixel_values[i : i + bs]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs)[0]
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values.append(ref_pixel_values_bs)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
if config.get('enable_multi_text_encoder', False):
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
ref_pixel_values_bs = ref_pixel_values[i : i + bs]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs)[0]
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values.append(ref_pixel_values_bs)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
else:
ref_latents = None
mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w")
bs = args.vae_mini_batch
@@ -1369,23 +1372,29 @@ def main():
mask_bs = mask_bs.sample()
new_mask.append(mask_bs)
mask = torch.cat(new_mask, dim = 0)
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
if ref_latents is not None:
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
else:
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
else:
# This way is quicker when batch grows up
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
new_ref_pixel_values_mini_batch = []
for j in range(0, ref_pixel_values.shape[2], sample_n_frames_bucket_interval):
ref_pixel_values_bs = ref_pixel_values[i : i + bs, :, j: j + sample_n_frames_bucket_interval, :, :]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs)[0]
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values_mini_batch.append(ref_pixel_values_bs)
new_ref_pixel_values_mini_batch = torch.cat(new_ref_pixel_values_mini_batch, dim = 2)
new_ref_pixel_values.append(new_ref_pixel_values_mini_batch)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
if config.get('enable_multi_text_encoder', False):
# This way is quicker when batch grows up
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
new_ref_pixel_values_mini_batch = []
for j in range(0, ref_pixel_values.shape[2], sample_n_frames_bucket_interval):
ref_pixel_values_bs = ref_pixel_values[i : i + bs, :, j: j + sample_n_frames_bucket_interval, :, :]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs)[0]
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values_mini_batch.append(ref_pixel_values_bs)
new_ref_pixel_values_mini_batch = torch.cat(new_ref_pixel_values_mini_batch, dim = 2)
new_ref_pixel_values.append(new_ref_pixel_values_mini_batch)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
else:
ref_latents = None
# This way is quicker when batch grows up
mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w")
@@ -1417,19 +1426,25 @@ def main():
new_mask_mini_batch = torch.cat(new_mask_mini_batch, dim = 2)
new_mask.append(new_mask_mini_batch)
mask = torch.cat(new_mask, dim = 0)
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
if ref_latents is not None:
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
else:
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
else:
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> (b f) c h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
ref_pixel_values_bs = ref_pixel_values[i : i + bs]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs.to(dtype=weight_dtype)).latent_dist
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values.append(ref_pixel_values_bs)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
ref_latents = rearrange(ref_latents, "(b f) c h w -> b c f h w", f=video_length)
if config.get('enable_multi_text_encoder', False):
ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> (b f) c h w")
bs = args.vae_mini_batch
new_ref_pixel_values = []
for i in range(0, ref_pixel_values.shape[0], bs):
ref_pixel_values_bs = ref_pixel_values[i : i + bs]
ref_pixel_values_bs = vae.encode(ref_pixel_values_bs.to(dtype=weight_dtype)).latent_dist
ref_pixel_values_bs = ref_pixel_values_bs.sample()
new_ref_pixel_values.append(ref_pixel_values_bs)
ref_latents = torch.cat(new_ref_pixel_values, dim = 0)
ref_latents = rearrange(ref_latents, "(b f) c h w -> b c f h w", f=video_length)
else:
ref_latents = None
mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> (b f) c h w")
bs = args.vae_mini_batch
@@ -1447,8 +1462,11 @@ def main():
mask, size=(mask_latents.size()[-2], mask_latents.size()[-1])
)
mask = rearrange(mask, "(b f) c h w -> b c f h w", f=video_length)
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
if ref_latents is not None:
ref_latents = ref_latents.expand_as(mask_latents)
inpaint_latents = torch.concat([mask, mask_latents, ref_latents], dim=1)
else:
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
with torch.no_grad():
clip_encoder_hidden_states = []