From e6336652aa0904bf0b2c9d97eb77028bbbe222ee Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Mon, 27 May 2024 13:54:29 +0800 Subject: [PATCH] update datasets loader --- easyanimate/data/dataset_video.py | 65 ++++++++++++++++++++++++------- easyanimate/ui/ui.py | 30 ++++++++------ 2 files changed, 68 insertions(+), 27 deletions(-) diff --git a/easyanimate/data/dataset_video.py b/easyanimate/data/dataset_video.py index 8095ffd..7831b4c 100644 --- a/easyanimate/data/dataset_video.py +++ b/easyanimate/data/dataset_video.py @@ -1,17 +1,26 @@ import csv +import gc import io import json import math import os import random +from contextlib import contextmanager +from threading import Thread +import albumentations +import cv2 import numpy as np import torch import torchvision.transforms as transforms from decord import VideoReader from einops import rearrange +from func_timeout import FunctionTimedOut, func_timeout +from PIL import Image +from torch.utils.data import BatchSampler, Sampler from torch.utils.data.dataset import Dataset +VIDEO_READER_TIMEOUT = 20 def get_random_mask(shape): f, c, h, w = shape @@ -53,6 +62,21 @@ def get_random_mask(shape): return mask +@contextmanager +def VideoReader_contextmanager(*args, **kwargs): + vr = VideoReader(*args, **kwargs) + try: + yield vr + finally: + del vr + gc.collect() + + +def get_video_reader_batch(video_reader, batch_index): + frames = video_reader.get_batch(batch_index).asnumpy() + return frames + + class WebVid10M(Dataset): def __init__( self, @@ -96,11 +120,11 @@ class WebVid10M(Dataset): batch_index = [random.randint(0, video_length - 1)] if not self.enable_bucket: - pixel_values = torch.from_numpy(video_reader.get_batch(batch_index).asnumpy()).permute(0, 3, 1, 2).contiguous() + pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous() pixel_values = pixel_values / 255. del video_reader else: - pixel_values = video_reader.get_batch(batch_index).asnumpy() + pixel_values = pixel_values if self.is_image: pixel_values = pixel_values[0] @@ -165,21 +189,32 @@ class VideoDataset(Dataset): video_dir = video_id else: video_dir = os.path.join(self.video_folder, video_id) - video_reader = VideoReader(video_dir) - video_length = len(video_reader) - - clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1) - start_idx = random.randint(0, video_length - clip_length) - batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int) - if not self.enable_bucket: - pixel_values = torch.from_numpy(video_reader.get_batch(batch_index).asnumpy()).permute(0, 3, 1, 2).contiguous() - pixel_values = pixel_values / 255. - del video_reader - else: - pixel_values = video_reader.get_batch(batch_index).asnumpy() + with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader: + video_length = len(video_reader) + + clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1) + start_idx = random.randint(0, video_length - clip_length) + batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int) - return pixel_values, name + try: + sample_args = (video_reader, batch_index) + pixel_values = func_timeout( + VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args + ) + except FunctionTimedOut: + raise ValueError(f"Read {idx} timeout.") + except Exception as e: + raise ValueError(f"Failed to extract frames from video. Error is {e}.") + + if not self.enable_bucket: + pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + del video_reader + else: + pixel_values = pixel_values + + return pixel_values, name def __len__(self): return self.length diff --git a/easyanimate/ui/ui.py b/easyanimate/ui/ui.py index 7dea125..b3df43a 100644 --- a/easyanimate/ui/ui.py +++ b/easyanimate/ui/ui.py @@ -93,11 +93,13 @@ class EasyAnimateController: if edition == "v1": self.inference_config = OmegaConf.load(os.path.join(self.config_dir, "easyanimate_video_motion_module_v1.yaml")) return gr.Dropdown.update(), gr.update(value="none"), gr.update(visible=True), gr.update(visible=True), \ - gr.update(visible=False), gr.update(value=80, minimum=40, maximum=80, step=1) + gr.update(visible=False), gr.update(value=512, minimum=384, maximum=704, step=32), \ + gr.update(value=512, minimum=384, maximum=704, step=32), gr.update(value=80, minimum=40, maximum=80, step=1) else: self.inference_config = OmegaConf.load(os.path.join(self.config_dir, "easyanimate_video_magvit_motion_module_v2.yaml")) return gr.Dropdown.update(), gr.update(value="none"), gr.update(visible=False), gr.update(visible=False), \ - gr.update(visible=True), gr.update(value=144, minimum=9, maximum=144, step=9) + gr.update(visible=True), gr.update(value=672, minimum=128, maximum=1280, step=16), \ + gr.update(value=384, minimum=128, maximum=1280, step=16), gr.update(value=144, minimum=9, maximum=144, step=9) def update_diffusion_transformer(self, diffusion_transformer_dropdown): print("Update diffusion transformer") @@ -227,7 +229,7 @@ class EasyAnimateController: torch.cuda.ipc_collect() if self.lora_model_path != "none": self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) - return gr.Video.update() + return gr.Image.update(), gr.Video.update() # lora part if self.lora_model_path != "none": @@ -274,7 +276,8 @@ def ui(): with gr.Blocks(css=css) as demo: gr.Markdown( """ - # EasyAnimate: Generate your animation easily + # EasyAnimate: Integrated generation of baseline scheme for videos and images. + Generate your videos easily [Github](https://github.com/aigc-apps/EasyAnimate/) """ ) @@ -374,8 +377,8 @@ def ui(): sampler_dropdown = gr.Dropdown(label="Sampling method", choices=list(scheduler_dict.keys()), value=list(scheduler_dict.keys())[0]) sample_step_slider = gr.Slider(label="Sampling steps", value=30, minimum=10, maximum=100, step=1) - width_slider = gr.Slider(label="Width", value=672, minimum=256, maximum=1024, step=32) - height_slider = gr.Slider(label="Height", value=384, minimum=256, maximum=1024, step=32) + width_slider = gr.Slider(label="Width", value=672, minimum=128, maximum=1280, step=16) + height_slider = gr.Slider(label="Height", value=384, minimum=128, maximum=1280, step=16) with gr.Row(): is_image = gr.Checkbox(False, label="Generate Image") length_slider = gr.Slider(label="Animation length", value=144, minimum=9, maximum=144, step=9) @@ -405,7 +408,9 @@ def ui(): motion_module_dropdown, motion_module_refresh_button, is_image, - length_slider + width_slider, + height_slider, + length_slider, ] ) generate_button.click( @@ -538,7 +543,8 @@ def ui_modelscope(edition, config_path, model_name, savedir_sample): with gr.Blocks(css=css) as demo: gr.Markdown( """ - # EasyAnimate: Generate your animation easily + # EasyAnimate: Integrated generation of baseline scheme for videos and images. + Generate your videos easily [Github](https://github.com/aigc-apps/EasyAnimate/) """ ) @@ -553,15 +559,15 @@ def ui_modelscope(edition, config_path, model_name, savedir_sample): sample_step_slider = gr.Slider(label="Sampling steps", value=30, minimum=10, maximum=100, step=1) if edition == "v1": - width_slider = gr.Slider(label="Width", value=512, minimum=384, maximum=704, step=64) - height_slider = gr.Slider(label="Height", value=512, minimum=384, maximum=704, step=64) + width_slider = gr.Slider(label="Width", value=512, minimum=384, maximum=704, step=32) + height_slider = gr.Slider(label="Height", value=512, minimum=384, maximum=704, step=32) with gr.Row(): is_image = gr.Checkbox(False, label="Generate Image", visible=False) length_slider = gr.Slider(label="Animation length", value=80, minimum=40, maximum=96, step=1) cfg_scale_slider = gr.Slider(label="CFG Scale", value=6.0, minimum=0, maximum=20) else: - width_slider = gr.Slider(label="Width", value=672, minimum=256, maximum=1024, step=32) - height_slider = gr.Slider(label="Height", value=384, minimum=256, maximum=1024, step=32) + width_slider = gr.Slider(label="Width", value=672, minimum=384, maximum=704, step=16) + height_slider = gr.Slider(label="Height", value=384, minimum=384, maximum=704, step=16) with gr.Row(): is_image = gr.Checkbox(False, label="Generate Image") length_slider = gr.Slider(label="Animation length", value=144, minimum=9, maximum=144, step=9)