This commit is contained in:
sayakpaul
2024-11-18 15:26:52 +05:30
parent 5ba510e061
commit 9852c3d994
6 changed files with 1976 additions and 26 deletions
+474
View File
@@ -0,0 +1,474 @@
import argparse
def _get_model_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
default=None,
required=True,
help="Path to pretrained model or model identifier from huggingface.co/models.",
)
parser.add_argument(
"--revision",
type=str,
default=None,
required=False,
help="Revision of pretrained model identifier from huggingface.co/models.",
)
parser.add_argument(
"--variant",
type=str,
default=None,
help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16",
)
parser.add_argument(
"--cache_dir",
type=str,
default=None,
help="The directory where the downloaded models and datasets will be stored.",
)
def _get_dataset_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--data_root",
type=str,
default=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",
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.",
)
def _get_validation_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--validation_prompt",
type=str,
default=None,
help="One or more prompt(s) that is used during validation to verify that the model is learning. Multiple validation prompts should be separated by the '--validation_prompt_seperator' string.",
)
parser.add_argument(
"--validation_images",
type=str,
default=None,
help="One or more image path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.",
)
parser.add_argument(
"--validation_prompt_separator",
type=str,
default=":::",
help="String that separates multiple validation prompts",
)
parser.add_argument(
"--num_validation_videos",
type=int,
default=1,
help="Number of videos that should be generated during validation per `validation_prompt`.",
)
parser.add_argument(
"--validation_epochs",
type=int,
default=50,
help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.",
)
parser.add_argument(
"--enable_model_cpu_offload",
action="store_true",
default=False,
help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.",
)
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(
"--lora_alpha",
type=int,
default=64,
help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.",
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >= 1.10.and an Nvidia Ampere GPU. "
"Default to the value of accelerate config of the current system or the flag passed with the `accelerate.launch` command. Use this "
"argument to override the accelerate config."
),
)
parser.add_argument(
"--output_dir",
type=str,
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,
default=4,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_train_epochs", type=int, default=1)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help="Total number of training steps to perform. If provided, overrides `--num_train_epochs`.",
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=(
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."
),
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
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",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-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",
help=(
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'
),
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=500,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of hard resets of the lr in cosine_with_restarts scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--enable_slicing",
action="store_true",
default=False,
help="Whether or not to use VAE slicing for saving memory.",
)
parser.add_argument(
"--enable_tiling",
action="store_true",
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:
parser.add_argument(
"--optimizer",
type=lambda s: s.lower(),
default="adam",
choices=["adam", "adamw", "prodigy", "came"],
help=("The optimizer type to use."),
)
parser.add_argument(
"--use_8bit",
action="store_true",
help="Whether or not to use 8-bit optimizers from `bitsandbytes` or `bitsandbytes`.",
)
parser.add_argument(
"--use_4bit",
action="store_true",
help="Whether or not to use 4-bit optimizers from `torchao`.",
)
parser.add_argument(
"--use_torchao", action="store_true", help="Whether or not to use the `torchao` backend for optimizers."
)
parser.add_argument(
"--beta1",
type=float,
default=0.9,
help="The beta1 parameter for the Adam and Prodigy optimizers.",
)
parser.add_argument(
"--beta2",
type=float,
default=0.95,
help="The beta2 parameter for the Adam and Prodigy optimizers.",
)
parser.add_argument(
"--beta3",
type=float,
default=None,
help="Coefficients for computing the Prodigy optimizer's stepsize using running averages. If set to None, uses the value of square root of beta2.",
)
parser.add_argument(
"--prodigy_decouple",
action="store_true",
help="Use AdamW style decoupled weight decay.",
)
parser.add_argument(
"--weight_decay",
type=float,
default=1e-04,
help="Weight decay to use for optimizer.",
)
parser.add_argument(
"--epsilon",
type=float,
default=1e-8,
help="Epsilon value for the Adam optimizer and Prodigy optimizers.",
)
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
parser.add_argument(
"--prodigy_use_bias_correction",
action="store_true",
help="Turn on Adam's bias correction.",
)
parser.add_argument(
"--prodigy_safeguard_warmup",
action="store_true",
help="Remove lr from the denominator of D estimate to avoid issues during warm-up stage.",
)
parser.add_argument(
"--use_cpu_offload_optimizer",
action="store_true",
help="Whether or not to use the CPUOffloadOptimizer from TorchAO to perform optimization step and maintain parameters on the CPU.",
)
parser.add_argument(
"--offload_gradients",
action="store_true",
help="Whether or not to offload the gradients to CPU when using the CPUOffloadOptimizer from TorchAO.",
)
def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--tracker_name", type=str, default=None, help="Project tracker name")
parser.add_argument(
"--push_to_hub",
action="store_true",
help="Whether or not to push the model to the Hub.",
)
parser.add_argument(
"--hub_token",
type=str,
default=None,
help="The token to use to push to the Model Hub.",
)
parser.add_argument(
"--hub_model_id",
type=str,
default=None,
help="The name of the repository to keep in sync with the local `output_dir`.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help="Directory where logs are stored.",
)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--nccl_timeout",
type=int,
default=600,
help="Maximum timeout duration before which allgather, or related, operations fail in multi-GPU/multi-node training settings.",
)
parser.add_argument(
"--report_to",
type=str,
default=None,
help=(
'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
),
)
def get_args():
parser = argparse.ArgumentParser(description="Simple example of a training script for Mochi-1.")
_get_model_args(parser)
_get_dataset_args(parser)
_get_training_args(parser)
_get_validation_args(parser)
_get_optimizer_args(parser)
_get_configuration_args(parser)
return parser.parse_args()
+444
View File
@@ -0,0 +1,444 @@
import random
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import pandas as pd
import torch
import torchvision.transforms as TT
from accelerate.logging import get_logger
from torch.utils.data import Dataset, Sampler
from torchvision import transforms
from torchvision.transforms import InterpolationMode
from torchvision.transforms.functional import resize
# 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")
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]
VAE_SPATIAL_SCALE_FACTOR = 8
VAE_TEMPORAL_SCALE_FACTOR = 6
class VideoDataset(Dataset):
def __init__(
self,
data_root: str,
dataset_file: Optional[str] = None,
caption_column: str = "text",
video_column: str = "video",
max_num_frames: int = 49,
id_token: Optional[str] = None,
height_buckets: List[int] = None,
width_buckets: List[int] = None,
frame_buckets: List[int] = None,
load_tensors: bool = False,
random_flip: Optional[float] = None,
image_to_video: bool = False,
) -> None:
super().__init__()
self.data_root = Path(data_root)
self.dataset_file = dataset_file
self.caption_column = caption_column
self.video_column = video_column
self.max_num_frames = max_num_frames
self.id_token = id_token or ""
self.height_buckets = height_buckets or HEIGHT_BUCKETS
self.width_buckets = width_buckets or WIDTH_BUCKETS
self.frame_buckets = frame_buckets or FRAME_BUCKETS
self.load_tensors = load_tensors
self.random_flip = random_flip
self.image_to_video = image_to_video
self.resolutions = [
(f, h, w) for h in self.height_buckets for w in self.width_buckets for f in self.frame_buckets
]
# Two methods of loading data are supported.
# - Using a CSV: caption_column and video_column must be some column in the CSV. One could
# make use of other columns too, such as a motion score or aesthetic score, by modifying the
# logic in CSV processing.
# - Using two files containing line-separate captions and relative paths to videos.
# For a more detailed explanation about preparing dataset format, checkout the README.
if dataset_file is None:
(
self.prompts,
self.video_paths,
) = self._load_dataset_from_local_path()
else:
(
self.prompts,
self.video_paths,
) = self._load_dataset_from_csv()
if len(self.video_paths) != len(self.prompts):
raise ValueError(
f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset."
)
self.video_transforms = transforms.Compose(
[
transforms.RandomHorizontalFlip(random_flip)
if random_flip
else transforms.Lambda(self.identity_transform),
transforms.Lambda(self.scale_transform),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
)
@staticmethod
def identity_transform(x):
return x
@staticmethod
def scale_transform(x):
return x / 255.0
def __len__(self) -> int:
return len(self.video_paths)
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 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
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],
},
}
def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]:
if not self.data_root.exists():
raise ValueError("Root folder for videos does not exist")
prompt_path = self.data_root.joinpath(self.caption_column)
video_path = self.data_root.joinpath(self.video_column)
if not prompt_path.exists() or not prompt_path.is_file():
raise ValueError(
"Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts."
)
if not video_path.exists() or not video_path.is_file():
raise ValueError(
"Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory."
)
with open(prompt_path, "r", encoding="utf-8") as file:
prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0]
with open(video_path, "r", encoding="utf-8") as file:
video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0]
if not self.load_tensors and any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file."
)
return prompts, video_paths
def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]:
df = pd.read_csv(self.dataset_file)
prompts = df[self.caption_column].tolist()
video_paths = df[self.video_column].tolist()
video_paths = [self.data_root.joinpath(line.strip()) for line in video_paths]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file."
)
return prompts, video_paths
def _preprocess_video(self, path: Path) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
r"""
Loads a single video, or latent and prompt embedding, based on initialization parameters.
If returning a video, returns a [F, C, H, W] video tensor, and None for the prompt embedding. Here,
F, C, H and W are the frames, channels, height and width of the input video.
If returning latent/embedding, returns a [F, C, H, W] latent, and the prompt embedding of shape [S, D].
F, C, H and W are the frames, channels, height and width of the latent, and S, D are the sequence length
and embedding dimension of prompt embeddings.
"""
if self.load_tensors:
return self._load_preprocessed_latents_and_embeds(path)
else:
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
indices = list(range(0, video_num_frames, video_num_frames // self.max_num_frames))
frames = video_reader.get_batch(indices)
frames = frames[: self.max_num_frames].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0)
image = frames[:1].clone() if self.image_to_video else None
return image, frames, None
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 VideoDatasetWithResizing(VideoDataset):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
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))
)
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)
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 = 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)
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]
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
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,
)
h, w = arr.shape[2], arr.shape[3]
arr = arr.squeeze(0)
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":
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
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))
)
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))
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]
class BucketSampler(Sampler):
r"""
PyTorch Sampler that groups 3D data by height, width and frames.
Args:
data_source (`VideoDataset`):
A PyTorch dataset object that is an instance of `VideoDataset`.
batch_size (`int`, defaults to `8`):
The batch size to use for training.
shuffle (`bool`, defaults to `True`):
Whether or not to shuffle the data in each batch before dispatching to dataloader.
drop_last (`bool`, defaults to `False`):
Whether or not to drop incomplete buckets of data after completely iterating over all data
in the dataset. If set to True, only batches that have `batch_size` number of entries will
be yielded. If set to False, it is guaranteed that all data in the dataset will be processed
and batches that do not have `batch_size` number of entries will also be yielded.
"""
def __init__(
self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False
) -> None:
self.data_source = data_source
self.batch_size = batch_size
self.shuffle = shuffle
self.drop_last = drop_last
self.buckets = {resolution: [] for resolution in data_source.resolutions}
self._raised_warning_for_drop_last = False
def __len__(self):
if self.drop_last and not self._raised_warning_for_drop_last:
self._raised_warning_for_drop_last = True
logger.warning(
"Calculating the length for bucket sampler is not possible when `drop_last` is set to True. This may cause problems when setting the number of epochs used for training."
)
return (len(self.data_source) + self.batch_size - 1) // self.batch_size
def __iter__(self):
for index, data in enumerate(self.data_source):
if data is not None:
video_metadata = data["video_metadata"]
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
self.buckets[(f, h, w)].append(data)
if len(self.buckets[(f, h, w)]) == self.batch_size:
if self.shuffle:
random.shuffle(self.buckets[(f, h, w)])
yield self.buckets[(f, h, w)]
del self.buckets[(f, h, w)]
self.buckets[(f, h, w)] = []
if self.drop_last:
return
for fhw, bucket in list(self.buckets.items()):
if len(bucket) == 0:
continue
if self.shuffle:
random.shuffle(bucket)
yield bucket
del self.buckets[fhw]
self.buckets[fhw] = []
+23
View File
@@ -0,0 +1,23 @@
compute_environment: LOCAL_MACHINE
debug: false
deepspeed_config:
gradient_accumulation_steps: 1
gradient_clipping: 1.0
offload_optimizer_device: cpu
offload_param_device: cpu
zero3_init_flag: false
zero_stage: 2
distributed_type: DEEPSPEED
downcast_bf16: 'no'
enable_cpu_affinity: false
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 1
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
+31 -26
View File
@@ -8,6 +8,7 @@ import pathlib
import queue
import traceback
import uuid
from contextlib import nullcontext
from concurrent.futures import ThreadPoolExecutor
from typing import Any, Dict, List, Optional, Union
@@ -25,7 +26,7 @@ from transformers import T5EncoderModel, T5Tokenizer
import decord # isort:skip
import sys
sys.path.append("..")
sys.path.append(".")
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
@@ -107,13 +108,13 @@ def get_args() -> Dict[str, Any]:
"--width_buckets",
nargs="+",
type=check_width,
default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536],
default=[256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536],
)
parser.add_argument(
"--frame_buckets",
nargs="+",
type=check_frames,
default=[49],
default=[84],
)
parser.add_argument(
"--random_flip",
@@ -149,11 +150,11 @@ def get_args() -> Dict[str, Any]:
required=True,
help="Path to output directory where preprocessed videos/latents/embeddings will be saved.",
)
parser.add_argument("--max_num_frames", type=int, default=49, help="Maximum number of frames in output video.")
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=226, help="Max sequence length of prompt embeddings."
"--max_sequence_length", type=int, default=256, help="Max sequence length of prompt embeddings."
)
parser.add_argument("--target_fps", type=int, default=8, help="Frame rate of output videos.")
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",
@@ -213,6 +214,8 @@ def _get_t5_prompt_embeds(
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.")
@@ -220,12 +223,13 @@ def _get_t5_prompt_embeds(
prompt_embeds = text_encoder(text_input_ids.to(device))[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, 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
return prompt_embeds, prompt_attention_mask
def encode_prompt(
@@ -233,13 +237,13 @@ def encode_prompt(
text_encoder: T5EncoderModel,
prompt: Union[str, List[str]],
num_videos_per_prompt: int = 1,
max_sequence_length: int = 226,
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 = _get_t5_prompt_embeds(
prompt_embeds, prompt_attention_mask = _get_t5_prompt_embeds(
tokenizer,
text_encoder,
prompt=prompt,
@@ -249,7 +253,7 @@ def encode_prompt(
dtype=dtype,
text_input_ids=text_input_ids,
)
return prompt_embeds
return prompt_embeds, prompt_attention_mask
def compute_prompt_embeddings(
@@ -261,8 +265,9 @@ def compute_prompt_embeddings(
dtype: torch.dtype,
requires_grad: bool = False,
):
if requires_grad:
prompt_embeds = encode_prompt(
ctx = nullcontext() if requires_grad else torch.no_grad()
with ctx:
prompt_embeds, prompt_attention_mask = encode_prompt(
tokenizer,
text_encoder,
prompts,
@@ -271,18 +276,7 @@ def compute_prompt_embeddings(
device=device,
dtype=dtype,
)
else:
with torch.no_grad():
prompt_embeds = encode_prompt(
tokenizer,
text_encoder,
prompts,
num_videos_per_prompt=1,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return prompt_embeds
return prompt_embeds, prompt_attention_mask
to_pil_image = transforms.ToPILImage(mode="RGB")
@@ -318,12 +312,14 @@ def serialize_artifacts(
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:
num_frames, height, width = videos.size(1), videos.size(3), videos.size(4)
metadata = [{"num_frames": num_frames, "height": height, "width": width}]
@@ -335,6 +331,7 @@ def serialize_artifacts(
(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)]
@@ -398,6 +395,7 @@ def main():
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)
@@ -405,6 +403,7 @@ def main():
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
@@ -543,9 +542,10 @@ 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 = compute_prompt_embeddings(
prompt_embeds, prompt_attention_mask = compute_prompt_embeddings(
tokenizer,
text_encoder,
prompts,
@@ -554,6 +554,7 @@ 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
@@ -570,12 +571,14 @@ def main():
"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,
}
)
@@ -608,6 +611,7 @@ def main():
("video_latents", "pt"),
("prompts", "txt"),
("prompt_embeds", "pt"),
("prompt_attention_mask", "pt"),
("videos", "txt"),
]:
tmp_subfolder = tmp_dir.joinpath(subfolder)
@@ -660,6 +664,7 @@ def main():
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",
+948
View File
@@ -0,0 +1,948 @@
# Copyright 2024 The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import gc
import logging
import math
import os
import shutil
from datetime import timedelta
from pathlib import Path
from typing import Any, Dict
import diffusers
import torch
import transformers
import wandb
import copy
from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
DistributedDataParallelKwargs,
InitProcessGroupKwargs,
ProjectConfiguration,
set_seed,
)
from diffusers import (
AutoencoderKLMochi,
FlowMatchEulerDiscreteScheduler,
MochiPipeline,
MochiTransformer3DModel,
)
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.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
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
import sys
sys.path.append("..")
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
from text_encoder import compute_prompt_embeddings # isort:skip
from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip
logger = get_logger(__name__)
def save_model_card(
repo_id: str,
videos=None,
base_model: str = None,
validation_prompt=None,
repo_folder=None,
fps=8,
):
widget_dict = []
if videos is not None:
for i, video in enumerate(videos):
export_to_video(video, os.path.join(repo_folder, f"final_video_{i}.mp4", fps=fps))
widget_dict.append(
{
"text": validation_prompt if validation_prompt else " ",
"output": {"url": f"video_{i}.mp4"},
}
)
model_description = f"""
# Mochi-1 Preview LoRA Finetune
<Gallery />
## Model description
This is a lora finetune of the Moch-1 preview model `{base_model}`.
The model was trained using [CogVideoX Factory](https://github.com/a-r-r-o-w/cogvideox-factory) - a repository containing memory-optimized training scripts for the CogVideoX, Mochi family of models using [TorchAO](https://github.com/pytorch/ao) and [DeepSpeed](https://github.com/microsoft/DeepSpeed). The scripts were adopted from [CogVideoX Diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/cogvideo/train_cogvideox_lora.py).
## Download model
[Download LoRA]({repo_id}/tree/main) in the Files & Versions tab.
## Usage
Requires the [🧨 Diffusers library](https://github.com/huggingface/diffusers) installed.
```py
TODO
```
For more details, including weighting, merging and fusing LoRAs, check the [documentation](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) on loading LoRAs in diffusers.
"""
model_card = load_or_create_model_card(
repo_id_or_path=repo_id,
from_training=True,
license="apache-2.0",
base_model=base_model,
prompt=validation_prompt,
model_description=model_description,
widget=widget_dict,
)
tags = [
"text-to-video",
"diffusers-training",
"diffusers",
"lora",
"mochi-1-preview",
"mochi-1-preview-diffusers",
"template:sd-lora",
]
model_card = populate_model_card(model_card, tags=tags)
model_card.save(os.path.join(repo_folder, "README.md"))
def log_validation(
accelerator: Accelerator,
pipe: MochiPipeline,
args: Dict[str, Any],
pipeline_args: Dict[str, Any],
epoch,
is_final_validation: bool = False,
):
logger.info(
f"Running validation... \n Generating {args.num_validation_videos} videos with prompt: {pipeline_args['prompt']}."
)
pipe = pipe.to(accelerator.device)
# run inference
generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed else None
videos = []
with torch.autocast(accelerator.device.type, torch.bfloat16, cache_enabled=False):
for _ in range(args.num_validation_videos):
video = pipe(**pipeline_args, generator=generator, output_type="np").frames[0]
videos.append(video)
for tracker in accelerator.trackers:
phase_name = "test" if is_final_validation else "validation"
if tracker.name == "wandb":
video_filenames = []
for i, video in enumerate(videos):
prompt = (
pipeline_args["prompt"][:25]
.replace(" ", "_")
.replace(" ", "_")
.replace("'", "_")
.replace('"', "_")
.replace("/", "_")
)
filename = os.path.join(args.output_dir, f"{phase_name}_video_{i}_{prompt}.mp4")
export_to_video(video, filename, fps=30)
video_filenames.append(filename)
tracker.log(
{
phase_name: [
wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}", fps=30)
for i, filename in enumerate(video_filenames)
]
}
)
return videos
class CollateFunction:
def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None:
self.weight_dtype = weight_dtype
self.load_tensors = load_tensors
def __call__(self, data: Dict[str, Any]) -> Dict[str, torch.Tensor]:
prompts = [x["prompt"] for x in data[0]]
prompt_attention_mask = None
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]])
videos = [x["video"] for x in data[0]]
videos = torch.stack(videos).to(dtype=self.weight_dtype, non_blocking=True)
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
def main(args):
if args.report_to == "wandb" and args.hub_token is not None:
raise ValueError(
"You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token."
" Please use `huggingface-cli login` to authenticate with the Hub."
)
if torch.backends.mps.is_available() and args.mixed_precision == "bf16":
# due to pytorch#99272, MPS does not yet support bfloat16.
raise ValueError(
"Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead."
)
logging_dir = Path(args.output_dir, args.logging_dir)
accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir)
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
init_process_group_kwargs = InitProcessGroupKwargs(backend="nccl", timeout=timedelta(seconds=args.nccl_timeout))
accelerator = Accelerator(
gradient_accumulation_steps=args.gradient_accumulation_steps,
mixed_precision=args.mixed_precision,
log_with=args.report_to,
project_config=accelerator_project_config,
kwargs_handlers=[ddp_kwargs, init_process_group_kwargs],
)
# Disable AMP for MPS.
if torch.backends.mps.is_available():
accelerator.native_amp = False
# Make one log on every process with the configuration for debugging.
logging.basicConfig(
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
datefmt="%m/%d/%Y %H:%M:%S",
level=logging.INFO,
)
logger.info(accelerator.state, main_process_only=False)
if accelerator.is_local_main_process:
transformers.utils.logging.set_verbosity_warning()
diffusers.utils.logging.set_verbosity_info()
else:
transformers.utils.logging.set_verbosity_error()
diffusers.utils.logging.set_verbosity_error()
# If passed along, set the training seed now.
if args.seed is not None:
set_seed(args.seed)
# Handle the repository creation
if accelerator.is_main_process:
if args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
if args.push_to_hub:
repo_id = create_repo(
repo_id=args.hub_model_id or Path(args.output_dir).name,
exist_ok=True,
).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()
text_encoder.requires_grad_(False)
text_encoder.to(accelerator.device, dtype=weight_dtype)
vae.requires_grad_(False)
vae.to(accelerator.device, dtype=weight_dtype)
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)
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
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 torch.backends.mps.is_available() and weight_dtype == torch.bfloat16:
# due to pytorch#99272, MPS does not yet support bfloat16.
raise ValueError(
"Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead."
)
transformer.requires_grad_(False)
transformer.to(accelerator.device, dtype=weight_dtype)
if args.gradient_checkpointing:
transformer.enable_gradient_checkpointing()
# now we will add new LoRA weights to the attention layers
transformer_lora_config = LoraConfig(
r=args.rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
def unwrap_model(model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
transformer_lora_layers_to_save = None
for model in models:
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
model = unwrap_model(model)
transformer_lora_layers_to_save = get_peft_model_state_dict(model)
else:
raise ValueError(f"unexpected save model: {model.__class__}")
# make sure to pop weight so that corresponding model is not saved again
if weights:
weights.pop()
MochiPipeline.save_lora_weights(
output_dir,
transformer_lora_layers=transformer_lora_layers_to_save,
)
def load_model_hook(models, input_dir):
transformer_ = None
# This is a bit of a hack but I don't know any other solution.
if not accelerator.distributed_type == DistributedType.DEEPSPEED:
while len(models) > 0:
model = models.pop()
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
transformer_ = unwrap_model(model)
else:
raise ValueError(f"Unexpected save model: {unwrap_model(model).__class__}")
else:
transformer_ = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path, subfolder="transformer"
)
transformer_.add_adapter(transformer_lora_config)
lora_state_dict = MochiPipeline.lora_state_dict(input_dir)
transformer_state_dict = {
f'{k.replace("transformer.", "")}': v for k, v in lora_state_dict.items() if k.startswith("transformer.")
}
transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict)
incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default")
if incompatible_keys is not None:
# check only for unexpected keys
unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
if unexpected_keys:
logger.warning(
f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. "
)
# 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_])
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
# Enable TF32 for faster training on Ampere GPUs,
# cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
if args.allow_tf32 and torch.cuda.is_available():
torch.backends.cuda.matmul.allow_tf32 = True
if args.scale_lr:
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)
transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters()))
# Optimization parameters
transformer_parameters_with_lr = {
"params": transformer_lora_parameters,
"lr": args.learning_rate,
}
params_to_optimize = [transformer_parameters_with_lr]
num_trainable_parameters = sum(param.numel() for model in params_to_optimize for param in model["params"])
use_deepspeed_optimizer = (
accelerator.state.deepspeed_plugin is not None
and "optimizer" in accelerator.state.deepspeed_plugin.deepspeed_config
)
use_deepspeed_scheduler = (
accelerator.state.deepspeed_plugin is not None
and "scheduler" in accelerator.state.deepspeed_plugin.deepspeed_config
)
optimizer = get_optimizer(
params_to_optimize=params_to_optimize,
optimizer_name=args.optimizer,
learning_rate=args.learning_rate,
beta1=args.beta1,
beta2=args.beta2,
beta3=args.beta3,
epsilon=args.epsilon,
weight_decay=args.weight_decay,
prodigy_decouple=args.prodigy_decouple,
prodigy_use_bias_correction=args.prodigy_use_bias_correction,
prodigy_safeguard_warmup=args.prodigy_safeguard_warmup,
use_8bit=args.use_8bit,
use_4bit=args.use_4bit,
use_torchao=args.use_torchao,
use_deepspeed=use_deepspeed_optimizer,
use_cpu_offload_optimizer=args.use_cpu_offload_optimizer,
offload_gradients=args.offload_gradients,
)
# Dataset and DataLoader
dataset_init_kwargs = {
"data_root": args.data_root,
"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,
}
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
)
collate_fn = CollateFunction(weight_dtype, args.load_tensors)
train_dataloader = DataLoader(
train_dataset,
batch_size=1,
sampler=BucketSampler(train_dataset, batch_size=args.train_batch_size, shuffle=True),
collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
pin_memory=args.pin_memory,
prefetch_factor=4,
)
# Scheduler and math around the number of training steps.
overrode_max_train_steps = False
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if args.max_train_steps is None:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
overrode_max_train_steps = True
if args.use_cpu_offload_optimizer:
lr_scheduler = None
accelerator.print(
"CPU Offload Optimizer cannot be used with DeepSpeed or builtin PyTorch LR Schedulers. If "
"you are training with those settings, they will be ignored."
)
else:
if use_deepspeed_scheduler:
from accelerate.utils import DummyScheduler
lr_scheduler = DummyScheduler(
name=args.lr_scheduler,
optimizer=optimizer,
total_num_steps=args.max_train_steps * accelerator.num_processes,
num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
)
else:
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
num_training_steps=args.max_train_steps * accelerator.num_processes,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
)
# Prepare everything with our `accelerator`.
transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
transformer, optimizer, train_dataloader, lr_scheduler
)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
if overrode_max_train_steps:
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
# We need to initialize the trackers we use, and also store our configuration.
# The trackers initializes automatically on the main process.
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
tracker_name = args.tracker_name or "mochi-1-lora"
accelerator.init_trackers(tracker_name, config=vars(args))
accelerator.print("===== Memory before training =====")
reset_memory(accelerator.device)
print_memory(accelerator.device)
# Train!
total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
accelerator.print("***** Running training *****")
accelerator.print(f" Num trainable parameters = {num_trainable_parameters}")
accelerator.print(f" Num examples = {len(train_dataset)}")
accelerator.print(f" Num batches each epoch = {len(train_dataloader)}")
accelerator.print(f" Num epochs = {args.num_train_epochs}")
accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}")
accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
accelerator.print(f" Gradient accumulation steps = {args.gradient_accumulation_steps}")
accelerator.print(f" Total optimization steps = {args.max_train_steps}")
global_step = 0
first_epoch = 0
# Potentially load in the weights and states from a previous save
if not args.resume_from_checkpoint:
initial_global_step = 0
else:
if args.resume_from_checkpoint != "latest":
path = os.path.basename(args.resume_from_checkpoint)
else:
# Get the most recent checkpoint
dirs = os.listdir(args.output_dir)
dirs = [d for d in dirs if d.startswith("checkpoint")]
dirs = sorted(dirs, key=lambda x: int(x.split("-")[1]))
path = dirs[-1] if len(dirs) > 0 else None
if path is None:
accelerator.print(
f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run."
)
args.resume_from_checkpoint = None
initial_global_step = 0
else:
accelerator.print(f"Resuming from checkpoint {path}")
accelerator.load_state(os.path.join(args.output_dir, path))
global_step = int(path.split("-")[1])
initial_global_step = global_step
first_epoch = global_step // num_update_steps_per_epoch
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=initial_global_step,
desc="Steps",
# Only show the progress bar once on each machine.
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()
torch.cuda.synchronize(accelerator.device)
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()
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
for epoch in range(first_epoch, args.num_train_epochs):
transformer.train()
for step, batch in enumerate(train_dataloader):
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"]
# 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).latent_dist
else:
latent_dist = DiagonalGaussianDistribution(videos)
videos = latent_dist.sample()
videos = videos[:, :vae_in_channels, ...] # to respect `in_channels` for the vae
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
videos = videos.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
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,
requires_grad=False,
)
else:
prompt_embeds = prompts.to(dtype=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)
# print(f"{timesteps=}")
# Add noise according to flow matching.
# zt = (1 - texp) * x + texp * z1
sigmas = get_sigmas(timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
noisy_model_input = (1.0 - sigmas) * model_input + sigmas * noise
# Predict the noise residual
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
loss = torch.mean(
(weighting * (model_pred.float() - target.float()) ** 2).reshape(batch_size, -1),
dim=1,
)
loss = loss.mean()
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.state.deepspeed_plugin is None:
optimizer.step()
optimizer.zero_grad()
if not args.use_cpu_offload_optimizer:
lr_scheduler.step()
# Checks if the accelerator has performed an optimization step behind the scenes
if accelerator.sync_gradients:
progress_bar.update(1)
global_step += 1
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
if global_step % args.checkpointing_steps == 0:
# _before_ saving state, check if this save would set us over the `checkpoints_total_limit`
if args.checkpoints_total_limit is not None:
checkpoints = os.listdir(args.output_dir)
checkpoints = [d for d in checkpoints if d.startswith("checkpoint")]
checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))
# before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints
if len(checkpoints) >= args.checkpoints_total_limit:
num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1
removing_checkpoints = checkpoints[0:num_to_remove]
logger.info(
f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints"
)
logger.info(f"Removing checkpoints: {', '.join(removing_checkpoints)}")
for removing_checkpoint in removing_checkpoints:
removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint)
shutil.rmtree(removing_checkpoint)
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
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,
}
)
progress_bar.set_postfix(**logs)
accelerator.log(logs, step=global_step)
if global_step >= args.max_train_steps:
break
if 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)
torch.cuda.synchronize(accelerator.device)
pipe = MochiPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
# torch_dtype=weight_dtype,
)
if args.enable_slicing:
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
if args.enable_model_cpu_offload:
pipe.enable_model_cpu_offload()
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
pipeline_args = {
"prompt": validation_prompt,
"guidance_scale": 4.5,
"height": args.height,
"width": args.width,
"max_sequence_length": 256,
}
log_validation(
pipe=pipe,
args=args,
accelerator=accelerator,
pipeline_args=pipeline_args,
epoch=epoch,
)
accelerator.print("===== Memory after validation =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
del pipe
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(accelerator.device)
accelerator.wait_for_everyone()
if accelerator.is_main_process:
transformer = unwrap_model(transformer)
dtype = (
torch.float16
if args.mixed_precision == "fp16"
else torch.bfloat16
if args.mixed_precision == "bf16"
else torch.float32
)
# transformer = transformer.to(dtype)
transformer_lora_layers = get_peft_model_state_dict(transformer)
MochiPipeline.save_lora_weights(
save_directory=args.output_dir,
transformer_lora_layers=transformer_lora_layers,
)
# Cleanup trained models to save memory
if args.load_tensors:
del transformer
else:
del transformer, text_encoder, vae
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(accelerator.device)
accelerator.print("===== Memory before testing =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
# Final test inference
pipe = MochiPipeline.from_pretrained(
args.pretrained_model_name_or_path,
revision=args.revision,
variant=args.variant,
# torch_dtype=weight_dtype,
)
pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(pipe.scheduler.config)
if args.enable_slicing:
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
if args.enable_model_cpu_offload:
pipe.enable_model_cpu_offload()
# Load LoRA weights
lora_scaling = args.lora_alpha / args.rank
pipe.load_lora_weights(args.output_dir, adapter_name="mochi-lora")
pipe.set_adapters(["mochi-lora"], [lora_scaling])
# Run inference
validation_outputs = []
if args.validation_prompt and args.num_validation_videos > 0:
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
pipeline_args = {
"prompt": validation_prompt,
"guidance_scale": 4.5,
"use_dynamic_cfg": args.use_dynamic_cfg,
"height": args.height,
"width": args.width,
}
video = log_validation(
accelerator=accelerator,
pipe=pipe,
args=args,
pipeline_args=pipeline_args,
epoch=epoch,
is_final_validation=True,
)
validation_outputs.extend(video)
accelerator.print("===== Memory after testing =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
torch.cuda.synchronize(accelerator.device)
if args.push_to_hub:
save_model_card(
repo_id,
videos=validation_outputs,
base_model=args.pretrained_model_name_or_path,
validation_prompt=args.validation_prompt,
repo_folder=args.output_dir,
fps=args.fps,
)
upload_folder(
repo_id=repo_id,
folder_path=args.output_dir,
commit_message="End of training",
ignore_patterns=["step_*", "epoch_*"],
)
accelerator.end_training()
if __name__ == "__main__":
args = get_args()
main(args)
+56
View File
@@ -0,0 +1,56 @@
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"
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 \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
--video_column $VIDEO_COLUMN \
--id_token BW_STYLE \
--height_buckets 480 \
--width_buckets 848 \
--frame_buckets 84 \
--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 mochi-lora \
--max_num_frames 84 \
--train_batch_size 1 \
--dataloader_num_workers 4 \
--max_train_steps 500 \
--checkpointing_steps 50 \
--gradient_accumulation_steps 4 \
--gradient_checkpointing \
--learning_rate 0.0001 \
--lr_scheduler constant \
--lr_warmup_steps 0 \
--lr_num_cycles 1 \
--enable_slicing \
--enable_tiling \
--optimizer adamw \
--beta1 0.9 \
--beta2 0.95 \
--beta3 0.99 \
--weight_decay 0.001 \
--max_grad_norm 1.0 \
--allow_tf32 \
--report_to wandb \
--push_to_hub \
--nccl_timeout 1800"
echo "Running command: $cmd"
eval $cmd
echo -ne "-------------------- Finished executing script --------------------\n\n"