From 2fde026d30e34dc029b0e8a7a494ff51e52e8e93 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Thu, 28 Nov 2024 12:33:25 +0530 Subject: [PATCH] updates --- training/mochi-1/args.py | 150 ++--- training/mochi-1/dataset_mochi.py | 245 -------- training/mochi-1/dataset_simple.py | 50 ++ training/mochi-1/embed.py | 111 ++++ training/mochi-1/prepare_dataset.py | 682 ----------------------- training/mochi-1/prepare_dataset.sh | 50 +- training/mochi-1/text_to_video_lora.py | 375 ++++--------- training/mochi-1/train.sh | 46 +- training/mochi-1/trim_and_crop_videos.py | 126 +++++ 9 files changed, 463 insertions(+), 1372 deletions(-) delete mode 100644 training/mochi-1/dataset_mochi.py create mode 100644 training/mochi-1/dataset_simple.py create mode 100644 training/mochi-1/embed.py delete mode 100644 training/mochi-1/prepare_dataset.py create mode 100644 training/mochi-1/trim_and_crop_videos.py diff --git a/training/mochi-1/args.py b/training/mochi-1/args.py index 1d41f5a..46a8a17 100644 --- a/training/mochi-1/args.py +++ b/training/mochi-1/args.py @@ -28,6 +28,11 @@ def _get_model_args(parser: argparse.ArgumentParser) -> None: default=None, help="The directory where the downloaded models and datasets will be stored.", ) + parser.add_argument( + "--cast_dit", + action="store_true", + help="If we should cast DiT params to a lower precision.", + ) def _get_dataset_args(parser: argparse.ArgumentParser) -> None: @@ -38,58 +43,12 @@ def _get_dataset_args(parser: argparse.ArgumentParser) -> None: help=("A folder containing the training data."), ) parser.add_argument( - "--dataset_file", - type=str, - default=None, - help=("Path to a CSV file if loading prompts/video paths using this format."), - ) - parser.add_argument( - "--video_column", - type=str, - default="video", - help="The column of the dataset containing videos. Or, the name of the file in `--data_root` folder containing the line-separated path to video data.", - ) - parser.add_argument( - "--caption_column", - type=str, - default="text", - help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.", - ) - parser.add_argument( - "--id_token", - type=str, - default=None, - help="Identifier token appended to the start of each prompt if provided.", - ) - parser.add_argument( - "--height_buckets", - nargs="+", - type=int, - default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536], - ) - parser.add_argument( - "--width_buckets", - nargs="+", - type=int, - default=[256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536], - ) - parser.add_argument( - "--frame_buckets", - nargs="+", - type=int, - default=[84], - ) - parser.add_argument( - "--load_tensors", - action="store_true", - help="Whether to use a pre-encoded tensor dataset of latents and prompt embeddings instead of videos and text prompts. The expected format is that saved by running the `prepare_dataset.py` script.", - ) - parser.add_argument( - "--random_flip", + "--caption_dropout", type=float, default=None, - help="If random horizontal flip augmentation is to be used, this should be the flip probability.", + help=("Probability to drop out captions randomly."), ) + parser.add_argument( "--dataloader_num_workers", type=int, @@ -140,15 +99,31 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None: default=False, help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.", ) + parser.add_argument( + "--fps", + type=int, + default=30, + help="FPS to use when serializing the output videos.", + ) + parser.add_argument( + "--height", + type=int, + default=480, + ) + parser.add_argument( + "--width", + type=int, + default=848, + ) def _get_training_args(parser: argparse.ArgumentParser) -> None: parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") - parser.add_argument("--rank", type=int, default=64, help="The rank for LoRA matrices.") + parser.add_argument("--rank", type=int, default=16, help="The rank for LoRA matrices.") parser.add_argument( "--lora_alpha", type=int, - default=64, + default=16, help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.", ) parser.add_argument( @@ -156,7 +131,7 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: nargs="+", type=str, default=["to_k", "to_q", "to_v", "to_out.0"], - help="Target modules to train LoRA for." + help="Target modules to train LoRA for.", ) parser.add_argument( "--mixed_precision", @@ -175,43 +150,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: default="mochi-lora", help="The output directory where the model predictions and checkpoints will be written.", ) - parser.add_argument( - "--height", - type=int, - default=480, - help="All input videos are resized to this height.", - ) - parser.add_argument( - "--width", - type=int, - default=848, - help="All input videos are resized to this width.", - ) - parser.add_argument( - "--video_reshape_mode", - type=str, - default=None, - help="All input videos are reshaped to this mode. Choose between ['center', 'random', 'none']", - ) - parser.add_argument("--fps", type=int, default=30, help="All input videos will be used at this FPS.") - parser.add_argument( - "--max_num_frames", - type=int, - default=84, - help="All input videos will be truncated to these many frames.", - ) - parser.add_argument( - "--skip_frames_start", - type=int, - default=0, - help="Number of frames to skip from the beginning of each input video. Useful if training data contains intro sequences.", - ) - parser.add_argument( - "--skip_frames_end", - type=int, - default=0, - help="Number of frames to skip from the end of each input video. Useful if training data contains outro sequences.", - ) parser.add_argument( "--train_batch_size", type=int, @@ -256,25 +194,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: default=1, help="Number of updates steps to accumulate before performing a backward/update pass.", ) - parser.add_argument( - "--weighting_scheme", - type=str, - default="none", - choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], - help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), - ) - parser.add_argument( - "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." - ) - parser.add_argument( - "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." - ) - parser.add_argument( - "--mode_scale", - type=float, - default=1.29, - help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", - ) parser.add_argument( "--gradient_checkpointing", action="store_true", @@ -283,19 +202,18 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--learning_rate", type=float, - default=1e-4, + default=2e-4, help="Initial learning rate (after the potential warmup period) to use.", ) parser.add_argument( "--scale_lr", action="store_true", - default=False, help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", ) parser.add_argument( "--lr_scheduler", type=str, - default="constant", + default="cosine", help=( 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' ' "constant", "constant_with_warmup"]' @@ -304,7 +222,7 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--lr_warmup_steps", type=int, - default=500, + default=200, help="Number of steps for the warmup in the lr scheduler.", ) parser.add_argument( @@ -331,12 +249,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: default=False, help="Whether or not to use VAE tiling for saving memory.", ) - parser.add_argument( - "--noised_image_dropout", - type=float, - default=0.05, - help="Image condition dropout probability when finetuning image-to-video.", - ) def _get_optimizer_args(parser: argparse.ArgumentParser) -> None: @@ -386,7 +298,7 @@ def _get_optimizer_args(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--weight_decay", type=float, - default=1e-04, + default=0.01, help="Weight decay to use for optimizer.", ) parser.add_argument( diff --git a/training/mochi-1/dataset_mochi.py b/training/mochi-1/dataset_mochi.py deleted file mode 100644 index 7bab1ce..0000000 --- a/training/mochi-1/dataset_mochi.py +++ /dev/null @@ -1,245 +0,0 @@ -from pathlib import Path -from typing import Any, Dict, Tuple - -import numpy as np -import torch -import torch.nn as nn -from accelerate.logging import get_logger -from torchvision import transforms -from torchvision.transforms import InterpolationMode -from torchvision.transforms.functional import resize - - -# Must import after torch because this can sometimes lead to a nasty segmentation fault, or stack smashing error -# Very few bug reports but it happens. Look in decord Github issues for more relevant information. -import decord # isort:skip - -decord.bridge.set_bridge("torch") - -import sys - - -sys.path.append("..") - -from dataset import VideoDataset as VDS - - -logger = get_logger(__name__) - -# TODO (sayakpaul): probably not all buckets are needed for Mochi-1? -HEIGHT_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536] -WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536] -FRAME_BUCKETS = [16, 24, 32, 48, 64, 80, 85] - -VAE_SPATIAL_SCALE_FACTOR = 8 -VAE_TEMPORAL_SCALE_FACTOR = 6 - -class VideoDataset(VDS): - def __init__(self, *args, **kwargs) -> None: - super().__init__(*args, **kwargs) - - random_flip = kwargs.get("random_flip", None) - self.video_transforms = transforms.Compose( - [ - transforms.RandomHorizontalFlip([(random_flip)]) - if random_flip - else transforms.Lambda(lambda x: x), - transforms.Lambda(self.scale_transform), - ] - ) - - def scale_transform(self, x): - return x / 127.5 - 1.0 - - # 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. - # The BucketSampler should ideally return indices. But, in the sampler, we'd like - # to have information about num_frames, height and width. Since this is not stored - # as metadata, we need to read the video to get this information. You could read this - # information without loading the full video in memory, but we do it anyway. In order - # to not load the video twice (once to get the metadata, and once to return the loaded video - # based on sampled indices), we cache it in the BucketSampler. When the sampler is - # to yield, we yield the cache data instead of indices. So, this special check ensures - # that data is not loaded a second time. PRs are welcome for improvements. - return index - - if self.load_tensors: - image_latents, video_latents, prompt_embeds, prompt_attention_mask = self._preprocess_video(self.video_paths[index]) - - # This is hardcoded for now. - # Output of the VAE encoding is 2 * output_channels and then it's - # temporal compression factor is 6. Initially, the VAE encodings will have - # 24 latent number of frames. So, if we were to train with a - # max frame size of 85 and frame bucket of [85], we need to have the following logic. - latent_num_frames = video_latents.size(0) - num_frames = ((latent_num_frames // 2) * (VAE_TEMPORAL_SCALE_FACTOR + 1) + 1) - - height = video_latents.size(2) * VAE_SPATIAL_SCALE_FACTOR - width = video_latents.size(3) * VAE_SPATIAL_SCALE_FACTOR - - return { - "prompt": prompt_embeds, - "prompt_attention_mask": prompt_attention_mask, - "image": image_latents, - "video": video_latents, - "video_metadata": { - "num_frames": num_frames, - "height": height, - "width": width, - }, - } - else: - image, video, _ = self._preprocess_video(self.video_paths[index]) - 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], - }, - } - - # 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" - - # The current path is something like: /a/b/c/d/videos/00001.mp4 - # We need to reach: /a/b/c/d/video_latents/00001.pt - image_latents_path = path.parent.parent.joinpath("image_latents") - video_latents_path = path.parent.parent.joinpath("video_latents") - embeds_path = path.parent.parent.joinpath("prompt_embeds") - attention_mask_path = path.parent.parent.joinpath("prompt_attention_mask") - - if ( - not video_latents_path.exists() - or not embeds_path.exists() - or not attention_mask_path.exists() - or (self.image_to_video and not image_latents_path.exists()) - ): - raise ValueError( - f"When setting the load_tensors parameter to `True`, it is expected that the `{self.data_root=}` contains three folders named `video_latents`, `prompt_embeds`, and `prompt_attention_mask`. However, these folders were not found. Please make sure to have prepared your data correctly using `prepare_data.py`. Additionally, if you're training image-to-video, it is expected that an `image_latents` folder is also present." - ) - - if self.image_to_video: - image_latent_filepath = image_latents_path.joinpath(pt_filename) - video_latent_filepath = video_latents_path.joinpath(pt_filename) - embeds_filepath = embeds_path.joinpath(pt_filename) - attention_mask_filepath = attention_mask_path.joinpath(pt_filename) - - if not video_latent_filepath.is_file() or not embeds_filepath.is_file() or not attention_mask_filepath.is_file(): - if self.image_to_video: - image_latent_filepath = image_latent_filepath.as_posix() - video_latent_filepath = video_latent_filepath.as_posix() - embeds_filepath = embeds_filepath.as_posix() - attention_mask_filepath = attention_mask_filepath.as_posix() - raise ValueError( - f"The file {video_latent_filepath=} or {embeds_filepath=} or {attention_mask_filepath=} could not be found. Please ensure that you've correctly executed `prepare_dataset.py`." - ) - - images = ( - torch.load(image_latent_filepath, map_location="cpu", weights_only=True) if self.image_to_video else None - ) - latents = torch.load(video_latent_filepath, map_location="cpu", weights_only=True) - embeds = torch.load(embeds_filepath, map_location="cpu", weights_only=True) - attention_masks = torch.load(attention_mask_filepath, map_location="cpu", weights_only=True) - - return images, latents, embeds, attention_masks - - - -class VideoDatasetWithFlexibleResize(VideoDataset): - def __init__(self, video_reshape_mode: str = None, *args, **kwargs) -> None: - super().__init__(*args, **kwargs) - if video_reshape_mode: - assert video_reshape_mode in ["center", "random"] - self.video_reshape_mode = video_reshape_mode - - def _preprocess_video(self, path: Path) -> torch.Tensor: - 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) - - nearest_frame_bucket = min( - self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)) - ) - frame_indices = list( - range( - 0, - video_num_frames, - 1 if video_num_frames < nearest_frame_bucket else video_num_frames // nearest_frame_bucket - ) - ) - frames = video_reader.get_batch(frame_indices) - - # Pad or truncate frames to match the bucket size - if video_num_frames < nearest_frame_bucket: - pad_size = nearest_frame_bucket - video_num_frames - frames = nn.functional.pad(frames, (0, 0, 0, 0, 0, 0, 0, pad_size)) - frames = frames.float() - else: - frames = frames[:nearest_frame_bucket].float() - frames = frames.permute(0, 3, 1, 2).contiguous() - - # Find nearest resolution and apply resizing - nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3]) - if self.video_reshape_mode in {"center", "random"}: - frames = self._resize_for_rectangle_crop(frames, nearest_res) - else: - frames = torch.stack([resize(frame, nearest_res) for frame in frames], dim=0) - - # Apply transformations - frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0) - - # Optionally extract the first frame as an image - image = frames[:1].clone() if self.image_to_video else None - - return image, frames, None - - def _find_nearest_resolution(self, height: int, width: int) -> Tuple[int, int]: - """ - Find the nearest resolution from the predefined list of resolutions. - """ - nearest_res = min(self.resolutions, key=lambda x: abs(x[1] - height) + abs(x[2] - width)) - return nearest_res[1], nearest_res[2] - - def _resize_for_rectangle_crop(self, arr: torch.Tensor, image_size: Tuple[int, int]) -> torch.Tensor: - """ - Resize frames for rectangular cropping. - - Args: - arr (torch.Tensor): The video frames tensor [N, C, H, W]. - image_size (Tuple[int, int]): The target resolution (height, width). - """ - if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]: - arr = resize( - arr, - size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])], - interpolation=InterpolationMode.BICUBIC, - ) - else: - arr = resize( - arr, - size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]], - interpolation=InterpolationMode.BICUBIC, - ) - - # Perform cropping - h, w = arr.shape[2], arr.shape[3] - delta_h, delta_w = h - image_size[0], w - image_size[1] - - if self.video_reshape_mode == "random": - top, left = np.random.randint(0, delta_h + 1), np.random.randint(0, delta_w + 1) - elif self.video_reshape_mode == "center": - top, left = delta_h // 2, delta_w // 2 - else: - raise NotImplementedError(f"Unsupported reshape mode: {self.video_reshape_mode}") - - return transforms.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1]) diff --git a/training/mochi-1/dataset_simple.py b/training/mochi-1/dataset_simple.py new file mode 100644 index 0000000..8cc6153 --- /dev/null +++ b/training/mochi-1/dataset_simple.py @@ -0,0 +1,50 @@ +""" +Taken from +https://github.com/genmoai/mochi/blob/main/demos/fine_tuner/dataset.py +""" + +from pathlib import Path + +import click +import torch +from torch.utils.data import DataLoader, Dataset + + +def load_to_cpu(x): + return torch.load(x, map_location=torch.device("cpu"), weights_only=True) + + +class LatentEmbedDataset(Dataset): + def __init__(self, file_paths, repeat=1): + self.items = [ + (Path(p).with_suffix(".latent.pt"), Path(p).with_suffix(".embed.pt")) + for p in file_paths + if Path(p).with_suffix(".latent.pt").is_file() and Path(p).with_suffix(".embed.pt").is_file() + ] + self.items = self.items * repeat + print(f"Loaded {len(self.items)}/{len(file_paths)} valid file pairs.") + + def __len__(self): + return len(self.items) + + def __getitem__(self, idx): + latent_path, embed_path = self.items[idx] + return load_to_cpu(latent_path), load_to_cpu(embed_path) + + +@click.command() +@click.argument("directory", type=click.Path(exists=True, file_okay=False)) +def process_videos(directory): + dir_path = Path(directory) + mp4_files = [str(f) for f in dir_path.glob("**/*.mp4") if not f.name.endswith(".recon.mp4")] + assert mp4_files, f"No mp4 files found" + + dataset = LatentEmbedDataset(mp4_files) + dataloader = DataLoader(dataset, batch_size=4, shuffle=True) + + for latents, embeds in dataloader: + print([(k, v.shape) for k, v in latents.items()]) + + +if __name__ == "__main__": + process_videos() diff --git a/training/mochi-1/embed.py b/training/mochi-1/embed.py new file mode 100644 index 0000000..ec35ebb --- /dev/null +++ b/training/mochi-1/embed.py @@ -0,0 +1,111 @@ +""" +Adapted from: +https://github.com/genmoai/mochi/blob/main/demos/fine_tuner/encode_videos.py +https://github.com/genmoai/mochi/blob/main/demos/fine_tuner/embed_captions.py +""" + +import click +import torch +import torchvision +from pathlib import Path +from diffusers import AutoencoderKLMochi, MochiPipeline +from transformers import T5EncoderModel, T5Tokenizer +from tqdm.auto import tqdm + + +def encode_videos(model: torch.nn.Module, vid_path: Path, shape: str): + T, H, W = [int(s) for s in shape.split("x")] + assert (T - 1) % 6 == 0, "Expected T to be 1 mod 6" + video, _, metadata = torchvision.io.read_video(str(vid_path), output_format="THWC", pts_unit="secs") + fps = metadata["video_fps"] + video = video.permute(3, 0, 1, 2) + og_shape = video.shape + assert video.shape[2] == H, f"Expected {vid_path} to have height {H}, got {video.shape}" + assert video.shape[3] == W, f"Expected {vid_path} to have width {W}, got {video.shape}" + assert video.shape[1] >= T, f"Expected {vid_path} to have at least {T} frames, got {video.shape}" + if video.shape[1] > T: + video = video[:, :T] + print(f"Trimmed video from {og_shape[1]} to first {T} frames") + video = video.unsqueeze(0) + video = video.float() / 127.5 - 1.0 + video = video.to(model.device) + + assert video.ndim == 5 + + with torch.inference_mode(): + with torch.autocast("cuda", dtype=torch.bfloat16): + ldist = model._encode(video) + + torch.save(dict(ldist=ldist), vid_path.with_suffix(".latent.pt")) + + +@click.command() +@click.argument("output_dir", type=click.Path(exists=True, file_okay=False, dir_okay=True, path_type=Path)) +@click.option( + "--model_id", + type=str, + help="Repo id. Should be genmo/mochi-1-preview", + default="genmo/mochi-1-preview", +) +@click.option("--shape", default="163x480x848", help="Shape of the video to encode") +@click.option("--overwrite", "-ow", is_flag=True, help="Overwrite existing latents and caption embeddings.") +def batch_process(output_dir: Path, model_id: Path, shape: str, overwrite: bool) -> None: + """Process all videos and captions in a directory using a single GPU.""" + # comment out when running on unsupported hardware + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + + # Get all video paths + video_paths = list(output_dir.glob("**/*.mp4")) + if not video_paths: + print(f"No MP4 files found in {output_dir}") + return + + text_paths = list(output_dir.glob("**/*.txt")) + if not text_paths: + print(f"No text files found in {output_dir}") + return + + # load the models + vae = AutoencoderKLMochi.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32).to("cuda") + text_encoder = T5EncoderModel.from_pretrained(model_id, subfolder="text_encoder") + tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer") + pipeline = MochiPipeline.from_pretrained( + model_id, text_encoder=text_encoder, tokenizer=tokenizer, transformer=None, vae=None + ).to("cuda") + + for idx, video_path in tqdm(enumerate(sorted(video_paths))): + print(f"Processing {video_path}") + try: + if video_path.with_suffix(".latent.pt").exists() and not overwrite: + print(f"Skipping {video_path}") + continue + + # encode videos. + encode_videos(vae, vid_path=video_path, shape=shape) + + # embed captions. + prompt_path = Path("/".join(str(video_path).split(".")[:-1]) + ".txt") + embed_path = prompt_path.with_suffix(".embed.pt") + + if embed_path.exists() and not overwrite: + print(f"Skipping {prompt_path} - embeddings already exist") + continue + + with open(prompt_path) as f: + text = f.read().strip() + with torch.inference_mode(): + conditioning = pipeline.encode_prompt(prompt=[text]) + + conditioning = {"prompt_embeds": conditioning[0], "prompt_attention_mask": conditioning[1]} + torch.save(conditioning, embed_path) + + except Exception as e: + import traceback + + traceback.print_exc() + print(f"Error processing {video_path}: {str(e)}") + + +if __name__ == "__main__": + batch_process() diff --git a/training/mochi-1/prepare_dataset.py b/training/mochi-1/prepare_dataset.py deleted file mode 100644 index 0e710db..0000000 --- a/training/mochi-1/prepare_dataset.py +++ /dev/null @@ -1,682 +0,0 @@ -#!/usr/bin/env python3 - -import argparse -import functools -import json -import os -import pathlib -import queue -import traceback -import uuid -from concurrent.futures import ThreadPoolExecutor -from contextlib import nullcontext -from typing import Any, Dict, List, Optional, Union - -import torch -import torch.distributed as dist -from dataset_mochi import VideoDatasetWithFlexibleResize -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 -from torchvision import transforms -from tqdm import tqdm -from transformers import T5EncoderModel, T5Tokenizer - - -import decord # isort:skip - -decord.bridge.set_bridge("torch") - -import sys - - -sys.path.append("..") - -from dataset import BucketSampler - - -logger = get_logger(__name__) - -DTYPE_MAPPING = { - "fp32": torch.float32, - "fp16": torch.float16, - "bf16": torch.bfloat16, -} - - -def check_height(x: Any) -> int: - x = int(x) - if x % 16 != 0: - raise argparse.ArgumentTypeError( - f"`--height_buckets` must be divisible by 16, but got {x} which does not fit criteria." - ) - return x - - -def check_width(x: Any) -> int: - x = int(x) - if x % 16 != 0: - raise argparse.ArgumentTypeError( - f"`--width_buckets` must be divisible by 16, but got {x} which does not fit criteria." - ) - return x - - -def check_frames(x: Any) -> int: - x = int(x) - if x % 4 != 0 and x % 4 != 1: - raise argparse.ArgumentTypeError( - f"`--frames_buckets` must be of form `4 * k` or `4 * k + 1`, but got {x} which does not fit criteria." - ) - return x - - -def get_args() -> Dict[str, Any]: - parser = argparse.ArgumentParser() - parser.add_argument( - "--model_id", - type=str, - default="genmo/mochi-1-preview", - help="Hugging Face model ID to use for tokenizer, text encoder and VAE.", - ) - parser.add_argument("--data_root", type=str, required=True, help="Path to where training data is located.") - parser.add_argument( - "--dataset_file", type=str, default=None, help="Path to CSV file containing metadata about training data." - ) - parser.add_argument( - "--caption_column", - type=str, - default="caption", - help="If using a CSV file via the `--dataset_file` argument, this should be the name of the column containing the captions. If using the folder structure format for data loading, this should be the name of the file containing line-separated captions (the file should be located in `--data_root`).", - ) - parser.add_argument( - "--video_column", - type=str, - default="video", - help="If using a CSV file via the `--dataset_file` argument, this should be the name of the column containing the video paths. If using the folder structure format for data loading, this should be the name of the file containing line-separated video paths (the file should be located in `--data_root`).", - ) - parser.add_argument( - "--id_token", - type=str, - default=None, - help="Identifier token appended to the start of each prompt if provided.", - ) - parser.add_argument( - "--height_buckets", - nargs="+", - type=check_height, - default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536], - ) - parser.add_argument( - "--width_buckets", - nargs="+", - type=check_width, - default=[256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536], - ) - parser.add_argument( - "--frame_buckets", - nargs="+", - type=check_frames, - default=[84], - ) - parser.add_argument( - "--random_flip", - type=float, - default=None, - help="If random horizontal flip augmentation is to be used, this should be the flip probability.", - ) - parser.add_argument( - "--dataloader_num_workers", - type=int, - default=0, - help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.", - ) - parser.add_argument( - "--pin_memory", - action="store_true", - help="Whether or not to use the pinned memory setting in pytorch dataloader.", - ) - parser.add_argument( - "--video_reshape_mode", - type=str, - default=None, - help="All input videos are reshaped to this mode. Choose between ['center', 'random', 'none']", - ) - parser.add_argument( - "--save_image_latents", - action="store_true", - help="Whether or not to encode and store image latents, which are required for image-to-video finetuning. The image latents are the first frame of input videos encoded with the VAE.", - ) - parser.add_argument( - "--output_dir", - type=str, - required=True, - help="Path to output directory where preprocessed videos/latents/embeddings will be saved.", - ) - parser.add_argument("--max_num_frames", type=int, default=84, help="Maximum number of frames in output video.") - parser.add_argument( - "--max_sequence_length", type=int, default=256, help="Max sequence length of prompt embeddings." - ) - parser.add_argument("--target_fps", type=int, default=30, help="Frame rate of output videos.") - parser.add_argument( - "--save_latents_and_embeddings", - action="store_true", - help="Whether to encode videos/captions to latents/embeddings and save them in pytorch serializable format.", - ) - parser.add_argument( - "--use_slicing", - action="store_true", - help="Whether to enable sliced encoding/decoding in the VAE. Only used if `--save_latents_and_embeddings` is also used.", - ) - parser.add_argument( - "--use_tiling", - action="store_true", - help="Whether to enable tiled encoding/decoding in the VAE. Only used if `--save_latents_and_embeddings` is also used.", - ) - parser.add_argument("--batch_size", type=int, default=1, help="Number of videos to process at once in the VAE.") - parser.add_argument( - "--num_decode_threads", - type=int, - default=0, - help="Number of decoding threads for `decord` to use. The default `0` means to automatically determine required number of threads.", - ) - parser.add_argument( - "--dtype", - type=str, - choices=["fp32", "fp16", "bf16"], - default="fp32", - help="Data type to use when generating latents and prompt embeddings.", - ) - parser.add_argument("--seed", type=int, default=42, help="Seed for reproducibility.") - parser.add_argument( - "--num_artifact_workers", type=int, default=4, help="Number of worker threads for serializing artifacts." - ) - return parser.parse_args() - - -def _get_t5_prompt_embeds( - tokenizer: T5Tokenizer, - text_encoder: T5EncoderModel, - prompt: Union[str, List[str]], - num_videos_per_prompt: int = 1, - max_sequence_length: int = 226, - device: Optional[torch.device] = None, - dtype: Optional[torch.dtype] = None, - text_input_ids=None, -): - prompt = [prompt] if isinstance(prompt, str) else prompt - batch_size = len(prompt) - - if tokenizer is not None: - text_inputs = tokenizer( - prompt, - padding="max_length", - max_length=max_sequence_length, - truncation=True, - add_special_tokens=True, - return_tensors="pt", - ) - text_input_ids = text_inputs.input_ids - prompt_attention_mask = text_inputs.attention_mask - prompt_attention_mask = prompt_attention_mask.bool() - else: - if text_input_ids is None: - raise ValueError("`text_input_ids` must be provided when the tokenizer is not specified.") - - prompt_embeds = text_encoder(text_input_ids.to(device))[0] - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - - _, seq_len, _ = prompt_embeds.shape - prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) - prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) - prompt_attention_mask = prompt_attention_mask.view(batch_size, -1) - prompt_attention_mask = prompt_attention_mask.repeat(num_videos_per_prompt, 1) - - return prompt_embeds, prompt_attention_mask - - -def encode_prompt( - tokenizer: T5Tokenizer, - text_encoder: T5EncoderModel, - prompt: Union[str, List[str]], - num_videos_per_prompt: int = 1, - max_sequence_length: int = 256, - device: Optional[torch.device] = None, - dtype: Optional[torch.dtype] = None, - text_input_ids=None, -): - prompt = [prompt] if isinstance(prompt, str) else prompt - prompt_embeds, prompt_attention_mask = _get_t5_prompt_embeds( - tokenizer, - text_encoder, - prompt=prompt, - num_videos_per_prompt=num_videos_per_prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - text_input_ids=text_input_ids, - ) - return prompt_embeds, prompt_attention_mask - - -def compute_prompt_embeddings( - tokenizer: T5Tokenizer, - text_encoder: T5EncoderModel, - prompts: List[str], - max_sequence_length: int, - device: torch.device, - dtype: torch.dtype, - requires_grad: bool = False, -): - ctx = nullcontext() if requires_grad else torch.no_grad() - with ctx: - prompt_embeds, prompt_attention_mask = encode_prompt( - tokenizer, - text_encoder, - prompts, - num_videos_per_prompt=1, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - return prompt_embeds, prompt_attention_mask - - -to_pil_image = transforms.ToPILImage(mode="RGB") - - -def save_image(image: torch.Tensor, path: pathlib.Path) -> None: - image = to_pil_image(image) - image.save(path) - - -def save_video(video: torch.Tensor, path: pathlib.Path, fps: int = 8) -> None: - video = [to_pil_image(frame) for frame in video] - export_to_video(video, path, fps=fps) - - -def save_prompt(prompt: str, path: pathlib.Path) -> None: - with open(path, "w", encoding="utf-8") as file: - file.write(prompt) - - -def save_metadata(metadata: Dict[str, Any], path: pathlib.Path) -> None: - with open(path, "w", encoding="utf-8") as file: - file.write(json.dumps(metadata)) - - -@torch.no_grad() -def serialize_artifacts( - batch_size: int, - fps: int, - images_dir: Optional[pathlib.Path] = None, - image_latents_dir: Optional[pathlib.Path] = None, - videos_dir: Optional[pathlib.Path] = None, - video_latents_dir: Optional[pathlib.Path] = None, - prompts_dir: Optional[pathlib.Path] = None, - prompt_embeds_dir: Optional[pathlib.Path] = None, - prompt_attention_mask_dir: Optional[pathlib.Path] = None, - images: Optional[torch.Tensor] = None, - image_latents: Optional[torch.Tensor] = None, - videos: Optional[torch.Tensor] = None, - video_latents: Optional[torch.Tensor] = None, - prompts: Optional[List[str]] = None, - prompt_embeds: Optional[torch.Tensor] = None, - prompt_attention_mask: Optional[torch.Tensor] = None -) -> None: - metadata = [] - for i in range(videos.size(0)): - video = videos[i:i+1] - if video.size(1) == 1: - print(f"{video_latents[i:i+1].shape=}") - metadata_dict = {"num_frames": video.size(1), "height": video.size(3), "width": video.size(4)} - metadata.append(metadata_dict) - - data_folder_mapper_list = [ - (images, images_dir, lambda img, path: save_image(img[0], path), "png"), - (image_latents, image_latents_dir, torch.save, "pt"), - (videos, videos_dir, functools.partial(save_video, fps=fps), "mp4"), - (video_latents, video_latents_dir, torch.save, "pt"), - (prompts, prompts_dir, save_prompt, "txt"), - (prompt_embeds, prompt_embeds_dir, torch.save, "pt"), - (prompt_attention_mask, prompt_attention_mask_dir, torch.save, "pt"), - (metadata, videos_dir, save_metadata, "txt"), - ] - filenames = [uuid.uuid4() for _ in range(batch_size)] - - for data, folder, save_fn, extension in data_folder_mapper_list: - if data is None: - continue - for slice, filename in zip(data, filenames): - if isinstance(slice, torch.Tensor): - slice = slice.clone().to("cpu") - path = folder.joinpath(f"{filename}.{extension}") - save_fn(slice, path) - - -def save_intermediates(output_queue: queue.Queue) -> None: - while True: - try: - item = output_queue.get(timeout=30) - if item is None: - break - serialize_artifacts(**item) - - except queue.Empty: - continue - - -@torch.no_grad() -def main(): - args = get_args() - set_seed(args.seed) - - output_dir = pathlib.Path(args.output_dir) - tmp_dir = output_dir.joinpath("tmp") - - output_dir.mkdir(parents=True, exist_ok=True) - tmp_dir.mkdir(parents=True, exist_ok=True) - - # Create task queue for non-blocking serializing of artifacts - output_queue = queue.Queue() - save_thread = ThreadPoolExecutor(max_workers=args.num_artifact_workers) - save_future = save_thread.submit(save_intermediates, output_queue) - - # Initialize distributed processing - if "LOCAL_RANK" in os.environ: - local_rank = int(os.environ["LOCAL_RANK"]) - torch.cuda.set_device(local_rank) - dist.init_process_group(backend="nccl") - world_size = dist.get_world_size() - rank = dist.get_rank() - else: - # Single GPU - local_rank = 0 - world_size = 1 - rank = 0 - torch.cuda.set_device(rank) - - # Create folders where intermediate tensors from each rank will be saved - images_dir = tmp_dir.joinpath(f"images/{rank}") - image_latents_dir = tmp_dir.joinpath(f"image_latents/{rank}") - videos_dir = tmp_dir.joinpath(f"videos/{rank}") - video_latents_dir = tmp_dir.joinpath(f"video_latents/{rank}") - prompts_dir = tmp_dir.joinpath(f"prompts/{rank}") - prompt_embeds_dir = tmp_dir.joinpath(f"prompt_embeds/{rank}") - prompt_attention_mask_dir = tmp_dir.joinpath(f"prompt_attention_mask/{rank}") - - images_dir.mkdir(parents=True, exist_ok=True) - image_latents_dir.mkdir(parents=True, exist_ok=True) - videos_dir.mkdir(parents=True, exist_ok=True) - video_latents_dir.mkdir(parents=True, exist_ok=True) - prompts_dir.mkdir(parents=True, exist_ok=True) - prompt_embeds_dir.mkdir(parents=True, exist_ok=True) - prompt_attention_mask_dir.mkdir(parents=True, exist_ok=True) - - weight_dtype = DTYPE_MAPPING[args.dtype] - target_fps = args.target_fps - - if weight_dtype is not None: - weight_dtype = torch.float32 - print("To get the best results, we set `weight_dtype` to `torch.float32`.") - - # 1. Dataset - dataset_init_kwargs = { - "data_root": args.data_root, - "dataset_file": args.dataset_file, - "video_reshape_mode": args.video_reshape_mode, - "caption_column": args.caption_column, - "video_column": args.video_column, - "max_num_frames": args.max_num_frames, - "id_token": args.id_token, - "height_buckets": args.height_buckets, - "width_buckets": args.width_buckets, - "frame_buckets": args.frame_buckets, - "load_tensors": False, - "random_flip": args.random_flip, - "image_to_video": args.save_image_latents, - } - dataset = VideoDatasetWithFlexibleResize(**dataset_init_kwargs) - original_dataset_size = len(dataset) - - # Split data among GPUs - if world_size > 1: - samples_per_gpu = original_dataset_size // world_size - start_index = rank * samples_per_gpu - end_index = start_index + samples_per_gpu - if rank == world_size - 1: - end_index = original_dataset_size # Make sure the last GPU gets the remaining data - - # Slice the data - dataset.prompts = dataset.prompts[start_index:end_index] - dataset.video_paths = dataset.video_paths[start_index:end_index] - else: - pass - - rank_dataset_size = len(dataset) - - # 2. Dataloader - def collate_fn(data): - prompts = [x["prompt"] for x in data[0]] - - images = None - if args.save_image_latents: - images = [x["image"] for x in data[0]] - images = torch.stack(images).to(dtype=weight_dtype, non_blocking=True) - - videos = [x["video"] for x in data[0]] - videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True) - - return { - "images": images, - "videos": videos, - "prompts": prompts, - } - - dataloader = DataLoader( - dataset, - batch_size=1, - sampler=BucketSampler(dataset, batch_size=args.batch_size, shuffle=True, drop_last=False), - collate_fn=collate_fn, - num_workers=args.dataloader_num_workers, - pin_memory=args.pin_memory, - ) - - # 3. Prepare models - device = f"cuda:{rank}" - - if args.save_latents_and_embeddings: - tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer") - text_encoder = T5EncoderModel.from_pretrained( - args.model_id, subfolder="text_encoder", torch_dtype=weight_dtype - ).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: - vae.enable_tiling() - - # 4. Compute latents and embeddings and save - if rank == 0: - iterator = tqdm( - dataloader, desc="Encoding", total=(rank_dataset_size + args.batch_size - 1) // args.batch_size - ) - else: - iterator = dataloader - - for step, batch in enumerate(iterator): - try: - images = None - image_latents = None - video_latents = None - prompt_embeds = None - - if args.save_image_latents: - images = batch["images"].to(device, non_blocking=True) - images = images.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] - - videos = batch["videos"].to(device, non_blocking=True) - videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] - - prompts = batch["prompts"] - - # Encode videos & images - # we run under autocast following the official recommendations of Mochi - with torch.autocast(device, torch.bfloat16, cache_enabled=False): - if args.save_latents_and_embeddings: - if args.use_slicing: - if args.save_image_latents: - encoded_slices = [vae._encode(image_slice) for image_slice in images.split(1)] - image_latents = torch.cat(encoded_slices) - image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype) - - encoded_slices = [vae._encode(video_slice) for video_slice in videos.split(1)] - video_latents = torch.cat(encoded_slices) - - else: - if args.save_image_latents: - image_latents = vae._encode(images) - image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype) - - video_latents = vae._encode(videos) - - video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype) - - # Encode prompts - prompt_embeds, prompt_attention_mask = compute_prompt_embeddings( - tokenizer, - text_encoder, - prompts, - args.max_sequence_length, - device, - weight_dtype, - requires_grad=False, - ) - - if images is not None: - images = (images.permute(0, 2, 1, 3, 4) + 1) / 2 - - videos = (videos.permute(0, 2, 1, 3, 4) + 1) / 2 - - output_queue.put( - { - "batch_size": len(prompts), - "fps": target_fps, - "images_dir": images_dir, - "image_latents_dir": image_latents_dir, - "videos_dir": videos_dir, - "video_latents_dir": video_latents_dir, - "prompts_dir": prompts_dir, - "prompt_embeds_dir": prompt_embeds_dir, - "prompt_attention_mask_dir": prompt_attention_mask_dir, - "images": images, - "image_latents": image_latents, - "videos": videos, - "video_latents": video_latents, - "prompts": prompts, - "prompt_embeds": prompt_embeds, - "prompt_attention_mask": prompt_attention_mask, - } - ) - - except Exception: - print("-------------------------") - print(f"An exception occurred while processing data: {rank=}, {world_size=}, {step=}") - traceback.print_exc() - print("-------------------------") - - # 5. Complete distributed processing - if world_size > 1: - dist.barrier() - dist.destroy_process_group() - - output_queue.put(None) - save_thread.shutdown(wait=True) - save_future.result() - - # 6. Combine results from each rank - if rank == 0: - print( - f"Completed preprocessing latents and embeddings. Temporary files from all ranks saved to `{tmp_dir.as_posix()}`" - ) - - # Move files from each rank to common directory - for subfolder, extension in [ - ("images", "png"), - ("image_latents", "pt"), - ("videos", "mp4"), - ("video_latents", "pt"), - ("prompts", "txt"), - ("prompt_embeds", "pt"), - ("prompt_attention_mask", "pt"), - ("videos", "txt"), - ]: - tmp_subfolder = tmp_dir.joinpath(subfolder) - combined_subfolder = output_dir.joinpath(subfolder) - combined_subfolder.mkdir(parents=True, exist_ok=True) - pattern = f"*.{extension}" - - for file in tmp_subfolder.rglob(pattern): - file.replace(combined_subfolder / file.name) - - # Remove temporary directories - def rmdir_recursive(dir: pathlib.Path) -> None: - for child in dir.iterdir(): - if child.is_file(): - child.unlink() - else: - rmdir_recursive(child) - dir.rmdir() - - rmdir_recursive(tmp_dir) - - # Combine prompts and videos into individual text files and single jsonl - prompts_folder = output_dir.joinpath("prompts") - prompts = [] - stems = [] - - for filename in prompts_folder.rglob("*.txt"): - with open(filename, "r") as file: - prompts.append(file.read().strip()) - stems.append(filename.stem) - - prompts_txt = output_dir.joinpath("prompts.txt") - videos_txt = output_dir.joinpath("videos.txt") - data_jsonl = output_dir.joinpath("data.jsonl") - - with open(prompts_txt, "w") as file: - for prompt in prompts: - file.write(f"{prompt}\n") - - with open(videos_txt, "w") as file: - for stem in stems: - file.write(f"videos/{stem}.mp4\n") - - with open(data_jsonl, "w") as file: - for prompt, stem in zip(prompts, stems): - video_metadata_txt = output_dir.joinpath(f"videos/{stem}.txt") - with open(video_metadata_txt, "r", encoding="utf-8") as metadata_file: - metadata = json.loads(metadata_file.read()) - - data = { - "prompt": prompt, - "prompt_embed": f"prompt_embeds/{stem}.pt", - "prompt_attention_mask": f"prompt_attention_mask/{stem}.pt", - "image": f"images/{stem}.png", - "image_latent": f"image_latents/{stem}.pt", - "video": f"videos/{stem}.mp4", - "video_latent": f"video_latents/{stem}.pt", - "metadata": metadata, - } - file.write(json.dumps(data) + "\n") - - print(f"Completed preprocessing. All files saved to `{output_dir.as_posix()}`") - - -if __name__ == "__main__": - main() diff --git a/training/mochi-1/prepare_dataset.sh b/training/mochi-1/prepare_dataset.sh index c2aba5c..7d6d064 100644 --- a/training/mochi-1/prepare_dataset.sh +++ b/training/mochi-1/prepare_dataset.sh @@ -1,49 +1,9 @@ #!/bin/bash -MODEL_ID="genmo/mochi-1-preview" +GPU_ID=0 +VIDEO_DIR=/home/sayak/cogvideox-factory/video-dataset-disney-organized +OUTPUT_DIR=videos_prepared -NUM_GPUS=1 +python trim_and_crop_videos.py $VIDEO_DIR $OUTPUT_DIR --num_frames=37 --resolution=480x848 --force_upsample -# 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="85" -MAX_NUM_FRAMES="85" -MAX_SEQUENCE_LENGTH=256 -TARGET_FPS=30 -BATCH_SIZE=4 -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 \ - --use_slicing \ - --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" +CUDA_VISIBLE_DEVICES=$GPU_ID python embed.py $OUTPUT_DIR --shape=37x480x848 \ 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 e73dcd9..874af27 100644 --- a/training/mochi-1/text_to_video_lora.py +++ b/training/mochi-1/text_to_video_lora.py @@ -13,16 +13,17 @@ # See the License for the specific language governing permissions and # limitations under the License. -import copy import gc -import json +import random +from glob import glob import logging import math import os import shutil +import torch.nn.functional as F from datetime import timedelta from pathlib import Path -from typing import Any, Dict +from typing import Any, Dict, Tuple, List import diffusers import torch @@ -44,11 +45,7 @@ from diffusers import ( ) from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.optimization import get_scheduler -from diffusers.training_utils import ( - cast_training_params, - compute_density_for_timestep_sampling, - compute_loss_weighting_for_sd3, -) +from diffusers.training_utils import cast_training_params from diffusers.utils import convert_unet_state_dict_to_peft, export_to_video from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card from diffusers.utils.torch_utils import is_compiled_module @@ -56,20 +53,17 @@ from huggingface_hub import create_repo, upload_folder from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict from torch.utils.data import DataLoader from tqdm.auto import tqdm -from transformers import AutoTokenizer, T5EncoderModel from args import get_args # isort:skip -from dataset_mochi import VideoDatasetWithFlexibleResize # isort:skip +from dataset_simple import LatentEmbedDataset import sys sys.path.append("..") -from dataset import BucketSampler # isort:skip -from text_encoder import compute_prompt_embeddings # isort:skip -from utils import get_gradient_norm, get_optimizer, print_memory, reset_memory # isort:skip +from utils import get_optimizer, print_memory, reset_memory # isort:skip logger = get_logger(__name__) @@ -81,7 +75,7 @@ def save_model_card( base_model: str = None, validation_prompt=None, repo_folder=None, - fps=8, + fps=30, ): widget_dict = [] if videos is not None and len(videos) > 0: @@ -196,29 +190,44 @@ def log_validation( return videos +# Adapted from the original code: +# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/src/genmo/mochi_preview/pipelines.py#L578 +def cast_dit(model, dtype): + for name, module in model.named_modules(): + if isinstance(module, torch.nn.Linear): + assert any( + n in name for n in ["time_embed", "proj_out", "blocks", "norm_out"] + ), f"Unexpected linear layer: {name}" + module.to(dtype=dtype) + elif isinstance(module, torch.nn.Conv2d): + module.to(dtype=dtype) + return model + + class CollateFunction: - def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None: - self.weight_dtype = weight_dtype - self.load_tensors = load_tensors + def __init__(self, caption_dropout: float = None) -> None: + self.caption_dropout = caption_dropout - def __call__(self, data: Dict[str, Any]) -> Dict[str, torch.Tensor]: - prompts = [x["prompt"] for x in data[0]] - prompt_attention_mask = None + def __call__(self, samples: List[Tuple[dict, torch.Tensor]]) -> Dict[str, torch.Tensor]: + ldists = torch.cat([data[0]["ldist"] for data in samples], dim=0) + z = DiagonalGaussianDistribution(ldists).sample() + assert torch.isfinite(z).all() - if self.load_tensors: - prompts = torch.stack(prompts).to(dtype=self.weight_dtype, non_blocking=True) - prompt_attention_mask = torch.stack([x["prompt_attention_mask"] for x in data[0]]) + # Sample noise which we will add to the samples. + eps = torch.randn_like(z) + sigma = torch.rand(z.shape[:1], device="cpu", dtype=torch.float32) - videos = [x["video"] for x in data[0]] - videos = torch.stack(videos).to(dtype=self.weight_dtype, non_blocking=True) + prompt_embeds = torch.cat([data[1]["prompt_embeds"] for data in samples], dim=0) + prompt_attention_mask = torch.cat([data[1]["prompt_attention_mask"] for data in samples], dim=0) + if self.caption_dropout and random.random() < self.caption_dropout: + prompt_embeds.zero_() + prompt_attention_mask = prompt_attention_mask.long() + prompt_attention_mask.zero_() + prompt_attention_mask = prompt_attention_mask.bool() - out_dict = { - "videos": videos, - "prompts": prompts, - } - if prompt_attention_mask is not None: - out_dict.update({"prompt_attention_mask": prompt_attention_mask}) - return out_dict + return dict( + z=z, eps=eps, sigma=sigma, prompt_embeds=prompt_embeds, prompt_attention_mask=prompt_attention_mask + ) def main(args): @@ -281,73 +290,41 @@ def main(args): ).repo_id # Prepare models and scheduler - if not args.load_tensors: - tokenizer = AutoTokenizer.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="tokenizer", - revision=args.revision, - ) - text_encoder = T5EncoderModel.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="text_encoder", - revision=args.revision, - ) - - vae = AutoencoderKLMochi.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="vae", - revision=args.revision, - variant=args.variant, - ) - if args.enable_slicing: - vae.enable_slicing() - if args.enable_tiling: - vae.enable_tiling() - - # keep things in FP32. - text_encoder.requires_grad_(False) - text_encoder.to(accelerator.device, dtype=torch.float32) - vae.requires_grad_(False) - vae.to(accelerator.device, dtype=torch.float32) - - load_dtype = torch.bfloat16 if "5b" in args.pretrained_model_name_or_path.lower() else torch.float16 transformer = MochiTransformer3DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer", - torch_dtype=load_dtype, revision=args.revision, variant=args.variant, ) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") - noise_scheduler_copy = copy.deepcopy(scheduler) + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, subfolder="scheduler" + ) vae_config = AutoencoderKLMochi.load_config(args.pretrained_model_name_or_path, subfolder="vae") - vae_in_channels = vae_config["latent_channels"] has_latents_mean = "latents_mean" in vae_config and vae_config["latents_mean"] is not None has_latents_std = "latents_std" in vae_config and vae_config["latents_std"] is not None + if has_latents_mean and has_latents_std: + mean = torch.tensor(vae_config["latents_mean"])[:, None, None, None] + std = torch.tensor(vae_config["latents_mean"])[:, None, None, None] - VAE_SCALING_FACTOR = vae_config["scaling_factor"] - - # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision - # as these weights are only used for inference, keeping weights in full precision is not required. weight_dtype = torch.float32 - if accelerator.state.deepspeed_plugin: - # DeepSpeed is handling precision, use what's in the DeepSpeed config - if ( - "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config - and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] - ): - weight_dtype = torch.float16 - if ( - "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config - and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] - ): - weight_dtype = torch.bfloat16 - else: - if accelerator.mixed_precision == "fp16": - weight_dtype = torch.float16 - elif accelerator.mixed_precision == "bf16": - weight_dtype = torch.bfloat16 + # if accelerator.state.deepspeed_plugin: + # # DeepSpeed is handling precision, use what's in the DeepSpeed config + # if ( + # "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config + # and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] + # ): + # weight_dtype = torch.float16 + # if ( + # "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config + # and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] + # ): + # weight_dtype = torch.bfloat16 + # else: + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: # due to pytorch#99272, MPS does not yet support bfloat16. @@ -355,11 +332,12 @@ def main(args): "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." ) - # keep the transformer in FP32. transformer.requires_grad_(False) - transformer.to(accelerator.device, torch.float32) + transformer.to(accelerator.device) if args.gradient_checkpointing: transformer.enable_gradient_checkpointing() + if args.cast_dit: + transformer = cast_dit(transformer, weight_dtype) # now we will add new LoRA weights to the attention layers transformer_lora_config = LoraConfig( @@ -492,43 +470,19 @@ def main(args): accelerator.print(f"Using {optimizer.__class__.__name__} optimizer.") # Dataset and DataLoader - if args.load_tensors and args.id_token: - with open(os.path.join(args.data_root, "data.jsonl")) as f: - contents = [json.loads(jline) for jline in f.read().splitlines()] - parsed_id_token = None - for content in contents: - if "id_token" in content: - parsed_id_token = content["id_token"] - if parsed_id_token is not None and parsed_id_token.strip() != args.id_token.strip(): - raise ValueError( - f"Parsed `id_token` from serialized metadata is {parsed_id_token} and provided `id_token` is {args.id_token}. They should match." - ) + train_vids = list(sorted(glob(f"{args.data_root}/*.mp4"))) + train_vids = [v for v in train_vids if not v.endswith(".recon.mp4")] + accelerator.print(f"Found {len(train_vids)} training videos in {args.data_root}") + assert len(train_vids) > 0, f"No training data found in {args.data_root}" - dataset_init_kwargs = { - "data_root": args.data_root, - "video_reshape_mode": args.video_reshape_mode, - "dataset_file": args.dataset_file, - "caption_column": args.caption_column, - "video_column": args.video_column, - "max_num_frames": args.max_num_frames, - "id_token": args.id_token, - "height_buckets": args.height_buckets, - "width_buckets": args.width_buckets, - "frame_buckets": args.frame_buckets, - "load_tensors": args.load_tensors, - "random_flip": args.random_flip, - } - train_dataset = VideoDatasetWithFlexibleResize(**dataset_init_kwargs) - # keeping things in FP32 for now. - collate_fn = CollateFunction(weight_dtype=torch.float32, load_tensors=args.load_tensors) + collate_fn = CollateFunction(caption_dropout=args.caption_dropout) + train_dataset = LatentEmbedDataset(train_vids, repeat=1) train_dataloader = DataLoader( train_dataset, - batch_size=1, - sampler=BucketSampler(train_dataset, batch_size=args.train_batch_size, shuffle=True), collate_fn=collate_fn, + batch_size=args.train_batch_size, num_workers=args.dataloader_num_workers, pin_memory=args.pin_memory, - prefetch_factor=4, ) # Scheduler and math around the number of training steps. @@ -636,27 +590,6 @@ def main(args): disable=not accelerator.is_local_main_process, ) - # For DeepSpeed training - model_config = transformer.module.config if hasattr(transformer, "module") else transformer.config - - if args.load_tensors: - gc.collect() - torch.cuda.empty_cache() - - def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): - sigmas = noise_scheduler_copy.sigmas.to(device=accelerator.device, dtype=dtype) - schedule_timesteps = noise_scheduler_copy.timesteps.to(accelerator.device) - timesteps = timesteps.to(accelerator.device) - step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] - - sigma = sigmas[step_indices].flatten() - if "invert_sigmas" in noise_scheduler_copy.config and noise_scheduler_copy.config.invert_sigmas: - # https://github.com/huggingface/diffusers/blob/99c0483b67427de467f11aa35d54678fd36a7ea2/src/diffusers/schedulers/scheduling_flow_match_euler_discrete.py#L209 - sigma = 1.0 - sigma - while len(sigma.shape) < n_dim: - sigma = sigma.unsqueeze(-1) - return sigma - for epoch in range(first_epoch, args.num_train_epochs): transformer.train() @@ -664,108 +597,45 @@ def main(args): models_to_accumulate = [transformer] with accelerator.accumulate(models_to_accumulate): - videos = batch["videos"].to(accelerator.device, non_blocking=True) - prompts = batch["prompts"] - if args.load_tensors: - prompt_attention_mask = batch["prompt_attention_mask"] + z = batch["z"] + # revisit + # if has_latents_mean and has_latents_std: + # z = (z - mean.to(z)) / std.to(z) - # Encode videos - if not args.load_tensors: - videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] - latent_dist = vae.encode(videos.to(vae.dtype)).latent_dist - else: - latent_dist = DiagonalGaussianDistribution(videos) - - videos = latent_dist.sample() - if has_latents_mean and has_latents_std: - latents_mean = ( - torch.tensor(vae_config["latents_mean"]).view(1, vae_in_channels, 1, 1, 1).to(videos.device, videos.dtype) - ) - latents_std = ( - torch.tensor(vae_config["latents_std"]).view(1, vae_in_channels, 1, 1, 1).to(videos.device, videos.dtype) - ) - videos = (videos - latents_mean) * VAE_SCALING_FACTOR / latents_std - else: - videos = videos * VAE_SCALING_FACTOR - - # keep in FP32 for now. - videos = videos.to(memory_format=torch.contiguous_format, dtype=torch.float32) - model_input = videos - - # Encode prompts - if not args.load_tensors: - prompt_embeds, prompt_attention_mask = compute_prompt_embeddings( - tokenizer, - text_encoder, - prompts, - model_config.max_text_seq_length, - accelerator.device, - weight_dtype=weight_dtype, - requires_grad=False, - ) - else: - prompt_embeds = prompts.to(weight_dtype) - prompt_attention_mask = prompt_attention_mask.to(accelerator.device) - - # Sample noise that will be added to the latents - noise = torch.randn_like(model_input) - batch_size, num_channels, num_frames, height, width = model_input.shape - - # Sample a random timestep for each image - # for weighting schemes where we sample timesteps non-uniformly - u = compute_density_for_timestep_sampling( - weighting_scheme=args.weighting_scheme, - batch_size=batch_size, - logit_mean=args.logit_mean, - logit_std=args.logit_std, - mode_scale=args.mode_scale, - ) - # indices = (u * noise_scheduler_copy.config.num_train_timesteps).long() - # timesteps = noise_scheduler_copy.timesteps[indices].to(device=model_input.device) - # revisit. - timesteps = (u * noise_scheduler_copy.config.num_train_timesteps) + eps = batch["eps"] + sigma = batch["sigma"] + prompt_embeds = batch["prompt_embeds"] + prompt_attention_mask = batch["prompt_attention_mask"] + sigma_bcthw = sigma[:, None, None, None, None] # [B, 1, 1, 1, 1] # Add noise according to flow matching. # zt = (1 - texp) * x + texp * z1 - sigmas = get_sigmas( - timesteps=noise_scheduler_copy.timesteps[timesteps.long()].to(device=model_input.device), - n_dim=model_input.ndim, - dtype=model_input.dtype - ) - noisy_model_input = (1.0 - sigmas) * model_input + sigmas * noise # do we need to revisit this? - noisy_model_input = noisy_model_input.to(weight_dtype) + z_sigma = (1 - sigma_bcthw) * z + sigma_bcthw * eps + ut = z - eps # Predict the noise residual - actual_num_train_timesteps = float(noise_scheduler_copy.config.num_train_timesteps) - timesteps = (1 - (timesteps / actual_num_train_timesteps)) * actual_num_train_timesteps # revisit - timesteps = timesteps.to(device=model_input.device) - model_pred = transformer( - hidden_states=noisy_model_input, - encoder_hidden_states=prompt_embeds, - encoder_attention_mask=prompt_attention_mask, - timestep=timesteps, - return_dict=False, - )[0] - - # these weighting schemes use a uniform timestep sampling - # and instead post-weight the loss - weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) - - # flow matching loss - # target = noise - model_input - target = model_input - noise # as discussed with Ajay - - loss = torch.mean( - (weighting * (model_pred.float() - target.float()) ** 2).reshape(batch_size, -1), - dim=1, - ) - loss = loss.mean() + # (1 - sigma) because of + # https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/src/genmo/mochi_preview/dit/joint_model/asymm_models_joint.py#L656 + # Also, we operate on the scaled version of the `timesteps` directly in the `diffusers` implementation. + timesteps = (1 - sigma) * scheduler.config.num_train_timesteps + with torch.autocast(accelerator.device.type, weight_dtype): + model_pred = transformer( + hidden_states=z_sigma, + encoder_hidden_states=prompt_embeds, + encoder_attention_mask=prompt_attention_mask, + timestep=timesteps, + return_dict=False, + )[0] + assert model_pred.shape == z.shape + loss = F.mse_loss(model_pred.float(), ut.float()) accelerator.backward(loss) - if accelerator.sync_gradients: - gradient_norm_before_clip = get_gradient_norm(transformer_lora_parameters) - accelerator.clip_grad_norm_(transformer_lora_parameters, args.max_grad_norm) - gradient_norm_after_clip = get_gradient_norm(transformer_lora_parameters) + # if accelerator.sync_gradients: + # no grad norm for now, following the original code + # https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner/train.py#L380 + # gradient_norm_before_clip = get_gradient_norm(transformer_lora_parameters) + # accelerator.clip_grad_norm_(transformer_lora_parameters, args.max_grad_norm) + # gradient_norm_after_clip = get_gradient_norm(transformer_lora_parameters) if accelerator.state.deepspeed_plugin is None: optimizer.step() @@ -807,14 +677,14 @@ def main(args): last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate logs = {"loss": loss.detach().item(), "lr": last_lr} - # gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555 - if accelerator.distributed_type != DistributedType.DEEPSPEED: - logs.update( - { - "gradient_norm_before_clip": gradient_norm_before_clip, - "gradient_norm_after_clip": gradient_norm_after_clip, - } - ) + # # gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555 + # if accelerator.distributed_type != DistributedType.DEEPSPEED: + # logs.update( + # { + # "gradient_norm_before_clip": gradient_norm_before_clip, + # "gradient_norm_after_clip": gradient_norm_after_clip, + # } + # ) progress_bar.set_postfix(**logs) accelerator.log(logs, step=global_step) @@ -822,13 +692,14 @@ def main(args): break if global_step >= args.max_train_steps: - break + break if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process: if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0: accelerator.print("===== Memory before validation =====") print_memory(accelerator.device) + transformer.eval() pipe = MochiPipeline.from_pretrained( args.pretrained_model_name_or_path, transformer=unwrap_model(transformer), @@ -848,12 +719,12 @@ def main(args): for validation_prompt in validation_prompts: pipeline_args = { "prompt": validation_prompt, - "guidance_scale": 4.5, + "guidance_scale": 6.0, + "num_inference_steps": 64, "height": args.height, "width": args.width, "max_sequence_length": 256, } - log_validation( pipe=pipe, args=args, @@ -866,13 +737,14 @@ def main(args): print_memory(accelerator.device) reset_memory(accelerator.device) - del pipe.text_encoder del pipe.vae del pipe gc.collect() torch.cuda.empty_cache() + transformer.train() + accelerator.wait_for_everyone() if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process: @@ -884,10 +756,7 @@ def main(args): ) # Cleanup trained models to save memory - if args.load_tensors: - del transformer - else: - del transformer, text_encoder, vae + del transformer gc.collect() torch.cuda.empty_cache() @@ -922,9 +791,11 @@ def main(args): for validation_prompt in validation_prompts: pipeline_args = { "prompt": validation_prompt, - "guidance_scale": 4.5, + "guidance_scale": 6.0, + "num_inference_steps": 64, "height": args.height, "width": args.width, + "max_sequence_length": 256, } video = log_validation( diff --git a/training/mochi-1/train.sh b/training/mochi-1/train.sh index 69c636b..d15ff69 100644 --- a/training/mochi-1/train.sh +++ b/training/mochi-1/train.sh @@ -1,50 +1,38 @@ +#!/bin/bash export NCCL_P2P_DISABLE=1 export TORCH_NCCL_ENABLE_MONITORING=0 GPU_IDS="2" -DATA_ROOT="/home/sayak/cogvideox-factory/video-dataset-disney/mochi-1/preprocessed-dataset" - -CAPTION_COLUMN="prompts.txt" -VIDEO_COLUMN="videos.txt" +DATA_ROOT="/home/sayak/cogvideox-factory/training/mochi-1/videos_prepared" +MODEL="genmo/mochi-1-preview" +OUTPUT_PATH=/raid/.cache/huggingface/sayak/mochi-lora/ cmd="accelerate launch --config_file deepspeed.yaml --gpu_ids $GPU_IDS text_to_video_lora.py \ - --pretrained_model_name_or_path genmo/mochi-1-preview \ + --pretrained_model_name_or_path $MODEL \ --data_root $DATA_ROOT \ - --caption_column $CAPTION_COLUMN \ - --video_column $VIDEO_COLUMN \ - --id_token BW_STYLE \ - --height_buckets 480 \ - --width_buckets 848 \ - --frame_buckets 85 \ - --load_tensors \ --seed 42 \ - --rank 64 \ - --lora_alpha 64 \ - --mixed_precision bf16 \ - --output_dir /raid/.cache/huggingface/sayak/mochi-lora/ \ - --max_num_frames 85 \ + --mixed_precision "bf16" \ + --output_dir $OUTPUT_PATH \ --train_batch_size 1 \ --dataloader_num_workers 4 \ - --max_train_steps 10 \ - --checkpointing_steps 50 \ + --pin_memory \ + --caption_dropout 0.1 \ + --max_train_steps 2000 \ + --checkpointing_steps 200 \ + --checkpoints_total_limit 1 \ --gradient_accumulation_steps 4 \ --gradient_checkpointing \ - --learning_rate 1e-5 \ - --lr_scheduler constant \ - --lr_warmup_steps 0 \ - --lr_num_cycles 1 \ --enable_slicing \ --enable_tiling \ + --enable_model_cpu_offload \ --optimizer adamw --use_8bit \ - --beta1 0.9 \ - --beta2 0.95 \ - --beta3 0.99 \ - --weight_decay 0.001 \ - --max_grad_norm 1.0 \ + --validation_prompt \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions\" \ + --validation_prompt_separator ::: \ + --num_validation_videos 1 \ + --validation_epochs 1 \ --allow_tf32 \ --report_to wandb \ - --push_to_hub \ --nccl_timeout 1800" echo "Running command: $cmd" diff --git a/training/mochi-1/trim_and_crop_videos.py b/training/mochi-1/trim_and_crop_videos.py new file mode 100644 index 0000000..0c6f411 --- /dev/null +++ b/training/mochi-1/trim_and_crop_videos.py @@ -0,0 +1,126 @@ +""" +Adapted from: +https://github.com/genmoai/mochi/blob/main/demos/fine_tuner/trim_and_crop_videos.py +""" + +from pathlib import Path +import shutil + +import click +from moviepy.editor import VideoFileClip +from tqdm import tqdm + + +@click.command() +@click.argument("folder", type=click.Path(exists=True, dir_okay=True)) +@click.argument("output_folder", type=click.Path(dir_okay=True)) +@click.option("--num_frames", "-f", type=float, default=30, help="Number of frames") +@click.option("--resolution", "-r", type=str, default="480x848", help="Video resolution") +@click.option("--force_upsample", is_flag=True, help="Force upsample.") +def truncate_videos(folder, output_folder, num_frames, resolution, force_upsample): + """Truncate all MP4 and MOV files in FOLDER to specified number of frames and resolution""" + input_path = Path(folder) + output_path = Path(output_folder) + output_path.mkdir(parents=True, exist_ok=True) + + # Parse target resolution + target_height, target_width = map(int, resolution.split("x")) + + # Calculate duration + duration = (num_frames / 30) + 0.09 + + # Find all MP4 and MOV files + video_files = ( + list(input_path.rglob("*.mp4")) + + list(input_path.rglob("*.MOV")) + + list(input_path.rglob("*.mov")) + + list(input_path.rglob("*.MP4")) + ) + + for file_path in tqdm(video_files): + try: + relative_path = file_path.relative_to(input_path) + output_file = output_path / relative_path.with_suffix(".mp4") + output_file.parent.mkdir(parents=True, exist_ok=True) + + click.echo(f"Processing: {file_path}") + video = VideoFileClip(str(file_path)) + + # Skip if video is too short + if video.duration < duration: + click.echo(f"Skipping {file_path} as it is too short") + continue + + # Skip if target resolution is larger than input + if target_width > video.w or target_height > video.h: + if force_upsample: + click.echo( + f"{file_path} as target resolution {resolution} is larger than input {video.w}x{video.h}. So, upsampling the video." + ) + video = video.resize(width=target_width, height=target_height) + else: + click.echo( + f"Skipping {file_path} as target resolution {resolution} is larger than input {video.w}x{video.h}" + ) + continue + + # First truncate duration + truncated = video.subclip(0, duration) + + # Calculate crop dimensions to maintain aspect ratio + target_ratio = target_width / target_height + current_ratio = truncated.w / truncated.h + + if current_ratio > target_ratio: + # Video is wider than target ratio - crop width + new_width = int(truncated.h * target_ratio) + x1 = (truncated.w - new_width) // 2 + final = truncated.crop(x1=x1, width=new_width).resize((target_width, target_height)) + else: + # Video is taller than target ratio - crop height + new_height = int(truncated.w / target_ratio) + y1 = (truncated.h - new_height) // 2 + final = truncated.crop(y1=y1, height=new_height).resize((target_width, target_height)) + + # Set output parameters for consistent MP4 encoding + output_params = { + "codec": "libx264", + "audio": False, # Disable audio + "preset": "medium", # Balance between speed and quality + "bitrate": "5000k", # Adjust as needed + } + + # Set FPS to 30 + final = final.set_fps(30) + + # Check for a corresponding .txt file + txt_file_path = file_path.with_suffix(".txt") + if txt_file_path.exists(): + output_txt_file = output_path / relative_path.with_suffix(".txt") + output_txt_file.parent.mkdir(parents=True, exist_ok=True) + shutil.copy(txt_file_path, output_txt_file) + click.echo(f"Copied {txt_file_path} to {output_txt_file}") + else: + # Print warning in bold yellow with a warning emoji + click.echo( + f"\033[1;33m⚠️ Warning: No caption found for {file_path}, using an empty caption. This may hurt fine-tuning quality.\033[0m" + ) + output_txt_file = output_path / relative_path.with_suffix(".txt") + output_txt_file.parent.mkdir(parents=True, exist_ok=True) + output_txt_file.touch() + + # Write the output file + final.write_videofile(str(output_file), **output_params) + + # Clean up + video.close() + truncated.close() + final.close() + + except Exception as e: + click.echo(f"\033[1;31m Error processing {file_path}: {str(e)}\033[0m", err=True) + raise + + +if __name__ == "__main__": + truncate_videos()