From 440dc257c85635af8ee93cdfbdd7f294ceeca317 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 19 Nov 2024 10:51:00 +0530 Subject: [PATCH] better reuse. --- training/mochi-1/dataset.py | 232 ++----------------------- training/mochi-1/text_to_video_lora.py | 4 +- 2 files changed, 19 insertions(+), 217 deletions(-) diff --git a/training/mochi-1/dataset.py b/training/mochi-1/dataset.py index 0a95ab7..fa109c4 100644 --- a/training/mochi-1/dataset.py +++ b/training/mochi-1/dataset.py @@ -1,13 +1,10 @@ -import random from pathlib import Path -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Dict, Tuple import numpy as np -import pandas as pd import torch import torchvision.transforms as TT from accelerate.logging import get_logger -from torch.utils.data import Dataset, Sampler from torchvision import transforms from torchvision.transforms import InterpolationMode from torchvision.transforms.functional import resize @@ -19,6 +16,12 @@ import decord # isort:skip decord.bridge.set_bridge("torch") +import sys +sys.path.append("..") + +from dataset import VideoDataset as VDS +from dataset import BucketSampler + logger = get_logger(__name__) # TODO (sayakpaul): probably not all buckets are needed for Mochi-1? @@ -29,84 +32,11 @@ FRAME_BUCKETS = [16, 24, 32, 48, 64, 80, 84] VAE_SPATIAL_SCALE_FACTOR = 8 VAE_TEMPORAL_SCALE_FACTOR = 6 -class VideoDataset(Dataset): - def __init__( - self, - data_root: str, - dataset_file: Optional[str] = None, - caption_column: str = "text", - video_column: str = "video", - max_num_frames: int = 49, - id_token: Optional[str] = None, - height_buckets: List[int] = None, - width_buckets: List[int] = None, - frame_buckets: List[int] = None, - load_tensors: bool = False, - random_flip: Optional[float] = None, - image_to_video: bool = False, - ) -> None: - super().__init__() - - self.data_root = Path(data_root) - self.dataset_file = dataset_file - self.caption_column = caption_column - self.video_column = video_column - self.max_num_frames = max_num_frames - self.id_token = id_token or "" - self.height_buckets = height_buckets or HEIGHT_BUCKETS - self.width_buckets = width_buckets or WIDTH_BUCKETS - self.frame_buckets = frame_buckets or FRAME_BUCKETS - self.load_tensors = load_tensors - self.random_flip = random_flip - self.image_to_video = image_to_video - - self.resolutions = [ - (f, h, w) for h in self.height_buckets for w in self.width_buckets for f in self.frame_buckets - ] - - # Two methods of loading data are supported. - # - Using a CSV: caption_column and video_column must be some column in the CSV. One could - # make use of other columns too, such as a motion score or aesthetic score, by modifying the - # logic in CSV processing. - # - Using two files containing line-separate captions and relative paths to videos. - # For a more detailed explanation about preparing dataset format, checkout the README. - if dataset_file is None: - ( - self.prompts, - self.video_paths, - ) = self._load_dataset_from_local_path() - else: - ( - self.prompts, - self.video_paths, - ) = self._load_dataset_from_csv() - - if len(self.video_paths) != len(self.prompts): - raise ValueError( - f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset." - ) - - self.video_transforms = transforms.Compose( - [ - transforms.RandomHorizontalFlip(random_flip) - if random_flip - else transforms.Lambda(self.identity_transform), - transforms.Lambda(self.scale_transform), - transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), - ] - ) - - @staticmethod - def identity_transform(x): - return x - - @staticmethod - def scale_transform(x): - return x / 255.0 - - def __len__(self) -> int: - return len(self.video_paths) +class VideoDataset(VDS): + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + # Overriding this because we calculate `num_frames` differently. def __getitem__(self, index: int) -> Dict[str, Any]: if isinstance(index, list): # Here, index is actually a list of data objects that we need to return. @@ -159,74 +89,7 @@ class VideoDataset(Dataset): }, } - def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]: - if not self.data_root.exists(): - raise ValueError("Root folder for videos does not exist") - - prompt_path = self.data_root.joinpath(self.caption_column) - video_path = self.data_root.joinpath(self.video_column) - - if not prompt_path.exists() or not prompt_path.is_file(): - raise ValueError( - "Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts." - ) - if not video_path.exists() or not video_path.is_file(): - raise ValueError( - "Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory." - ) - - with open(prompt_path, "r", encoding="utf-8") as file: - prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0] - with open(video_path, "r", encoding="utf-8") as file: - video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0] - - if not self.load_tensors and any(not path.is_file() for path in video_paths): - raise ValueError( - f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." - ) - - return prompts, video_paths - - def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]: - df = pd.read_csv(self.dataset_file) - prompts = df[self.caption_column].tolist() - video_paths = df[self.video_column].tolist() - video_paths = [self.data_root.joinpath(line.strip()) for line in video_paths] - - if any(not path.is_file() for path in video_paths): - raise ValueError( - f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." - ) - - return prompts, video_paths - - def _preprocess_video(self, path: Path) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - r""" - Loads a single video, or latent and prompt embedding, based on initialization parameters. - - If returning a video, returns a [F, C, H, W] video tensor, and None for the prompt embedding. Here, - F, C, H and W are the frames, channels, height and width of the input video. - - If returning latent/embedding, returns a [F, C, H, W] latent, and the prompt embedding of shape [S, D]. - F, C, H and W are the frames, channels, height and width of the latent, and S, D are the sequence length - and embedding dimension of prompt embeddings. - """ - if self.load_tensors: - return self._load_preprocessed_latents_and_embeds(path) - else: - video_reader = decord.VideoReader(uri=path.as_posix()) - video_num_frames = len(video_reader) - - indices = list(range(0, video_num_frames, video_num_frames // self.max_num_frames)) - frames = video_reader.get_batch(indices) - frames = frames[: self.max_num_frames].float() - frames = frames.permute(0, 3, 1, 2).contiguous() - frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0) - - image = frames[:1].clone() if self.image_to_video else None - - return image, frames, None - + # Overriding this because we need `prompt_attention_mask`. def _load_preprocessed_latents_and_embeds(self, path: Path) -> Tuple[torch.Tensor, torch.Tensor]: filename_without_ext = path.name.split(".")[0] pt_filename = f"{filename_without_ext}.pt" @@ -274,6 +137,10 @@ class VideoDataset(Dataset): return images, latents, embeds, attention_masks +# We need the `VideoDatasetWithResizing` and `VideoDatasetWithResizeAndRectangleCrop` classes to subclass from +# the new `VideoDataset` class defined in this file. And also because of the changes in +# `_preprocess_video()` (how we handle `nearest_frame_bucket`). + class VideoDatasetWithResizing(VideoDataset): def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) @@ -373,69 +240,4 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset): def _find_nearest_resolution(self, height, width): nearest_res = min(self.resolutions, key=lambda x: abs(x[1] - height) + abs(x[2] - width)) - return nearest_res[1], nearest_res[2] - - -class BucketSampler(Sampler): - r""" - PyTorch Sampler that groups 3D data by height, width and frames. - - Args: - data_source (`VideoDataset`): - A PyTorch dataset object that is an instance of `VideoDataset`. - batch_size (`int`, defaults to `8`): - The batch size to use for training. - shuffle (`bool`, defaults to `True`): - Whether or not to shuffle the data in each batch before dispatching to dataloader. - drop_last (`bool`, defaults to `False`): - Whether or not to drop incomplete buckets of data after completely iterating over all data - in the dataset. If set to True, only batches that have `batch_size` number of entries will - be yielded. If set to False, it is guaranteed that all data in the dataset will be processed - and batches that do not have `batch_size` number of entries will also be yielded. - """ - - def __init__( - self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False - ) -> None: - self.data_source = data_source - self.batch_size = batch_size - self.shuffle = shuffle - self.drop_last = drop_last - - self.buckets = {resolution: [] for resolution in data_source.resolutions} - - self._raised_warning_for_drop_last = False - - def __len__(self): - if self.drop_last and not self._raised_warning_for_drop_last: - self._raised_warning_for_drop_last = True - logger.warning( - "Calculating the length for bucket sampler is not possible when `drop_last` is set to True. This may cause problems when setting the number of epochs used for training." - ) - return (len(self.data_source) + self.batch_size - 1) // self.batch_size - - def __iter__(self): - for index, data in enumerate(self.data_source): - 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)] = [] - - if self.drop_last: - return - - for fhw, bucket in list(self.buckets.items()): - if len(bucket) == 0: - continue - if self.shuffle: - random.shuffle(bucket) - yield bucket - del self.buckets[fhw] - self.buckets[fhw] = [] + return nearest_res[1], nearest_res[2] \ No newline at end of file diff --git a/training/mochi-1/text_to_video_lora.py b/training/mochi-1/text_to_video_lora.py index 9f42366..d0d1736 100644 --- a/training/mochi-1/text_to_video_lora.py +++ b/training/mochi-1/text_to_video_lora.py @@ -56,11 +56,11 @@ from transformers import AutoTokenizer, T5EncoderModel from args import get_args # isort:skip import sys -sys.path.append("..") +sys.path.append(".") from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip from text_encoder import compute_prompt_embeddings # isort:skip -from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip +from utils import get_gradient_norm, get_optimizer, print_memory, reset_memory # isort:skip logger = get_logger(__name__)