From 5ba510e0612e0baac7de5b9667bbc85609f2a6ed Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Sun, 17 Nov 2024 10:18:20 +0530 Subject: [PATCH] updates --- training/dataset.py | 54 ++++++++++++++++------------- training/mochi-1/prepare_dataset.py | 18 +++++----- training/mochi-1/prepare_dataset.sh | 48 +++++++++++++++++++++++++ 3 files changed, 88 insertions(+), 32 deletions(-) create mode 100644 training/mochi-1/prepare_dataset.sh diff --git a/training/dataset.py b/training/dataset.py index 7ace5c0..1d250dc 100644 --- a/training/dataset.py +++ b/training/dataset.py @@ -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 diff --git a/training/mochi-1/prepare_dataset.py b/training/mochi-1/prepare_dataset.py index 6e5fea9..777f7ae 100644 --- a/training/mochi-1/prepare_dataset.py +++ b/training/mochi-1/prepare_dataset.py @@ -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: diff --git a/training/mochi-1/prepare_dataset.sh b/training/mochi-1/prepare_dataset.sh new file mode 100644 index 0000000..8081ac2 --- /dev/null +++ b/training/mochi-1/prepare_dataset.sh @@ -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"