This commit is contained in:
sayakpaul
2024-11-22 23:34:17 +05:30
parent 9a15eaee41
commit 2dbddd566e
5 changed files with 87 additions and 103 deletions
+67 -66
View File
@@ -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]
return transforms.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1])
+5 -8
View File
@@ -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
+2 -2
View File
@@ -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
+10 -20
View File
@@ -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()
+3 -7
View File
@@ -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 \