From 1409d478d53fbc41d015ba8d6b86e6bbbb355e53 Mon Sep 17 00:00:00 2001 From: sayakpaul Date: Tue, 19 Nov 2024 10:21:28 +0530 Subject: [PATCH] updates. --- training/mochi-1/args.py | 7 +++++++ training/mochi-1/dataset.py | 15 ++++++--------- training/mochi-1/prepare_dataset.py | 9 +++++---- training/mochi-1/prepare_dataset.sh | 5 +++-- training/mochi-1/text_to_video_lora.py | 2 +- 5 files changed, 22 insertions(+), 16 deletions(-) diff --git a/training/mochi-1/args.py b/training/mochi-1/args.py index 25248e4..ee7c463 100644 --- a/training/mochi-1/args.py +++ b/training/mochi-1/args.py @@ -151,6 +151,13 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None: default=64, help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.", ) + parser.add_argument( + "--target_modules", + nargs="+", + type=str, + default=["to_k", "to_q", "to_v", "to_out.0"], + help="Target modules to train LoRA for." + ) parser.add_argument( "--mixed_precision", type=str, diff --git a/training/mochi-1/dataset.py b/training/mochi-1/dataset.py index 09b9e2c..0a95ab7 100644 --- a/training/mochi-1/dataset.py +++ b/training/mochi-1/dataset.py @@ -128,9 +128,7 @@ class VideoDataset(Dataset): # 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. - # print(f"{video_latents.shape=}") latent_num_frames = video_latents.size(0) - # print(f"{latent_num_frames=}") num_frames = (latent_num_frames // 2) * (VAE_TEMPORAL_SCALE_FACTOR + 1) height = video_latents.size(2) * VAE_SPATIAL_SCALE_FACTOR @@ -288,11 +286,10 @@ class VideoDatasetWithResizing(VideoDataset): 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)) + [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, ) - if video_num_frames < nearest_frame_bucket: - # TODO: we could handle this by padding zero frames or duplicating the existing frames? - return None, None, None frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket)) frames = video_reader.get_batch(frame_indices) @@ -355,10 +352,10 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset): 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)) + [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, ) - if video_num_frames < nearest_frame_bucket: - return None, None, None frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket)) diff --git a/training/mochi-1/prepare_dataset.py b/training/mochi-1/prepare_dataset.py index 2f08b55..d9da99a 100644 --- a/training/mochi-1/prepare_dataset.py +++ b/training/mochi-1/prepare_dataset.py @@ -321,8 +321,11 @@ def serialize_artifacts( prompt_embeds: Optional[torch.Tensor] = None, prompt_attention_mask: Optional[torch.Tensor] = None ) -> None: - num_frames, height, width = videos.size(1), videos.size(3), videos.size(4) - metadata = [{"num_frames": num_frames, "height": height, "width": width}] + metadata = [] + for i in range(videos.size(0)): + video = videos[i:i+1] + 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"), @@ -542,7 +545,6 @@ def main(): video_latents = vae._encode(videos) video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype) - print(f"{video_latents.shape=}") # Encode prompts prompt_embeds, prompt_attention_mask = compute_prompt_embeddings( @@ -554,7 +556,6 @@ def main(): weight_dtype, requires_grad=False, ) - print(f"{prompt_attention_mask.shape=}") if images is not None: images = (images.permute(0, 2, 1, 3, 4) + 1) / 2 diff --git a/training/mochi-1/prepare_dataset.sh b/training/mochi-1/prepare_dataset.sh index 8081ac2..02786e5 100644 --- a/training/mochi-1/prepare_dataset.sh +++ b/training/mochi-1/prepare_dataset.sh @@ -11,11 +11,11 @@ VIDEO_COLUMN="videos.txt" OUTPUT_DIR="/home/sayak/cogvideox-factory/video-dataset-disney/mochi-1/preprocessed-dataset" HEIGHT_BUCKETS="480" WIDTH_BUCKETS="848" -FRAME_BUCKETS="84" +FRAME_BUCKETS="1 84" MAX_NUM_FRAMES="84" MAX_SEQUENCE_LENGTH=256 TARGET_FPS=30 -BATCH_SIZE=1 +BATCH_SIZE=4 DTYPE=fp32 # To create a folder-style dataset structure without pre-encoding videos and captions @@ -35,6 +35,7 @@ CMD_WITHOUT_PRE_ENCODING="\ --max_sequence_length $MAX_SEQUENCE_LENGTH \ --target_fps $TARGET_FPS \ --batch_size $BATCH_SIZE \ + --use_slicing \ --dtype $DTYPE " diff --git a/training/mochi-1/text_to_video_lora.py b/training/mochi-1/text_to_video_lora.py index fd48f19..9f42366 100644 --- a/training/mochi-1/text_to_video_lora.py +++ b/training/mochi-1/text_to_video_lora.py @@ -354,7 +354,7 @@ def main(args): r=args.rank, lora_alpha=args.lora_alpha, init_lora_weights=True, - target_modules=["to_k", "to_q", "to_v", "to_out.0"], + target_modules=args.target_modules, ) transformer.add_adapter(transformer_lora_config)