support float16 && add reference && fix bug in training
This commit is contained in:
@@ -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).
|
||||
@@ -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).
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
"""Modified from https://github.com/kijai/ComfyUI-EasyAnimateWrapper/blob/main/nodes.py
|
||||
"""
|
||||
import gc
|
||||
import os
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user