update datasets loader

This commit is contained in:
bubbliiiing
2024-05-27 13:54:29 +08:00
parent f4286edee0
commit e6336652aa
2 changed files with 68 additions and 27 deletions
+50 -15
View File
@@ -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
+18 -12
View File
@@ -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)