This commit is contained in:
sayakpaul
2024-11-29 10:48:54 +05:30
parent 2fde026d30
commit ced8558eb0
7 changed files with 347 additions and 698 deletions
+96
View File
@@ -0,0 +1,96 @@
# Simple Mochi-1 finetuner
Now you can make Mochi-1 your own with `diffusers`, too 🤗 🧨
We provide a minimal and faithful reimplementation of the [Mochi-1 original fine-tuner](https://github.com/genmoai/mochi/tree/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner). As usual, we leverage `peft` for things LoRA in our implementation.
## Getting started
Install the dependencies: `pip install -r requirements.txt`. Also make sure your `diffusers` installation is from the current `main`.
Download a demo dataset:
```bash
huggingface-cli download \
--repo-type dataset sayakpaul/video-dataset-disney-organized \
--local-dir video-dataset-disney-organized
```
The dataset follows the directory structure expected by the subsequent scripts. In particular, it follows what's prescribed [here](https://github.com/genmoai/mochi/tree/main/demos/fine_tuner#1-collect-your-videos-and-captions):
```bash
video_1.mp4
video_1.txt -- One-paragraph description of video_1
video_2.mp4
video_2.txt -- One-paragraph description of video_2
...
```
Then run (be sure to check the paths accordingly):
```bash
bash prepare_dataset.sh
```
We can adjust `num_frames` and `resolution`. By default, in `prepare_dataset.sh`, we use `--force_upsample`. This means if the original video resolution is smaller than the requested resolution, we will upsample the video.
> [!IMPORTANT]
> It's important to have a resolution of at least 480x848 to satisy Mochi-1's requirements.
Now, we're ready to fine-tune. To launch, run:
```bash
bash train.sh
```
You can disable intermediate validation by:
```diff
- --validation_prompt "..." \
- --validation_prompt_separator ::: \
- --num_validation_videos 1 \
- --validation_epochs 1 \
```
We haven't rigorously tested but without validation enabled, this script should run under 40GBs of GPU VRAM.
To use the LoRA checkpoint:
```py
from diffusers import MochiPipeline
from diffusers.utils import export_to_video
import torch
pipe = MochiPipeline.from_pretrained("genmo/mochi-1-preview")
pipe.load_lora_weights("path-to-lora")
pipe.enable_model_cpu_offload()
pipeline_args = {
"prompt": "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",
"guidance_scale": 6.0,
"num_inference_steps": 64,
"height": 480,
"width": 848,
"max_sequence_length": 256,
"output_type": "np",
}
with torch.autocast("cuda", torch.bfloat16)
video = pipe(**pipeline_args).frames[0]
export_to_video(video)
```
## Known limitations
(Contributions are welcome 🤗)
Our script currently doesn't leverage `accelerate` and some of its consequences are detailed below:
* No support for distributed training.
* No intermediate checkpoint saving and loading support.
* `train_batch_size > 1` are supported but can potentially lead to OOMs because we currently don't have gradient accumulation support.
**Misc**:
* We're aware of the quality issues in the `diffusers` implementation of Mochi-1. This is being fixed in [this PR](https://github.com/huggingface/diffusers/pull/10033).
* `embed.py` script is non-batched.
+25 -155
View File
@@ -1,3 +1,9 @@
"""
Default values taken from
https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner/configs/lora.yaml
when applicable.
"""
import argparse
@@ -33,6 +39,11 @@ def _get_model_args(parser: argparse.ArgumentParser) -> None:
action="store_true",
help="If we should cast DiT params to a lower precision.",
)
parser.add_argument(
"--compile_dit",
action="store_true",
help="If we should cast DiT params to a lower precision.",
)
def _get_dataset_args(parser: argparse.ArgumentParser) -> None:
@@ -93,6 +104,18 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None:
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_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(
"--enable_model_cpu_offload",
action="store_true",
@@ -133,17 +156,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=["to_k", "to_q", "to_v", "to_out.0"],
help="Target modules to train LoRA for.",
)
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,
@@ -163,37 +175,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
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(
"--gradient_checkpointing",
action="store_true",
@@ -210,45 +191,12 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
action="store_true",
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_scheduler",
type=str,
default="cosine",
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=200,
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.",
)
def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
@@ -256,78 +204,15 @@ def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
"--optimizer",
type=lambda s: s.lower(),
default="adam",
choices=["adam", "adamw", "prodigy", "came"],
choices=["adam", "adamw"],
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.999,
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=0.01,
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:
@@ -349,12 +234,6 @@ def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
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",
@@ -363,20 +242,11 @@ def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
" 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.'
),
help="If logging to wandb."
)
-23
View File
@@ -1,23 +0,0 @@
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
+4 -2
View File
@@ -1,9 +1,11 @@
#!/bin/bash
GPU_ID=0
VIDEO_DIR=/home/sayak/cogvideox-factory/video-dataset-disney-organized
VIDEO_DIR=video-dataset-disney-organized
OUTPUT_DIR=videos_prepared
NUM_FRAMES=37
RESOLUTION=480x848
python trim_and_crop_videos.py $VIDEO_DIR $OUTPUT_DIR --num_frames=37 --resolution=480x848 --force_upsample
python trim_and_crop_videos.py $VIDEO_DIR $OUTPUT_DIR --num_frames=$NUM_FRAMES --resolution=$RESOLUTION --force_upsample
CUDA_VISIBLE_DEVICES=$GPU_ID python embed.py $OUTPUT_DIR --shape=37x480x848
+7
View File
@@ -0,0 +1,7 @@
peft
transformers
wandb
torch
torchvision
moviepy
click
+207 -507
View File
@@ -16,41 +16,22 @@
import gc
import random
from glob import glob
import logging
import math
import os
import shutil
import torch.nn.functional as F
from datetime import timedelta
import numpy as np
from pathlib import Path
from typing import Any, Dict, Tuple, List
import diffusers
import torch
import transformers
import wandb
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 import FlowMatchEulerDiscreteScheduler, MochiPipeline, MochiTransformer3DModel
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.optimization import get_scheduler
from diffusers.training_utils import cast_training_params
from diffusers.utils import convert_unet_state_dict_to_peft, export_to_video
from diffusers.utils import 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 peft import LoraConfig, get_peft_model_state_dict
from torch.utils.data import DataLoader
from tqdm.auto import tqdm
@@ -63,10 +44,23 @@ import sys
sys.path.append("..")
from utils import get_optimizer, print_memory, reset_memory # isort:skip
from utils import print_memory, reset_memory # isort:skip
logger = get_logger(__name__)
# Taken from
# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner/train.py#L139
def get_cosine_annealing_lr_scheduler(
optimizer: torch.optim.Optimizer,
warmup_steps: int,
total_steps: int,
):
def lr_lambda(step):
if step < warmup_steps:
return float(step) / float(max(1, warmup_steps))
else:
return 0.5 * (1 + np.cos(np.pi * (step - warmup_steps) / (total_steps - warmup_steps)))
return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
def save_model_card(
@@ -84,7 +78,7 @@ def save_model_card(
widget_dict.append(
{
"text": validation_prompt if validation_prompt else " ",
"output": {"url": f"video_{i}.mp4"},
"output": {"url": f"final_video_{i}.mp4"},
}
)
@@ -138,54 +132,53 @@ For more details, including weighting, merging and fusing LoRAs, check the [docu
def log_validation(
accelerator: Accelerator,
pipe: MochiPipeline,
args: Dict[str, Any],
pipeline_args: Dict[str, Any],
epoch,
wandb_run: str = None,
is_final_validation: bool = False,
):
logger.info(
print(
f"Running validation... \n Generating {args.num_validation_videos} videos with prompt: {pipeline_args['prompt']}."
)
phase_name = "test" if is_final_validation else "validation"
if not args.enable_model_cpu_offload:
pipe = pipe.to(accelerator.device)
pipe = pipe.to("cuda")
# run inference
generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed else None
generator = torch.manual_seed(args.seed) if args.seed else None
videos = []
with torch.autocast(accelerator.device.type, torch.bfloat16, cache_enabled=False):
with torch.autocast("cuda", 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)
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)
]
}
)
if wandb_run:
wandb.log(
{
phase_name: [
wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}", fps=30)
for i, filename in enumerate(video_filenames)
]
}
)
return videos
@@ -231,63 +224,24 @@ class CollateFunction:
def main(args):
if not torch.cuda.is_available():
raise ValueError("Not supported without CUDA.")
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.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
if args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
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
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
transformer = MochiTransformer3DModel.from_pretrained(
@@ -300,44 +254,14 @@ def main(args):
args.pretrained_model_name_or_path, subfolder="scheduler"
)
vae_config = AutoencoderKLMochi.load_config(args.pretrained_model_name_or_path, subfolder="vae")
has_latents_mean = "latents_mean" in vae_config and vae_config["latents_mean"] is not None
has_latents_std = "latents_std" in vae_config and vae_config["latents_std"] is not None
if has_latents_mean and has_latents_std:
mean = torch.tensor(vae_config["latents_mean"])[:, None, None, None]
std = torch.tensor(vae_config["latents_mean"])[:, None, None, None]
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)
transformer.to("cuda")
if args.gradient_checkpointing:
transformer.enable_gradient_checkpointing()
if args.cast_dit:
transformer = cast_dit(transformer, weight_dtype)
transformer = cast_dit(transformer, torch.bfloat16)
if args.compile_dit:
transformer.compile()
# now we will add new LoRA weights to the attention layers
transformer_lora_config = LoraConfig(
@@ -348,131 +272,25 @@ def main(args):
)
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
# 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
)
args.learning_rate = args.learning_rate * args.train_batch_size
# only upcast trainable parameters (LoRA) into fp32
cast_training_params([transformer], dtype=torch.float32)
# Prepare optimizer
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,
)
accelerator.print(f"Using {optimizer.__class__.__name__} optimizer.")
num_trainable_parameters = sum(param.numel() for param in transformer_lora_parameters)
optimizer = torch.optim.AdamW(transformer_lora_parameters, lr=args.learning_rate, weight_decay=args.weight_decay)
# Dataset and DataLoader
train_vids = list(sorted(glob(f"{args.data_root}/*.mp4")))
train_vids = [v for v in train_vids if not v.endswith(".recon.mp4")]
accelerator.print(f"Found {len(train_vids)} training videos in {args.data_root}")
print(f"Found {len(train_vids)} training videos in {args.data_root}")
assert len(train_vids) > 0, f"No training data found in {args.data_root}"
collate_fn = CollateFunction(caption_dropout=args.caption_dropout)
@@ -485,46 +303,19 @@ def main(args):
pin_memory=args.pin_memory,
)
# Scheduler and math around the number of training steps.
# LR 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)
num_update_steps_per_epoch = len(train_dataloader)
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
lr_scheduler = get_cosine_annealing_lr_scheduler(
optimizer, warmup_steps=args.lr_warmup_steps, total_steps=args.max_train_steps
)
# 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)
num_update_steps_per_epoch = len(train_dataloader)
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
@@ -532,80 +323,43 @@ def main(args):
# 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:
wandb_run = None
if args.report_to == "wandb":
tracker_name = args.tracker_name or "mochi-1-lora"
accelerator.init_trackers(tracker_name, config=vars(args))
wandb_run = wandb.init(project=tracker_name, config=vars(args))
accelerator.print("===== Memory before training =====")
reset_memory(accelerator.device)
print_memory(accelerator.device)
print("===== Memory before training =====")
reset_memory("cuda")
print_memory("cuda")
# Train!
total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
total_batch_size = args.train_batch_size
print("***** Running training *****")
print(f" Num trainable parameters = {num_trainable_parameters}")
print(f" Num examples = {len(train_dataset)}")
print(f" Num batches each epoch = {len(train_dataloader)}")
print(f" Num epochs = {args.num_train_epochs}")
print(f" Instantaneous batch size per device = {args.train_batch_size}")
print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
print(f" Total optimization steps = {args.max_train_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,
initial=global_step,
desc="Steps",
# Only show the progress bar once on each machine.
disable=not accelerator.is_local_main_process,
)
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):
z = batch["z"]
# revisit
# if has_latents_mean and has_latents_std:
# z = (z - mean.to(z)) / std.to(z)
eps = batch["eps"]
sigma = batch["sigma"]
prompt_embeds = batch["prompt_embeds"]
prompt_attention_mask = batch["prompt_attention_mask"]
with torch.no_grad():
z = batch["z"].to("cuda")
eps = batch["eps"].to("cuda")
sigma = batch["sigma"].to("cuda")
prompt_embeds = batch["prompt_embeds"].to("cuda")
prompt_attention_mask = batch["prompt_attention_mask"].to("cuda")
sigma_bcthw = sigma[:, None, None, None, None] # [B, 1, 1, 1, 1]
# Add noise according to flow matching.
@@ -613,80 +367,35 @@ def main(args):
z_sigma = (1 - sigma_bcthw) * z + sigma_bcthw * eps
ut = z - eps
# Predict the noise residual
# (1 - sigma) because of
# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/src/genmo/mochi_preview/dit/joint_model/asymm_models_joint.py#L656
# Also, we operate on the scaled version of the `timesteps` directly in the `diffusers` implementation.
timesteps = (1 - sigma) * scheduler.config.num_train_timesteps
with torch.autocast(accelerator.device.type, weight_dtype):
model_pred = transformer(
hidden_states=z_sigma,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=timesteps,
return_dict=False,
)[0]
assert model_pred.shape == z.shape
loss = F.mse_loss(model_pred.float(), ut.float())
accelerator.backward(loss)
# if accelerator.sync_gradients:
# no grad norm for now, following the original code
# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner/train.py#L380
# gradient_norm_before_clip = get_gradient_norm(transformer_lora_parameters)
# accelerator.clip_grad_norm_(transformer_lora_parameters, args.max_grad_norm)
# gradient_norm_after_clip = get_gradient_norm(transformer_lora_parameters)
with torch.autocast("cuda", torch.bfloat16):
model_pred = transformer(
hidden_states=z_sigma,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=timesteps,
return_dict=False,
)[0]
assert model_pred.shape == z.shape
loss = F.mse_loss(model_pred.float(), ut.float())
loss.backward()
if accelerator.state.deepspeed_plugin is None:
optimizer.step()
optimizer.zero_grad()
optimizer.step()
optimizer.zero_grad()
lr_scheduler.step()
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}")
progress_bar.update(1)
global_step += 1
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 wandb_run:
wandb_run.log(logs, step=global_step)
if global_step >= args.max_train_steps:
break
@@ -694,82 +403,15 @@ def main(args):
if global_step >= args.max_train_steps:
break
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0:
accelerator.print("===== Memory before validation =====")
print_memory(accelerator.device)
transformer.eval()
pipe = MochiPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
)
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": 6.0,
"num_inference_steps": 64,
"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.text_encoder
del pipe.vae
del pipe
gc.collect()
torch.cuda.empty_cache()
transformer.train()
accelerator.wait_for_everyone()
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
transformer = unwrap_model(transformer)
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
del transformer
gc.collect()
torch.cuda.empty_cache()
# Final test inference
validation_outputs = []
if args.validation_prompt and args.num_validation_videos > 0:
accelerator.print("===== Memory before testing =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0:
print("===== Memory before validation =====")
print_memory("cuda")
transformer.eval()
pipe = MochiPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=transformer,
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
)
@@ -781,12 +423,6 @@ def main(args):
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_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
pipeline_args = {
@@ -797,40 +433,104 @@ def main(args):
"width": args.width,
"max_sequence_length": 256,
}
video = log_validation(
accelerator=accelerator,
log_validation(
pipe=pipe,
args=args,
pipeline_args=pipeline_args,
epoch=epoch,
is_final_validation=True,
wandb_run=wandb_run,
)
validation_outputs.extend(video)
accelerator.print("===== Memory after testing =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
torch.cuda.synchronize(accelerator.device)
print("===== Memory after validation =====")
print_memory("cuda")
reset_memory("cuda")
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,
del pipe.text_encoder
del pipe.vae
del pipe
gc.collect()
torch.cuda.empty_cache()
transformer.train()
transformer.eval()
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
del transformer
gc.collect()
torch.cuda.empty_cache()
# Final test inference
validation_outputs = []
if args.validation_prompt and args.num_validation_videos > 0:
print("===== Memory before testing =====")
print_memory("cuda")
reset_memory("cuda")
pipe = MochiPipeline.from_pretrained(
args.pretrained_model_name_or_path,
revision=args.revision,
variant=args.variant,
)
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_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
pipeline_args = {
"prompt": validation_prompt,
"guidance_scale": 6.0,
"num_inference_steps": 64,
"height": args.height,
"width": args.width,
"max_sequence_length": 256,
}
video = log_validation(
pipe=pipe,
args=args,
pipeline_args=pipeline_args,
epoch=epoch,
wandb_run=wandb_run,
is_final_validation=True,
)
upload_folder(
repo_id=repo_id,
folder_path=args.output_dir,
commit_message="End of training",
ignore_patterns=["step_*", "epoch_*", "*.bin", "*.pt"],
)
accelerator.print(f"Params pushed to {repo_id}.")
validation_outputs.extend(video)
accelerator.end_training()
print("===== Memory after testing =====")
print_memory("cuda")
reset_memory("cuda")
torch.cuda.synchronize("cuda")
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_*", "*.bin", "*.pt"],
)
print(f"Params pushed to {repo_id}.")
if __name__ == "__main__":
+8 -11
View File
@@ -2,38 +2,35 @@
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
GPU_IDS="2"
GPU_IDS="0"
DATA_ROOT="/home/sayak/cogvideox-factory/training/mochi-1/videos_prepared"
DATA_ROOT="videos_prepared"
MODEL="genmo/mochi-1-preview"
OUTPUT_PATH=/raid/.cache/huggingface/sayak/mochi-lora/
OUTPUT_PATH="mochi-lora"
cmd="accelerate launch --config_file deepspeed.yaml --gpu_ids $GPU_IDS text_to_video_lora.py \
cmd="CUDA_VISIBLE_DEVICES=$GPU_IDS python text_to_video_lora_simple.py \
--pretrained_model_name_or_path $MODEL \
--cast_dit \
--data_root $DATA_ROOT \
--seed 42 \
--mixed_precision "bf16" \
--output_dir $OUTPUT_PATH \
--train_batch_size 1 \
--dataloader_num_workers 4 \
--pin_memory \
--caption_dropout 0.1 \
--max_train_steps 2000 \
--checkpointing_steps 200 \
--checkpoints_total_limit 1 \
--gradient_accumulation_steps 4 \
--gradient_checkpointing \
--enable_slicing \
--enable_tiling \
--enable_model_cpu_offload \
--optimizer adamw --use_8bit \
--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\" \
--optimizer adamw \
--validation_prompt \"A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions\" \
--validation_prompt_separator ::: \
--num_validation_videos 1 \
--validation_epochs 1 \
--allow_tf32 \
--report_to wandb \
--nccl_timeout 1800"
--push_to_hub"
echo "Running command: $cmd"
eval $cmd