From 039d67acf19c90ad105a8ef533650fbfaadcd74d Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Wed, 24 Jul 2024 13:19:03 +0800 Subject: [PATCH] support float16 && add reference && fix bug in training --- README.md | 1 + README_zh-CN.md | 1 + app.py | 9 +- comfyui/comfyui_nodes.py | 2 + easyanimate/data/dataset_image_video.py | 4 +- easyanimate/models/autoencoder_magvit.py | 17 +++- easyanimate/ui/ui.py | 40 +++++--- .../vae/ldm/modules/vaemodules/common.py | 3 + predict_i2v.py | 14 ++- predict_t2v.py | 14 ++- requirements.txt | 10 +- scripts/train.py | 96 +++++++++++-------- 12 files changed, 143 insertions(+), 68 deletions(-) diff --git a/README.md b/README.md index 1098136..3f32275 100644 --- a/README.md +++ b/README.md @@ -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). \ No newline at end of file diff --git a/README_zh-CN.md b/README_zh-CN.md index 8f66c18..ee4f074 100644 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -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). diff --git a/app.py b/app.py index a0992d4..331a7ab 100644 --- a/app.py +++ b/app.py @@ -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( diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index b28f01f..189d4ab 100644 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -1,3 +1,5 @@ +"""Modified from https://github.com/kijai/ComfyUI-EasyAnimateWrapper/blob/main/nodes.py +""" import gc import os diff --git a/easyanimate/data/dataset_image_video.py b/easyanimate/data/dataset_image_video.py index aafd7e7..44ab80b 100644 --- a/easyanimate/data/dataset_image_video.py +++ b/easyanimate/data/dataset_image_video.py @@ -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 diff --git a/easyanimate/models/autoencoder_magvit.py b/easyanimate/models/autoencoder_magvit.py index 0470d67..ba36536 100644 --- a/easyanimate/models/autoencoder_magvit.py +++ b/easyanimate/models/autoencoder_magvit.py @@ -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) diff --git a/easyanimate/ui/ui.py b/easyanimate/ui/ui.py index 5ba475b..5987f14 100644 --- a/easyanimate/ui/ui.py +++ b/easyanimate/ui/ui.py @@ -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, diff --git a/easyanimate/vae/ldm/modules/vaemodules/common.py b/easyanimate/vae/ldm/modules/vaemodules/common.py index a49999d..401dad1 100755 --- a/easyanimate/vae/ldm/modules/vaemodules/common.py +++ b/easyanimate/vae/ldm/modules/vaemodules/common.py @@ -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): diff --git a/predict_i2v.py b/predict_i2v.py index 4ebfb3b..edbbbd4 100644 --- a/predict_i2v.py +++ b/predict_i2v.py @@ -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" diff --git a/predict_t2v.py b/predict_t2v.py index 22ea102..9f2685b 100644 --- a/predict_t2v.py +++ b/predict_t2v.py @@ -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, diff --git a/requirements.txt b/requirements.txt index e6cc69d..1d66be0 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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 \ No newline at end of file diff --git a/scripts/train.py b/scripts/train.py index 322c037..6f67701 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -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 = []