From 2dbddd566e49c600d701dd405d30aad21af07147 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Fri, 22 Nov 2024 23:34:17 +0530 Subject: [PATCH] updates --- training/mochi-1/dataset_mochi.py | 133 +++++++++++++------------ training/mochi-1/prepare_dataset.py | 13 +-- training/mochi-1/prepare_dataset.sh | 4 +- training/mochi-1/text_to_video_lora.py | 30 ++---- training/mochi-1/train.sh | 10 +- 5 files changed, 87 insertions(+), 103 deletions(-) diff --git a/training/mochi-1/dataset_mochi.py b/training/mochi-1/dataset_mochi.py index 2e58f3a..d2af84b 100644 --- a/training/mochi-1/dataset_mochi.py +++ b/training/mochi-1/dataset_mochi.py @@ -3,10 +3,11 @@ from typing import Any, Dict, Tuple import numpy as np import torch -import torchvision.transforms as TT +from torchvision import transforms from accelerate.logging import get_logger from torchvision.transforms import InterpolationMode from torchvision.transforms.functional import resize +import torch.nn as nn # Must import after torch because this can sometimes lead to a nasty segmentation fault, or stack smashing error @@ -25,7 +26,7 @@ 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, 84] +FRAME_BUCKETS = [16, 24, 32, 48, 64, 80, 85] VAE_SPATIAL_SCALE_FACTOR = 8 VAE_TEMPORAL_SCALE_FACTOR = 6 @@ -34,6 +35,19 @@ 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): @@ -55,9 +69,9 @@ class VideoDataset(VDS): # 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 84 and frame bucket of [84], we need to have the following logic. + # 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) + 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 @@ -135,13 +149,13 @@ class VideoDataset(VDS): 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: +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: @@ -151,36 +165,56 @@ class VideoDatasetWithResizing(VideoDataset): video_num_frames = len(video_reader) nearest_frame_bucket = min( - [bucket for bucket in self.frame_buckets if bucket <= video_num_frames], - key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)), - default=1, + 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 + ) ) - - 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() + + # 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]) - frames_resized = torch.stack([resize(frame, nearest_res) for frame in frames], dim=0) - frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0) + 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, width): + 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] - -class VideoDatasetWithResizeAndRectangleCrop(VideoDataset): - def __init__(self, video_reshape_mode: str = "center", *args, **kwargs) -> None: - super().__init__(*args, **kwargs) - self.video_reshape_mode = video_reshape_mode - - def _resize_for_rectangle_crop(self, arr, image_size): - reshape_mode = self.video_reshape_mode + 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, @@ -194,48 +228,15 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset): interpolation=InterpolationMode.BICUBIC, ) + # Perform cropping h, w = arr.shape[2], arr.shape[3] - arr = arr.squeeze(0) + delta_h, delta_w = h - image_size[0], w - image_size[1] - delta_h = h - image_size[0] - delta_w = w - image_size[1] - - if reshape_mode == "random" or reshape_mode == "none": - top = np.random.randint(0, delta_h + 1) - left = np.random.randint(0, delta_w + 1) - elif reshape_mode == "center": + 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 - arr = TT.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1]) - return arr + raise NotImplementedError(f"Unsupported reshape mode: {self.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( - [bucket for bucket in self.frame_buckets if bucket <= video_num_frames], - key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)), - default=1, - ) - - 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() - - nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3]) - frames_resized = self._resize_for_rectangle_crop(frames, nearest_res) - frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0) - - image = frames[:1].clone() if self.image_to_video else None - - return image, frames, None - - 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] \ No newline at end of file + return transforms.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1]) \ No newline at end of file diff --git a/training/mochi-1/prepare_dataset.py b/training/mochi-1/prepare_dataset.py index 0826c27..864e382 100644 --- a/training/mochi-1/prepare_dataset.py +++ b/training/mochi-1/prepare_dataset.py @@ -21,7 +21,7 @@ from torch.utils.data import DataLoader from torchvision import transforms from tqdm import tqdm from transformers import T5EncoderModel, T5Tokenizer -from dataset_mochi import VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop +from dataset_mochi import VideoDatasetWithFlexibleResize import decord # isort:skip @@ -326,6 +326,8 @@ def serialize_artifacts( 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) @@ -421,6 +423,7 @@ def main(): 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, @@ -432,13 +435,7 @@ def main(): "random_flip": args.random_flip, "image_to_video": args.save_image_latents, } - if args.video_reshape_mode is None: - dataset = VideoDatasetWithResizing(**dataset_init_kwargs) - else: - dataset = VideoDatasetWithResizeAndRectangleCrop( - video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs - ) - + dataset = VideoDatasetWithFlexibleResize(**dataset_init_kwargs) original_dataset_size = len(dataset) # Split data among GPUs diff --git a/training/mochi-1/prepare_dataset.sh b/training/mochi-1/prepare_dataset.sh index 02786e5..c2aba5c 100644 --- a/training/mochi-1/prepare_dataset.sh +++ b/training/mochi-1/prepare_dataset.sh @@ -11,8 +11,8 @@ 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="1 84" -MAX_NUM_FRAMES="84" +FRAME_BUCKETS="85" +MAX_NUM_FRAMES="85" MAX_SEQUENCE_LENGTH=256 TARGET_FPS=30 BATCH_SIZE=4 diff --git a/training/mochi-1/text_to_video_lora.py b/training/mochi-1/text_to_video_lora.py index 4ffe908..26091e2 100644 --- a/training/mochi-1/text_to_video_lora.py +++ b/training/mochi-1/text_to_video_lora.py @@ -54,7 +54,7 @@ from tqdm.auto import tqdm from transformers import AutoTokenizer, T5EncoderModel from args import get_args # isort:skip -from dataset_mochi import VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip +from dataset_mochi import VideoDatasetWithFlexibleResize # isort:skip import sys sys.path.append("..") @@ -309,7 +309,6 @@ def main(args): variant=args.variant, ) scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") - # noise_scheduler_copy = FlowMatchEulerDiscreteScheduler.from_config(scheduler.config, invert_sigmas=False) noise_scheduler_copy = copy.deepcopy(scheduler) vae_config = AutoencoderKLMochi.load_config(args.pretrained_model_name_or_path, subfolder="vae") @@ -347,7 +346,8 @@ def main(args): ) transformer.requires_grad_(False) - transformer.to(accelerator.device, dtype=weight_dtype) + # transformer.to(accelerator.device, dtype=weight_dtype) + transformer.to(accelerator.device) if args.gradient_checkpointing: transformer.enable_gradient_checkpointing() @@ -423,9 +423,8 @@ def main(args): # Make sure the trainable params are in float32. This is again needed since the base models # are in `weight_dtype`. More details: # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 - if args.mixed_precision == "fp16": - # only upcast trainable parameters (LoRA) into fp32 - cast_training_params([transformer_]) + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer_]) accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) @@ -439,11 +438,8 @@ def main(args): args.learning_rate = ( args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes ) - - # Make sure the trainable params are in float32. - if args.mixed_precision == "fp16": - # only upcast trainable parameters (LoRA) into fp32 - cast_training_params([transformer], dtype=torch.float32) + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer], dtype=torch.float32) transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters())) @@ -488,6 +484,7 @@ def main(args): # Dataset and DataLoader 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, @@ -499,15 +496,8 @@ def main(args): "load_tensors": args.load_tensors, "random_flip": args.random_flip, } - if args.video_reshape_mode is None: - train_dataset = VideoDatasetWithResizing(**dataset_init_kwargs) - else: - train_dataset = VideoDatasetWithResizeAndRectangleCrop( - video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs - ) - + train_dataset = VideoDatasetWithFlexibleResize(**dataset_init_kwargs) collate_fn = CollateFunction(weight_dtype, args.load_tensors) - train_dataloader = DataLoader( train_dataset, batch_size=1, @@ -635,7 +625,6 @@ def main(args): 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) - # notice the reverse. step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] sigma = sigmas[step_indices].flatten() @@ -944,6 +933,7 @@ def main(args): commit_message="End of training", ignore_patterns=["step_*", "epoch_*", "*.bin", "*.pt"], ) + accelerator.print(f"Params pushed to {repo_id}.") accelerator.end_training() diff --git a/training/mochi-1/train.sh b/training/mochi-1/train.sh index e9851cb..69c636b 100644 --- a/training/mochi-1/train.sh +++ b/training/mochi-1/train.sh @@ -16,21 +16,17 @@ cmd="accelerate launch --config_file deepspeed.yaml --gpu_ids $GPU_IDS text_to_v --id_token BW_STYLE \ --height_buckets 480 \ --width_buckets 848 \ - --frame_buckets 84 \ + --frame_buckets 85 \ --load_tensors \ - --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 \ --seed 42 \ --rank 64 \ --lora_alpha 64 \ --mixed_precision bf16 \ --output_dir /raid/.cache/huggingface/sayak/mochi-lora/ \ - --max_num_frames 84 \ + --max_num_frames 85 \ --train_batch_size 1 \ --dataloader_num_workers 4 \ - --max_train_steps 500 \ + --max_train_steps 10 \ --checkpointing_steps 50 \ --gradient_accumulation_steps 4 \ --gradient_checkpointing \