This commit is contained in:
sayakpaul
2024-11-17 10:18:20 +05:30
parent a9adf22386
commit 5ba510e061
3 changed files with 88 additions and 32 deletions
+30 -24
View File
@@ -22,8 +22,8 @@ decord.bridge.set_bridge("torch")
logger = get_logger(__name__)
HEIGHT_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
FRAME_BUCKETS = [16, 24, 32, 48, 64, 80]
WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536]
FRAME_BUCKETS = [16, 24, 32, 48, 64, 80, 84]
class VideoDataset(Dataset):
@@ -144,17 +144,17 @@ class VideoDataset(Dataset):
}
else:
image, video, _ = self._preprocess_video(self.video_paths[index])
return {
"prompt": self.id_token + self.prompts[index],
"image": image,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
if video is not None:
return {
"prompt": self.id_token + self.prompts[index],
"image": image,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]:
if not self.data_root.exists():
@@ -276,12 +276,15 @@ class VideoDatasetWithResizing(VideoDataset):
else:
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames))
)
if video_num_frames < nearest_frame_bucket:
# TODO: we could handle this by padding zero frames or duplicating the existing frames?
return None, None, None
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
frames = video_reader.get_batch(frame_indices)
frames = frames[:nearest_frame_bucket].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
@@ -344,6 +347,8 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
nearest_frame_bucket = min(
self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames))
)
if video_num_frames < nearest_frame_bucket:
return None, None, None
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
@@ -404,16 +409,17 @@ class BucketSampler(Sampler):
def __iter__(self):
for index, data in enumerate(self.data_source):
video_metadata = data["video_metadata"]
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
if data is not None:
video_metadata = data["video_metadata"]
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
self.buckets[(f, h, w)].append(data)
if len(self.buckets[(f, h, w)]) == self.batch_size:
if self.shuffle:
random.shuffle(self.buckets[(f, h, w)])
yield self.buckets[(f, h, w)]
del self.buckets[(f, h, w)]
self.buckets[(f, h, w)] = []
self.buckets[(f, h, w)].append(data)
if len(self.buckets[(f, h, w)]) == self.batch_size:
if self.shuffle:
random.shuffle(self.buckets[(f, h, w)])
yield self.buckets[(f, h, w)]
del self.buckets[(f, h, w)]
self.buckets[(f, h, w)] = []
if self.drop_last:
return
+10 -8
View File
@@ -13,7 +13,7 @@ from typing import Any, Dict, List, Optional, Union
import torch
import torch.distributed as dist
from diffusers import AutoencoderKL
from diffusers import AutoencoderKLMochi
from diffusers.training_utils import set_seed
from diffusers.utils import export_to_video, get_logger
from torch.utils.data import DataLoader
@@ -24,7 +24,9 @@ from transformers import T5EncoderModel, T5Tokenizer
import decord # isort:skip
from ..dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
import sys
sys.path.append("..")
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
decord.bridge.set_bridge("torch")
@@ -485,12 +487,12 @@ def main():
tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer")
text_encoder = T5EncoderModel.from_pretrained(
args.model_id, subfolder="text_encoder", torch_dtype=weight_dtype
)
text_encoder = text_encoder.to(device)
vae = AutoencoderKL.from_pretrained(args.model_id, subfolder="vae", torch_dtype=weight_dtype)
vae = vae.to(device)
).to(device)
vae = AutoencoderKLMochi.from_pretrained(
args.model_id, subfolder="vae", torch_dtype=weight_dtype
).to(device)
if args.use_slicing:
vae.enable_slicing()
if args.use_tiling:
+48
View File
@@ -0,0 +1,48 @@
#!/bin/bash
MODEL_ID="genmo/mochi-1-preview"
NUM_GPUS=1
# For more details on the expected data format, please refer to the README.
DATA_ROOT="/home/sayak/cogvideox-factory/video-dataset-disney" # This needs to be the path to the base directory where your videos are located.
CAPTION_COLUMN="prompt.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/home/sayak/cogvideox-factory/video-dataset-disney/mochi-1/preprocessed-dataset"
HEIGHT_BUCKETS="480"
WIDTH_BUCKETS="848"
FRAME_BUCKETS="84"
MAX_NUM_FRAMES="84"
MAX_SEQUENCE_LENGTH=256
TARGET_FPS=30
BATCH_SIZE=1
DTYPE=fp32
# To create a folder-style dataset structure without pre-encoding videos and captions
# For Image-to-Video finetuning, make sure to pass `--save_image_latents`
CMD_WITHOUT_PRE_ENCODING="\
torchrun --nproc_per_node=$NUM_GPUS \
prepare_dataset.py \
--model_id $MODEL_ID \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
--video_column $VIDEO_COLUMN \
--output_dir $OUTPUT_DIR \
--height_buckets $HEIGHT_BUCKETS \
--width_buckets $WIDTH_BUCKETS \
--frame_buckets $FRAME_BUCKETS \
--max_num_frames $MAX_NUM_FRAMES \
--max_sequence_length $MAX_SEQUENCE_LENGTH \
--target_fps $TARGET_FPS \
--batch_size $BATCH_SIZE \
--dtype $DTYPE
"
CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_latents_and_embeddings"
# Select which you'd like to run
CMD=$CMD_WITH_PRE_ENCODING
echo "===== Running \`$CMD\` ====="
eval $CMD
echo -ne "===== Finished running script =====\n"